资讯动态

PyTorch unfold操作详解:从im2col到卷积与显存优化

发布时间:2026/10/3 1:01:03 来源:尧图企业网站定制
1. 从“手写卷积怎么写才不丢人”说起unfold解决的是访存问题如果让我在深度学习里挑一个“名字不太起眼但几乎所有视觉模型都绕不开”的底层操作我会选 unfold。很多朋友第一次看到torch.nn.functional.unfold时都会愣一下它到底在展开什么为什么要展展开之后又拿去干嘛这篇文章就专门把这件事讲透——包括 unfold 的数学形态、实际用法、和 Conv2d / Fold 的关系以及我自己踩过的维度错乱、显存爆炸的坑。适合刚学 CNN 不久、准备深度学习面试、或者想手写自定义算子的人阅读。先亮个观点unfold 本质上是一个数据搬运操作。它没有可学习参数也不做乘加运算只做切片和重排但它决定了后续所有矩阵运算的排列方式。你可以把它理解成“一条把大图按固定窗口切成小块的流水线”窗口怎么切、按什么顺序摆、最后放到什么形状里全部由 unfold 的参数决定。这个理解一旦建立很多困惑会自然消失。1.1 第一版卷积循环嵌套里藏着访存灾难如果你自己从零写卷积最自然的写法是什么遍历 batch、遍历输出通道、遍历输出高度和宽度、再遍历输入通道和卷积核内的 Kh×Kw 个点逐点相乘累加。这个版本在 CPU 上跑小图完全没问题逻辑清晰课程作业能交。但真要放到 GPU 上做大规模并行问题立刻暴露相邻输出位置用到的输入区域高度重叠每个线程的访存模式不规整数据在 shared memory 和寄存器之间的复用率很低。GPU 擅长的是“大量线程同时对连续整齐的数据做同样操作”而不是这种密集交叉的切片式读取。于是大家开始想能不能不把卷积当成一堆嵌套循环来做做法其实很老派把窗口搬出来让卷积变成矩阵乘法。这就是 im2colimage to column算法。先把每个滑动窗口内的元素按固定顺序拍成一个列向量所有窗口的列向量拼成一个大矩阵然后再用一个由卷积核权重重排成的矩阵去乘它。矩阵乘法可以交给 cuBLAS 这种把性能压到极致的库剩下的问题就变成“如何把数据搬得又快又整齐”。用一定内存开销换取并行效率是卷积底层实现最经典的取舍。1.2 unfold 就是框架化、可微分的 im2colPyTorch 的F.unfold就是对这个过程的官方封装。它接收一张(N, C, H, W)的特征图根据你给的 kernel_size、stride、dilation、padding把每个滑动窗口的内容取出来展平再按窗口索引排成一列最终输出一个形状为(N, C×Kh×Kw, L)的张量其中 L 是窗口总数。整个过程完全可微分反向传播时梯度能自动传回原图这一点对深度学习框架至关重要你可以在自定义网络层里放心使用 unfold不用担心梯度断掉。顺便回答一个很多人会问的“为什么不用 Python 循环”因为循环里的每个窗口切片操作在 GPU 上无法充分发挥并行能力而 unfold 把窗口切分变成一次底层的、高度优化的数据重排可以一次性搬运大量连续数据。你用循环写的逻辑和 unfold 完全一致但性能差好几个数量级。后面我会专门用代码验证“循环版本”和“unfold 版本”结果一致让大家在语义上彻底放心。2. 看懂 unfold 的输出形状维度变化背后是份“窗口台账”2.1 一个 4×4 的例子胜过十行定义先跑一个最小例子把输入设成连续的 0 到 15方便肉眼看窗口内容import torch import torch.nn.functional as F x torch.arange(16, dtypetorch.float32).reshape(1, 1, 4, 4) u F.unfold(x, kernel_size2, stride2) print(u.shape) # torch.Size([1, 4, 4])输入是(1, 1, 4, 4)的单通道矩阵[[ 0, 1, 2, 3], [ 4, 5, 6, 7], [ 8, 9, 10, 11], [12, 13, 14, 15]]用kernel_size2, stride2能切出 4 个不重叠的 2×2 窗口。输出形状是(1, 4, 4)这里的4需要拆成两部分理解第二维的C×Kh×Kw 1×2×2 4是单个窗口展平后的向量长度第三维的4是窗口个数。每个窗口按“从左到右、从上到下”的规则展平成列向量所以四列分别是第 0 列[0, 1, 4, 5]对应左上角窗口第 1 列[2, 3, 6, 7]对应右上角窗口第 2 列[8, 9, 12, 13]对应左下角窗口第 3 列[10, 11, 14, 15]对应右下角窗口。当通道数大于 1 时窗口内展平顺序也遵循“先沿高度方向、再沿宽度方向、最后跨通道”的规则。所以第二维的排列顺序是通道 0 的 Kh×Kw 区域、通道 1 的 Kh×Kw 区域依次往后。2.2 输出列数就是卷积输出尺寸的那套公式L的计算方式跟 Conv2d 的输出尺寸公式完全一致L_h floor((H 2*padding_h - dilation_h*(kernel_h - 1) - 1) / stride_h 1) L_w floor((W 2*padding_w - dilation_w*(kernel_w - 1) - 1) / stride_w 1) L L_h * L_w如果你设置了 padding公式里的padding也会影响窗口数量。比如输入(8, 8)、kernel_size3、stride2、padding1那么L_h L_w floor((8 2 - 2 - 1) / 2 1) floor(7 / 2) 1 4总共 16 个窗口。这个数字和F.conv2d在相同配置下的输出高宽完全一致所以 unfold 天然适合和卷积配对使用。我经常看到有人用 unfold 时随手填参数然后发现后续矩阵乘法的形状对不上其实大部分错误都可以通过先算一遍 L 来避免。2.3 为什么顺序是 (N, C×Kh×Kw, L)为了喂给矩阵乘法输出为什么不直接给(N, L, C×Kh×Kw)而是把通道和核大小合并后放在第二维这是历史包袱也是工程优化。考虑一个标准卷积权重是(Cout, Cin, Kh, Kw)把它重排成(Cout, Cin×Kh×Kw)之后可以直接和 unfold 输出的矩阵做一次批量矩阵乘法W_reshaped (Cout, Cin*Kh*Kw) col (Cin*Kh*Kw, L) - (Cout, L)col的每一列是一个窗口内存上Cin*Kh*Kw这一段是连续的正好对应矩阵乘法的 K 维。这种布局从 Caffe 时代就开始用专门为了方便后端 GEMM 库高效访问PyTorch 延续了这个约定。理解这个布局后你在做自定义算子或手写卷积时内存排布会清晰很多。3. 手写一遍 unfold代码和直觉对齐3.1 官方 API 的低层逻辑F.unfold的完整签名是torch.nn.functional.unfold(input, kernel_size, dilation1, padding0, stride1)input 通常是(N, C, H, W)。返回值是四维变三维(N, C×Kh×Kw, L)。注意它不像 Conv2d 那样自动把通道数映射到输出通道它只负责“切窗口”不负责“算特征”。真正的计算发生在你拿到窗口矩阵之后。3.2 用 Python 循环复现 unfold为了确认语义理解正确我用最原始的切片方式复现一遍H, W 4, 4 kh kw 2 stride 2 cols [] for i in range(0, H - kh 1, stride): for j in range(0, W - kw 1, stride): patch x[0, 0, i:ikh, j:jkw].reshape(-1) cols.append(patch) manual torch.stack(cols, dim1).unsqueeze(0) print(manual.shape) # torch.Size([1, 4, 4]) assert torch.allclose(manual, u), handwritten version should match unfold这个循环版本就是 unfold 最朴素的定义按照滑动窗口顺序把每个 patch 拍平然后作为列拼起来。唯一区别是 unfold 在底层实现里做了内存优化和并行化但对语义没有任何影响。当你心里对 unfold 产生怀疑时写一个这样的小循环去对照永远是最快的验证方式。3.3 用 unfold 复现一次 Conv2d下面用 unfold 完整复现一个Conv2d包括 padding 和 stridex torch.randn(2, 3, 8, 8) conv torch.nn.Conv2d(3, 5, kernel_size3, stride2, padding1, biasFalse) with torch.no_grad(): w conv.weight # (5, 3, 3, 3) cols F.unfold(x, kernel_size3, stride2, padding1) # cols shape: (2, 3*3*3, 4*4) (2, 27, 16) w_flat w.reshape(5, -1) # (5, 27) out torch.matmul(w_flat, cols) # (2, 5, 16) out out.reshape(2, 5, 4, 4) out_conv conv(x) print(torch.allclose(out, out_conv, atol1e-5)) # True输出跟官方卷积完全一致。这解释了为什么有时候我们称卷积本质是“GEMM 加上一次 unfold”权重被展开成矩阵输入也被展开成矩阵剩下的就是一次大矩阵乘法。实际的高性能卷积库不一定真的把 unfold 显式落盘他们会用“隐式 GEMM”的方式把取窗口的操作融合进计算核心里但原理上仍然跟 unfold 同源。4. fold 操作展开之后总要有人收拾残局4.1 fold 把一个 patch 矩阵填回原图重叠区域无情累加F.fold可以看作 unfold 的逆过程把(N, C×Kh×Kw, L)的窗口矩阵填回(N, C, H, W)的特征图。但它不是严格的数学逆运算而是“散点累加”x torch.ones(1, 1, 4, 4) u F.unfold(x, kernel_size2, stride1) # (1, 4, 9) y F.fold(u, output_size(4, 4), kernel_size2, stride1) print(y)输出看起来是这样tensor([[[[1., 2., 2., 1.], [2., 4., 4., 2.], [2., 4., 4., 2.], [1., 2., 2., 1.]]]])中间位置的元素被四个窗口同时覆盖fold 会把所有覆盖它的 patch 内容加起来所以结果是 4。边缘位置覆盖次数少所以是 2 或 1。fold 不会自动做平均它只做累加。这个行为继承自信号处理里的 overlap-add 方法在重建窗口化信号时很常用但如果你以为它是 unfold 的严格逆运算就会得到意料之外的结果。4.2 想取平均怎么办维护一个计数矩阵如果你想做“滑窗统计后还原”比如对每个窗口求均值再拼回原图那一定要自己维护一个归一化矩阵count F.fold( F.unfold(torch.ones_like(x), kernel_size2, stride1), output_size(4, 4), kernel_size2, stride1, ) avg y / count用全 1 输入走一遍 unfold 再 fold得到的就是每个位置的窗口覆盖次数。把累加结果除以这个覆盖次数就得到逐元素的平均。这个小技巧在处理重叠窗口的局部统计时非常实用我在做滑动窗口类算法时基本每次都用到。4.3 unfold 和 fold 的梯度就是互相传scatter_add 是核心很多人在自定义层里不敢用 unfold担心梯度传播出问题。其实完全不用担心F.unfold的反向传播本质是把输出梯度按照同样的窗口位置累加回输入这恰好就是一次F.fold而F.fold的反向传播本质是把输出梯度按窗口切分出去这恰好就是一次F.unfold。框架内部用类似 scatter_add 的方式实现“从哪里来回哪里去”。如果你自己写低层算子梯度回传最难的部分往往是“把梯度按位置加回去”而不是矩阵求导。这也是我建议用 gradcheck 验证自定义算子的原因后面会详细说。5. 实战视角这些模型结构里全有 unfold 的影子5.1 写自定义池化或局部算子时的“万能底板”当你需要的不是标准卷积而是某种自定义的局部算子时unfold 是很好的底板。比如局部响应归一化、局部标准差、局部直方图特征都可以先把窗口切出来然后在L维或窗口特征维上做任意操作。举个例子用 unfold 实现最大池化x_pad F.pad(x, (1, 1, 1, 1), valuefloat(-inf)) patches F.unfold(x_pad, kernel_size3, stride2) max_out patches.max(dim1).values max_out max_out.reshape(2, 5, 4, 4) # 前提是你提前算好输出尺寸这种写法比手写循环清晰得多而且如果你想在池化过程里保存每个窗口最大值的位置索引只需要额外对patches.argmax(dim1)做一个反推灵活度很高。5.2 Patch Embedding 的两种写法完全等价Vision Transformer 里的 patch embedding 是 unfold 思想最典型的现代应用。通常大家用nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size)实现因为它快、显存友好。但从数学上看它等价于先 unfold再做一次线性投影patches F.unfold(x, kernel_size16, stride16) # (B, C*256, L) patches patches.transpose(1, 2) # (B, L, C*256) embedding linear_proj(patches) # (B, L, embed_dim)Conv2d的每个输出位置本质上就是对输入 patch 做一次线性组合组合系数就是卷积核权重。所以卷积实现 patch embedding 和“unfold Linear”之间只是计算路径不同数学表达完全等价。理解这一点后你就能解释为什么改 ViT 时有人用 Conv2d、有人用 unfold还能自由地在两种写法之间切换。5.3 一维信号和时间序列里的滑窗样本unfold 不只属于图像。做时间序列预测时最常规的操作就是把一段长信号切成很多定长窗口作为训练样本。PyTorch 的Tensor.unfold方法可以直接沿某一维切窗口x torch.arange(10, dtypetorch.float32) windows x.unfold(0, 4, 2) # shape (4, 4)注意这里和F.unfold的返回形状不同Tensor.unfold沿指定维度增加一个窗口大小的尾维返回的是原张量的视图不会复制数据。它更适合快速预览和轻量级滑窗而F.unfold返回的是显式重排后的矩阵适合跟矩阵乘法衔接。两者底层思想同源但使用场景不同。6. 认真踩坑维度顺序、padding组合和显存一个都不能少6.1 padding 和 dilation 都在参数里但输出尺寸很容易算错F.unfold本身支持 padding 和 dilation并不需要你手动 pad。但很多人会忽略一件事padding 会影响L的大小而且 fold 时的 padding 必须和 unfold 保持一致否则填回的位置会错位。我的建议是凡是涉及 unfold/fold 成对使用的代码先把公式抄在注释里再用一个小 shape 手动验证一遍不然后续矩阵乘法的形状错误会非常难查。6.2 transpose 之后才能 reshape顺序错一个维度后面全是噪声这是我在实际代码里遇过最多的问题。假设你想把(N, C×Kh×Kw, L)转成(N, L, C×Kh×Kw)接一个全连接层# 正确写法 patches patches.transpose(1, 2) # 错误写法看起来也是 (N, L, D)但元素顺序完全错了 patches patches.reshape(N, L, -1)原因在于 reshape 按内存顺序重新解释数据。原张量的内存顺序是先填满第二维C×Kh×Kw再切到第三维的下一个位置直接 reshape 会把第二个窗口的前几个元素和第一个窗口的后几个元素混在一起。所以遇到这种需求一律先transpose或者permute确认维度顺序后再 reshape。6.3 显存是 unfold 的最大敌人先算账再动手unfold 最大的代价是显存。举一个极端例子输入1×64×512×512的 float32 特征图本身约占 64MB。如果 kernel_size7、stride1、padding3输出窗口数量是 512×512每列长度是 64×49 3136展开后约 8.2 亿个元素占 3.3GB 显存。同样的特征图直接过卷积可能只需要几十 MB 的临时空间因为底层库可以走隐式 GEMM不真正展开。遇到显存不够时最直接的手段是分 batch 处理把输入按 batch 维度切块每个小块做 unfold 和后续计算再拼接结果。其次是可以考虑用卷积替代“unfold 线性变换”让底层库自动选择更优的实现路径。如果必须用 unfold那就在写代码之前先算一算展开后的张量大小提前做好切块计划。6.4 gradcheck怀疑前向或反向写错时让它替你体检如果你基于 unfold 做了自定义封装或者手动实现了类似逻辑强烈建议用torch.autograd.gradcheck验证一遍x torch.randn(1, 3, 8, 8, dtypetorch.float64, requires_gradTrue) def func(t): return F.unfold(t, kernel_size3, stride2, padding1) torch.autograd.gradcheck(func, x, eps1e-6)gradcheck 会用数值差分去近似雅可比矩阵跟你的反向传播结果做对比。它能同时验证前向和反向是否正确是排查自定义算子问题的第一工具。我在写一些临时的滑窗逻辑时也会拿它快速验证自己的包装没有破坏梯度流省下不少调试时间。最后说一个跟我工作习惯有关的小事我在自己的代码里总是先写一行注释“unfold 之后第二维是 C×Kh×Kw第三维是空间位置”每次 reshape 前先打印 shape再动手。这句话帮我避免了很多次静默出错。如果这篇文章只能让你记住一句话我希望是unfold 是展平滑动窗口fold 是散点累加窗口重叠时 fold 累加而不是平均。记牢这一句大部分坑都不会再踩。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑