☰
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 擅长的是“大量线程同时对连续整齐的数据做同样操作”,而不是这种密集交叉的切片式读取。

于是大家开始想,能不能不把卷积当成一堆嵌套循环来做?做法其实很老派:把窗口搬出来,让卷积变成矩阵乘法。这就是 im2col(image to column)算法。先把每个滑动窗口内的元素按固定顺序拍成一个列向量,所有窗口的列向量拼成一个大矩阵,然后再用一个由卷积核权重重排成的矩阵去乘它。矩阵乘法可以交给 cuBLAS 这种把性能压到极致的库,剩下的问题就变成“如何把数据搬得又快又整齐”。用一定内存开销换取并行效率,是卷积底层实现最经典的取舍。

1.2 unfold 就是框架化、可微分的 im2col

PyTorch 的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, dtype=torch.float32).reshape(1, 1, 4, 4) u = F.unfold(x, kernel_size=2, stride=2) 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_size=2, stride=2,能切出 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_size=3、stride=2、padding=1,那么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, dilation=1, padding=0, stride=1)

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:i+kh, j:j+kw].reshape(-1) cols.append(patch) manual = torch.stack(cols, dim=1).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 和 stride:

x = torch.randn(2, 3, 8, 8) conv = torch.nn.Conv2d(3, 5, kernel_size=3, stride=2, padding=1, bias=False) with torch.no_grad(): w = conv.weight # (5, 3, 3, 3) cols = F.unfold(x, kernel_size=3, stride=2, padding=1) # 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, atol=1e-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_size=2, stride=1) # (1, 4, 9) y = F.fold(u, output_size=(4, 4), kernel_size=2, stride=1) 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_size=2, stride=1), output_size=(4, 4), kernel_size=2, stride=1, ) 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), value=float('-inf')) patches = F.unfold(x_pad, kernel_size=3, stride=2) max_out = patches.max(dim=1).values max_out = max_out.reshape(2, 5, 4, 4) # 前提是你提前算好输出尺寸

这种写法比手写循环清晰得多,而且如果你想在池化过程里保存每个窗口最大值的位置索引,只需要额外对patches.argmax(dim=1)做一个反推,灵活度很高。

5.2 Patch Embedding 的两种写法完全等价

Vision Transformer 里的 patch embedding 是 unfold 思想最典型的现代应用。通常大家用nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)实现,因为它快、显存友好。但从数学上看,它等价于先 unfold,再做一次线性投影:

patches = F.unfold(x, kernel_size=16, stride=16) # (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, dtype=torch.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_size=7、stride=1、padding=3,输出窗口数量是 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, dtype=torch.float64, requires_grad=True) def func(t): return F.unfold(t, kernel_size=3, stride=2, padding=1) torch.autograd.gradcheck(func, x, eps=1e-6)

gradcheck 会用数值差分去近似雅可比矩阵,跟你的反向传播结果做对比。它能同时验证前向和反向是否正确,是排查自定义算子问题的第一工具。我在写一些临时的滑窗逻辑时,也会拿它快速验证自己的包装没有破坏梯度流,省下不少调试时间。

最后说一个跟我工作习惯有关的小事:我在自己的代码里总是先写一行注释,“unfold 之后第二维是 C×Kh×Kw,第三维是空间位置”,每次 reshape 前先打印 shape,再动手。这句话帮我避免了很多次静默出错。如果这篇文章只能让你记住一句话,我希望是:unfold 是展平滑动窗口,fold 是散点累加,窗口重叠时 fold 累加而不是平均。记牢这一句,大部分坑都不会再踩。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询