1. 从一个需求说起为什么需要unfold()先说一个我去年处理数据时遇到的实际问题。当时在做一个时间序列的预测模型原始数据是一维的股价序列长度几万步。模型需要把每60步的历史窗口作为一个样本喂进去每个样本还要和下一个时刻的真实值对应起来。如果用最笨的for循环切片几万个样本跑下来光是数据预处理就得好几分钟而且在DataLoader里每次都要重新切片训练时GPU经常等CPU喂数据气得我想砸键盘。后来翻文档时注意到Tensor.unfold()这个API试了一下整个预处理从几分钟压缩到了几百毫秒而且代码干净得不像话。从那以后凡是要做滑窗、分块、提取局部区域的操作我第一个想到的就是unfold()。这个函数其实做的事情非常单纯沿着指定的维度用固定大小的窗口以固定的步长滑动把窗口里的元素抽出来拼成一个新的维度。本质上就是滑动窗口切片的向量化实现。它不涉及任何梯度计算以外的特殊逻辑纯PyTorch原生支持在CPU和GPU上都能高效运行。如果你接触过卷积神经网络会发现unfold()和卷积操作天然有联系——卷积本质上就是在做unfold再加上乘法和加法。理解了unfold()你对Conv2d底层原理的理解也会上一个台阶。这篇文章适合以下几类读者做时间序列滑窗特征提取的做图像patch划分比如ViT、Swin Transformer的前处理的想深入理解卷积底层实现的以及所有被for循环数据预处理折磨过的PyTorch用户。我会从函数签名讲起对比各种常见用法再给出几个可以直接抄的实战例子。2. unfold()的官方定义与核心理解2.1 看一下函数签名在PyTorch文档里unfold()是Tensor的一个方法签名长这样Tensor.unfold(dimension, size, step) - Tensor三个参数的含义很直白dimension要在哪个维度上滑动窗口整数值可以是负数索引和Python索引规则一致。size窗口大小正整数表示每次取多少个元素。step滑动步长正整数表示窗口每次往后挪多远。返回值是一个新的张量形状会发生变化。关键点在于unfold()是在指定维度上新增一个维度用来放窗口内容。假设输入tensor在第dimension维的长度是L那么输出在该维的长度变为输出长度 floor((L - size) / step) 1前提是L size否则窗口放不下会直接报错。这个公式和卷积输出尺寸的计算公式是一模一样的后面我会专门解释它们的关系。2.2 一个最简单的一维例子先看最基础的一维情况import torch x torch.arange(10) # tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9]) y x.unfold(0, 4, 2) print(y) print(y.shape)输出结果tensor([[0, 1, 2, 3], [2, 3, 4, 5], [4, 5, 6, 7], [6, 7, 8, 9]]) torch.Size([4, 4])我来把执行过程拆解开一维tensor长度为10窗口大小4步长2。第一次窗口取索引0~3得到[0,1,2,3]然后窗口右移2步取索引2~5得到[2,3,4,5]再右移2步取索引4~7最后取索引6~9。按公式计算(10-4)/2 1 4正好产生4个窗口。这里输出变成了二维张量第二个维度最后一维就是窗口内容维度。换句话说在哪个维度上unfold那个维度就会被“拆掉”新生成的窗口维度会追加到张量的最后。这个“新维度在最后”的特性非常关键尤其在做多维张量操作时很多人就是因为没搞懂这一点而踩坑。2.3 多维张量时的维度变化规律二维以上的情况稍微复杂一些但规律其实很简单。我们用一个3x5的矩阵来演示x torch.arange(15).reshape(3, 5) print(x)这个矩阵长这样tensor([[ 0, 1, 2, 3, 4], [ 5, 6, 7, 8, 9], [10, 11, 12, 13, 14]])先沿着第0维行方向unfold窗口大小2步长1y x.unfold(0, 2, 1) print(y.shape) # torch.Size([2, 5, 2])结果维度变到了3维原第0维长度3变成了2窗口维度2追加到了最后。具体含义是从行方向看每次取两行窗口高度2步长1一共得到(3-2)/112个窗口每个窗口是2x5的矩阵。再沿着第1维列方向unfold窗口大小2步长2z x.unfold(1, 2, 2) print(z.shape) # torch.Size([3, 2, 2])原第1维长度5变成了2窗口维度2追加到最后。一共有(5-2)/212个窗口每个窗口是3x2的矩阵。如果想同时对行和列做滑窗可以连续调用w x.unfold(0, 2, 1).unfold(1, 2, 1) print(w.shape) # torch.Size([2, 4, 2, 2])这里先对第0维unfold得到[2,5,2]再对当前张量的第1维也就是原第1维unfold窗口大小2步长1得到[2,4,2,2]。这实际上就是在做图像patch提取——把3x5的矩阵拆成了若干2x2的小块行方向2个列方向4个总共8个patch。记住这句话unfold()永远只影响你指定的那个维度其他维度保持不变新窗口维度统一追加到末尾。这是理解多维unfold的钥匙。3. 关键点一unfold()与卷积/im2col的等价关系3.1 输出长度公式为什么这么眼熟如果你写过CNN一定见过卷积输出尺寸公式输出尺寸 floor((输入尺寸 2padding - dilation(kernel_size-1) - 1) / stride 1)当padding0、dilation1时这个公式简化为输出尺寸 floor((输入尺寸 - kernel_size) / stride 1)这跟unfold()的输出长度公式完全一致窗口大小对应kernel_size步长对应stride。所以可以认为unfold()就是在做卷积的第一步im2col。卷积操作实际可以拆成两个阶段第一阶段把输入特征图按卷积核大小和步长切分成无数个小块im2col这一阶段不涉及任何数值计算只是数据重组第二阶段把每个小块和卷积核做逐元素乘法并求和得到输出特征图上的一个值。unfold()就是第一阶段的标准实现。3.2 用一个具体例子验证假设有一张1通道、5x5的输入图卷积核3x3步长1。用unfold()可以这样提取所有滑动窗口x torch.arange(25, dtypetorch.float32).reshape(1, 1, 5, 5) # 将每个3x3的窗口展平成9维向量 unfolded x.unfold(2, 3, 1).unfold(3, 3, 1) print(unfolded.shape) # torch.Size([1, 1, 3, 3, 3, 3]) # 合并窗口内部的维度得到 [batch, C, num_h, num_w, k*k] unfolded unfolded.contiguous().view(1, 1, 3, 3, 9) print(unfolded.shape) # torch.Size([1, 1, 3, 3, 9])这里连续对第2维高度H和第3维宽度W做了unfold窗口大小3步长1得到[1,1,3,3,3,3]。最前面两个3是输出特征图的空间网格5-313最后两个3是窗口内部的3x3内容。最后把窗口内部的3x3展平成9维就得到了im2col的标准形式。如果你用PyTorch的F.unfold函数一步就能得到类似的结果import torch.nn.functional as F unfolded F.unfold(x, kernel_size3, stride1) print(unfolded.shape) # torch.Size([1, 9, 9])F.unfold返回的形状是[batch, C*k*k, L]其中L是所有窗口的总数量这里是3x39。两种方式殊途同归但底层逻辑完全一致。3.3 为什么理解这个关系有价值理解了unfold与卷积的关系有几个实际的收益。第一手写自定义卷积层时可以直接用unfold加速。如果你需要实现一个非标准卷积操作——比如局部连接层、不需要共享权重的卷积、或对每个窗口应用不同操作的层unfold()比F.conv2d灵活得多。第二调试卷积行为时多了一个可视化工具。遇到卷积结果和预期不符的情况可以用unfold()把窗口拆出来逐个检查数据流向定位问题比对着特征图瞎猜高效得多。第三理解许多前沿模型的数据流。ViT、Swin Transformer等视觉Transformer模型的第一步都是把图像切成patch很多实现直接用x.unfold(2, patch_size, patch_size).unfold(3, patch_size, patch_size)看懂unfold这些代码在你眼里就变成了透明的。4. 关键点二窗口维度永远在最后的坑与对策4.1 一个典型的踩坑场景前面提到过unfold()新增的窗口维度会放在张量的最后面。这在某些场景下会导致维度顺序不符合预期。举个例子。假设你有一个[B, T, D]形状的时间序列批量数据B是batch大小T是时间步数D是每个时间步的特征维度。你想沿着时间维度提取长度为W的滑窗步长为Sx torch.randn(4, 100, 8) # B4, T100, D8 y x.unfold(1, 10, 5) print(y.shape) # torch.Size([4, 19, 8, 10])这里窗口维度最后一个10在最后而特征维度8变成了倒数第二维。如果你期望的结果是[B, num_windows, W, D]也就是窗口内的时间步在特征前面这个维度顺序就和预期不一致。这时候需要手动用permute调整维度顺序y y.permute(0, 1, 3, 2) # [4, 19, 10, 8]或者如果你之后要把窗口数据送入某些期望[B, W, D]的模块就得先合并batch和窗口维度y y.permute(0, 1, 3, 2).reshape(-1, 10, 8)4.2 什么时候窗口维度在最后反而方便虽然维度顺序有时候需要调整但窗口维度在最后也有它的便利之处。最典型的是图像patch提取。一张[B, C, H, W]的图片连续对H和W做unfold之后得到[B, C, num_h, num_w, patch_h, patch_w]。如果你想把每个patch展平成一维向量可以直接把最后两个维度合并patches x.unfold(2, patch_h, patch_h).unfold(3, patch_w, patch_w) # [B, C, num_h, num_w, patch_h, patch_w] patches patches.reshape(B, C, num_patches, patch_h * patch_w)这里不需要permute因为展平的操作天然就是把相邻维度合并窗口维度在最末尾反而让这一步变得非常自然。再比如用unfold()实现局部平均池化的可视化验证窗口维度在最后直接对最后一维求均值就行不需要关心窗口内容在哪个维度。4.3 批量滑窗时的内存优化经验使用unfold()最需要注意的一个实际问题是内存开销。unfold()虽然避免了for循环但它会显式地把每个窗口的数据复制出来存储开销大约是窗口个数 × 窗口大小 × 每个元素的字节数。举个例子一个[128, 128]的矩阵用窗口大小[32, 32]步长[16, 16]去unfold会产生(128-32)/161 7行窗口和7列窗口一共49个窗口每个窗口32x321024个元素。展平后就是49 x 1024的矩阵原数据是16384个元素unfold之后是50176个元素放大了约3倍。这个放大倍数在某些极端场景下很吓人。如果输入的sar值太大会崩内存尤其当你在高分辨率图像上切大量小patch时。我的经验是如果窗口数量多到资源吃紧宁可把数据切成多个batch分批处理不要一次性全unfold否则一个数据预处理就能把整个训练进程干崩。另一个内存优化技巧是如果后续只需要窗口的聚合结果比如均值、最大值用unfold()展开后再聚合不是最佳方案。因为展开过程的中间张量会占用额外内存而直接使用F.avg_pool或F.max_pool这类操作可以在不显式展开窗口的情况下完成计算内存效率高得多。unfold适合的是“需要显式看到每个窗口内容”的场景如果只是做池化用池化API才是正解。5. 拿unfold()做时间序列滑窗的实战演练5.1 需求描述与数据准备回到文章开头说的那个场景一维时间序列长度为10000需要构造训练样本。每个样本是连续60个时间步的历史数据预测下一个时间步的值。步长设为20也就是样本之间允许部分重叠数据利用率更高。数据准备要解决两个问题一是生成输入矩阵X形状是[num_samples, 60]二是生成标签向量y形状是[num_samples]对应每个窗口之后的一个值。先造一份模拟数据import torch # 模拟10000步的时间序列 data torch.sin(torch.linspace(0, 100, 10000)) torch.randn(10000) * 0.1 # 参数 window_size 60 step 20 # 用unfold()构造滑窗 samples data.unfold(0, window_size, step) print(samples.shape) # torch.Size([498, 60])计算一下样本数量(10000 - 60) / 20 1 498正好。标签也很容易构造每个样本预测窗口后的下一个值也就是索引从window_size开始每隔step取一个值labels data[window_size:].unfold(0, 1, step).squeeze(1) # 或者直接用切片索引 labels data[window_size::step] # 更简洁等等这里有个细节需要注意。如果用data[window_size::step]得到的是从第60个位置开始每隔20步取一个值一共(10000 - 60)/20 1 497个值——不对应该仔细算一下。索引0~59是第一个窗口它的标签是索引60第二个窗口索引20~79标签是索引80……第498个窗口的起始索引是20*4979940窗口范围9940~9999标签是索引10000但索引10000越界了。所以最后一个窗口没有对应标签正确的做法是只取前497个窗口# 先算出能完整构造训练对的窗口数量 num_samples (data.shape[0] - window_size) // step # 497 samples data.unfold(0, window_size, step)[:num_samples] labels data[window_size::step][:num_samples] print(samples.shape, labels.shape) # torch.Size([497, 60]) torch.Size([497])这算是一个很容易踩的细节当序列长度不能被步长整除时unfold()会保留最后一个完整的窗口但这个窗口之后可能没有足够的后续数据来构造标签。所以要么数据处理时提前裁剪要么构造标签时同步对齐长度。5.2 借助unfold()构造多维特征的扩展如果时间序列不只是单变量而是多变量的比如同时记录价格、成交量、持仓量数据形状是[T, D]unfold()依然能优雅处理data_multi torch.randn(10000, 5) # T10000, D5 # 沿时间维度滑窗 samples data_multi.unfold(0, window_size, step) print(samples.shape) # torch.Size([498, 5, 60])注意窗口维度跑到最后了形状是[num_windows, D, window_size]。如果你期望的形状是[num_windows, window_size, D]需要permutesamples samples.permute(0, 2, 1) # [498, 60, 5]这个形状在送入LSTM或Transformer等序列模型时更加自然——第一维是batch第二维是序列长度第三维是输入特征。5.3 结合DataLoader的完整流程把上面的处理组合进完整的数据加载流程可以写出这样的代码from torch.utils.data import Dataset, DataLoader class TimeSeriesDataset(Dataset): def __init__(self, data, window_size60, step20): self.data data self.window_size window_size self.step step # 用unfold一次算好所有窗口 self.samples data.unfold(0, window_size, step) # 计算有效样本数保留有标签的部分 self.num_samples (data.shape[0] - window_size) // step self.samples self.samples[:self.num_samples] self.labels data[window_size::step][:self.num_samples] def __len__(self): return self.num_samples def __getitem__(self, idx): return self.samples[idx], self.labels[idx] dataset TimeSeriesDataset(data, window_size60, step20) dataloader DataLoader(dataset, batch_size32, shuffleTrue) # 测试一批数据 for x_batch, y_batch in dataloader: print(x_batch.shape, y_batch.shape) break实测下来data.unfold(0, window_size, step)这一行就能替代原本几十行的for循环切片逻辑而且因为底层是连续内存的view操作加上可能的copy速度提升非常明显。5.4 内存占用实测对比我在一台16GB内存的机器上测试过序列长度100万窗口大小60步长20直接unfold出来的中间张量约49997 x 60个float元素约12MB完全无压力。但如果步长设为1窗口数量接近100万中间张量就变成了999941 x 60约240MB会明显感觉到内存上升。所以做高重叠滑窗时优先级应该是先考虑stride大一点如果必须低步长考虑用as_strided更底层、更省内存但更危险再不行就分批处理。unfold()虽好也不是万能药。6. 拿unfold()做图像Patch化与反向恢复6.1 图像patch化的两种写法对比视觉Transformer类的模型第一步几乎都是把图片切成固定大小的patch。假设输入是[B, C, H, W]patch大小是16x16步长16不重叠。用unfold()实现B, C, H, W 2, 3, 64, 64 x torch.randn(B, C, H, W) patch_size 16 # 方法一连续两次unfold patches x.unfold(2, patch_size, patch_size).unfold(3, patch_size, patch_size) print(patches.shape) # [2, 3, 4, 4, 16, 16] # 方法二用F.unfold一步到位 import torch.nn.functional as F patches_f F.unfold(x, kernel_sizepatch_size, stridepatch_size) print(patches_f.shape) # [2, 768, 16]7683*16*16164*4个patch方法一得到的是保留空间网格结构的形式更直观方法二得到的是展平后的形式更方便直接拼进Transformer。两种写法在不同实现中都有使用。如果是重叠patch窗口大小大于步长只需要把stride参数改一下逻辑完全一样。比如Swin Transformer的patch merging之前常常会有重叠窗口划分x.unfold(2, 8, 4).unfold(3, 8, 4)就能切出步长4、窗口8的重叠patch。6.2 Patch的展平与维度重排从方法一的[B, C, num_h, num_w, ph, pw]转换到Transformer输入需要的[B, num_patches, C*ph*pw]需要两步操作B, C, H, W 2, 3, 64, 64 patch_size 16 x torch.randn(B, C, H, W) patches x.unfold(2, patch_size, patch_size).unfold(3, patch_size, patch_size) # [B, C, num_h, num_w, ph, pw] # 先把窗口内部展平让每个patch变成一维向量 patches patches.reshape(B, C, -1, patch_size * patch_size) # [B, C, num_patches, ph*pw] # 调整维度顺序到 [B, num_patches, C*ph*pw] patches patches.permute(0, 2, 1, 3).reshape(B, -1, C * patch_size * patch_size) print(patches.shape) # torch.Size([2, 16, 768])注意这里的reshape和permute顺序先reshape把[ph, pw]合并成ph*pw再permute把patch数量维度提前最后reshape统一合并通道和窗口内容。顺序搞反会导致数据错乱这是很多初学者会犯的错误。实际调试时我会先做一个只有4个像素的小图比如1x1x2x2手动算出期望值再用同样的过程跑一遍对比输出确认维度变换逻辑正确后才能应用到真实数据上。6.3 从patch还原图像折叠操作有拆就有合。把patch重新拼回完整图像PyTorch提供了F.fold与F.unfold是对偶操作patches_f F.unfold(x, kernel_sizepatch_size, stridepatch_size) # [2, 768, 16] # 还原成原图 x_reconstructed F.fold(patches_f, output_size(H, W), kernel_sizepatch_size, stridepatch_size) print(x_reconstructed.shape) # torch.Size([2, 3, 64, 64])如果patch之间没有重叠fold可以完全无损还原原图。如果patch之间有重叠情况会变复杂重叠区域会有多个patch的值相加还原时需要除以重叠计数矩阵这个过程PyTorch的fold已经自动处理了实际上它默认就是累加再归一化。对于带重叠的滑动窗口重建我一般不建议直接用fold除非你清楚累加语义否则容易得到意料之外的结果。6.4 一些图像相关的使用心得在ViT类模型的前处理中用F.unfold比手动reshape更加通用因为它天然支持重叠、空洞等复杂采样模式。如果你的patch划分是标准的非重叠网格直接reshape其实更快因为不产生额外的窗口复制# 非重叠patch的快速写法不经过unfold x x.reshape(B, C, H // patch_size, patch_size, W // patch_size, patch_size) x x.permute(0, 2, 4, 3, 5, 1) # [B, num_h, num_w, ph, pw, C] x x.reshape(B, -1, patch_size * patch_size * C)这个方法比unfold快因为它只是改变步长和视图没有数据复制。但缺点是一旦patch有重叠或窗口不规整这个方法就失效了。所以两条路都掌握按需选择就行。7. unfold()的进阶操作与原理补充7.1 与as_strided的对比PyTorch里还有一个底层函数叫as_strided也能实现类似的滑动窗口提取。相比unfold()as_strided允许你直接修改张量的stride元数据完全不复制数据因此内存开销几乎为零。以相同的时间序列滑窗为例data torch.arange(100, dtypetorch.float32) window_size 10 step 5 # 使用as_strided samples data.as_strided( size((100 - 10) // 5 1, 10), stride(5, 1) ) print(samples.shape) # torch.Size([19, 10])注意这个samples和原数据data共享底层内存。好处是省内存速度极快坏处是修改samples会影响原数据而且一旦维度或步长算错可能会访问到不安全的内存区域导致难以排查的bug。我的建议是如果是生产代码或对安全性要求高的场景优先用unfold()如果数据量大到unfold()内存吃紧且你清楚自己在做什么可以用as_strided做极致优化。对绝大多数人来说unfold()是性能和安全的平衡点。7.2 梯度流与反向传播unfold()是一个线性操作本身是可导的。这意味着你可以把unfold()用在网络的前向传播中梯度可以正常回传。举个例子如果你想实现一个“可学习的patch分组层”可以先unfold再对每个patch做加权求和x torch.randn(2, 3, 16, 16, requires_gradTrue) patches x.unfold(2, 8, 4).unfold(3, 8, 4) print(patches.shape) # [2, 3, 3, 3, 8, 8] # 假设每个patch学习一个权重 weight torch.randn(3, 3, 8, 8, requires_gradTrue) # 简化对patch内部求加权和广播相乘再求和 out (patches * weight.unsqueeze(0).unsqueeze(2)).sum(dim(-2, -1)) out.mean().backward() print(x.grad is not None) # True在实际使用反向传播时完全不用自己手写梯度PyTorch的autograd引擎会正确处理unfold()及其逆操作如果有的话的梯度传播。这为在神经网络内部使用滑动窗口操作铺平了道路。7.3 one-hot编码与unfold结合的小技巧在某些NLP或序列处理场景里需要对one-hot向量做局部n-gram提取。unfold()在这类特征工程中也很顺手vocab_size 10 seq torch.tensor([1, 3, 5, 2, 4, 6, 8, 0]) one_hot F.one_hot(seq, num_classesvocab_size).float() print(one_hot.shape) # [8, 10] # 提取连续3个位置的one-hot拼起来作为特征 ngram one_hot.unfold(0, 3, 1) print(ngram.shape) # [6, 10, 3] # 展平成 [6, 30] ngram_flat ngram.reshape(6, -1)这个技巧在传统NLP/序列特征工程里非常实用一两行代码就能把整个语料的n-gram特征矩阵构造出来比for循环一个个拼接快几个数量级。7.4 当size或step参数不合法时的报错unfold()对参数有一点基本要求不满足时会直接抛错。常见的报错场景有这么几种窗口大小大于该维度长度比如一个长度为10的一维张量unfold(0, 12, 1)会报RuntimeError: maximum size for tensor at dimension。窗口大小或步长为0或负数会报RuntimeError: size/steps must be positive。在维度为0的张量上做unfold很明显没法切窗口。遇到这些报错不要慌先检查参数是否符合公式L size再确认step 0基本都能解决。另外dimension参数支持负数索引x.unfold(-1, size, step)等价于x.unfold(最后一个维度, size, step)这在写通用工具函数时很好用。8. 常见问题与避坑清单8.1 汇总容易踩的坑根据我自己的经历以及帮别人review代码时见到的各种问题这里整理一份高频问题清单问题现象解决方案维度顺序混乱窗口维出现在末尾后续操作报错或结果错乱用permute把窗口维度移到期望的位置或者在代码里先打印shape确认最后一个窗口没有对应标签时间序列滑窗后标签长度比样本少1个构造标签时按(L - window_size) // step对齐长度内存暴涨窗口数量大时中间张量占用太大增大step、分批处理、或改用as_strided与reshape混用时数据错乱permute和reshape顺序不对导致数据串位牢记reshape先合并相邻维度permute先改变维度顺序顺序不能反输出长度比预想少忘了公式里要1用手动小例子验算或先打印L-size//step 1在batch多个样本时滑动步长超出预期把step传成了窗口总数或者传错参数用step1的小测试确认窗口数量和内容8.2 调试unfold()的实用技巧当你对某次unfold结果不确定时我建议用这个方法快速验证造一个元素值有规律的小张量比如torch.arange(25).reshape(5,5)然后做unfold并打印结果。因为元素值是0~24递增的每个窗口的值一眼就能看出来是不是正确切分。还有一个小技巧用等价写法验证。对于一维张量xx.unfold(0, size, step)[i]等价于x[i*step : i*stepsize]。如果你对第i个窗口的内容有疑问可以直接用切片取出来对比这样能把抽象的维度操作变成看得见摸得着的索引操作。8.3 性能调优建议在处理超大张量时unfold()的性能一般不是瓶颈瓶颈通常出现在后续的reshape和内存拷贝上。以下是我实测有效的一些优化手段如果多次对同一数据做不同size的滑窗考虑只做一次unfold然后在这个结果上做切片或reshape避免重复复制数据。如果窗口数据之后要喂进卷积层优先考虑直接调用F.conv2d让底层优化过的kernel来干活unfold只适合需要显式窗口内容的场景。在GPU上unfold的速度提升比CPU更明显因为数据并行度高所以如果数据量大建议把数据先搬到GPU再做预处理。9. 进一步扩展用unfold()实现自定义滑窗算子如果你已经掌握了上面的内容可以试试一个综合练习用unfold()实现一个不调用F.conv2d的二维平均池化。def avg_pool_with_unfold(x, kernel_size, stride): # x: [B, C, H, W] B, C, H, W x.shape # unfold出所有窗口 windows x.unfold(2, kernel_size, stride).unfold(3, kernel_size, stride) # [B, C, out_h, out_w, k, k] pooled windows.mean(dim(-2, -1)) return pooled x torch.randn(1, 3, 8, 8) out avg_pool_with_unfold(x, 2, 2) print(out.shape) # torch.Size([1, 3, 4, 4]) # 验证结果与F.avg_pool一致 import torch.nn.functional as F out_ref F.avg_pool2d(x, kernel_size2, stride2) print(torch.allclose(out, out_ref)) # True这个例子虽然简单但把unfold()的核心能力都串起来了滑窗、维度变换、聚合计算。在此基础上你可以继续扩展出自定义局部归一化、局部对比度增强等操作。在我的实际工作中有好几个“非标准”的网络层都是靠unfold()拼出来的——比如不需要共享权重的局部全连接层比如带重叠区域的局部特征融合模块。每次用unfold()替代手工for循环代码缩短一大半运行速度提升几十倍这种爽感可能只有被数据预处理折磨过的人才能真正体会。我个人的经验是凡是遇到“滑动窗口”这四个字先想想能不能用unfold()解决这几乎已经成了我的条件反射。希望这篇文章也能让你在遇到类似需求时少走一些弯路。