☰
PyTorch F.unfold完全指南:从im2col到视觉Transformer窗口自注意力
2026/10/3 1:07:05 网站建设 项目流程

前阵子有个做视觉Transformer的朋友问我,PyTorch里F.unfold到底返回了个什么东西,为什么一个四维张量进去,出来变成了三维,窗口里的顺序又是怎么排的。我当时就说,这个操作别看API简单,它几乎是所有“滑动窗口类”模型的地基,从最传统的卷积网络到Swin Transformer里的窗口自注意力,背后都在悄悄用同一个套路:把局部区域抽出来,拼成一个大矩阵,然后该做矩阵乘法做矩阵乘法,该算注意力算注意力。本文就把unfold这个操作彻底讲透,覆盖它的数学形态、与卷积/im2col的关系、在视觉Transformer里的应用、手写实现细节,以及我实际踩过的各种坑,适合刚入门深度学习、正在啃“动手深度学习”这类教程、或者想在自注意力网络里自己实现窗口机制的读者。

1. 先搞明白:unfold到底从“展开”里得到了什么

1.1 用一张4x4的图,手工推一遍展开过程

很多人第一次看unfold的文档会懵,因为它的行为和直觉里的“reshape”完全不一样。我建议这样理解:unfold就是“按滑动窗口把图切开,然后每个窗口拉成一列”。它不是把整张图重新排列,而是把每个窗口位置上的局部像素都独立复制一份出来。

假设输入是一张单通道4x4的图,像素值就是1到16的自然数排列。如果我们指定窗口大小是2x2、步长为1,那么这张图上能放的窗口位置一共有多少?横向上4减2加1等于3个位置,纵向上也是3个,所以一共9个窗口。此时unfold的输出形状是(1, 4, 9),其中第一维是batch,第二维是C * kernel_h * kernel_w = 1 * 2 * 2 = 4,第三维是窗口数量9。

这9个窗口分别是左上角(1,2,5,6)、右上角的(2,3,6,7)、下面的(3,4,7,8),然后换到第二行窗口(5,6,9,10)……每个窗口都被拉成一个长度为4的列向量。你把这9个列向量按顺序拼在一起,就得到unfold的输出。这种操作在传统图像处理里叫im2col,全称是“image to column”,意思就是把图像块转成列,unfold其实就是深度学习框架里对im2col的封装。

1.2 为什么要把每个窗口单独复制出来,而不是用循环

初学者最容易问的一个问题是:我也可以用两层for循环去切patch,为什么非要搞一个unfold出来?关键在于两点。第一是计算效率,for循环在Python里是解释执行的,慢得离谱,而unfold底层是C++实现,一次调用就能把整张图的所有窗口都切好;第二是批量处理能力,unfold能同时处理整个batch的所有图、所有通道,而手写循环要处理这些维度会非常痛苦。

还有一个更微妙的好处:unfold让数据在内存里的排布变得规整,后续想做矩阵乘法、卷积、注意力计算都能直接拼成一个大矩阵操作。在GPU时代,矩阵乘法是亲儿子,任何算得慢的操作只要能被改写成矩阵乘法,性能都会突飞猛进。你后面会看到,卷积层能被训练得这么快,靠的就是这一步“展开”加上一个矩阵乘法。

2. unfold在实战里的核心应用:从卷积到视觉Transformer

2.1 卷积的底层真相:im2col加矩阵乘法

如果你去翻PyTorch的卷积实现,会发现F.conv2d内部并不是真的在GPU上做那种“滑窗相乘再累加”的直观操作。更常见的高效实现路径是:先把输入特征图用im2col展开成一个大矩阵,再把卷积核reshape成另一个矩阵,两者做矩阵乘法,最后再reshape回输出尺寸。这里im2col展开出来的矩阵,就是unfold的产物。

举个例子,输入是(N, C, H, W),卷积核大小是K,展开后得到(N, C*K*K, L),其中L是输出特征图上的位置数。同时把卷积核从(out_channels, C, K, K)reshape成(out_channels, C*K*K)。两个矩阵一乘,得到的(N, out_channels, L)再reshape回(N, out_channels, H_out, W_out),卷积就做完了。这就是为什么很多讲“动手深度学习”这类实战教程时,会把unfold和矩阵乘法结合起来,手动实现一个可微的卷积层。

理解这个真相有什么好处?第一,你写自定义算子时知道去哪找性能瓶颈;第二,调试网络时能把“卷积输出不对”的问题定位到展开阶段还是矩阵乘法阶段;第三,你会明白为什么卷积核大小和stride变了,展开矩阵的行数、列数会怎么变,所有维度都清清楚楚。

2.2 视觉Transformer里的patch embedding,本质上也是一个unfold

ViT把图像切成一个个patch,然后对每个patch做线性投影,得到token序列。很多初学Transformer的人会奇怪,为什么代码里用的是Conv2d(patch_size, patch_size, stride=patch_size)而不是一个显式的切patch操作?因为卷积的im2col展开和patch切分在数学上是等价的,而且卷积层还能顺带完成线性投影。

如果你把这里换成unfold,流程会变得更透明:先用F.unfold把图像切成重叠或非重叠的patch,形状是(N, C*P*P, L),然后转置成(N, L, C*P*P),再接一个nn.Linear(C*P*P, embed_dim),得到(N, L, embed_dim)的token序列。这两个写法结果几乎一致,区别只在于Conv2d在底层额外把这个过程优化过了而已。所以你在看ViT源码时,如果看到别人用einops.rearrange或F.unfold来实现patch embedding,不要觉得奇怪,它们都是同一个操作的不同马甲。

2.3 窗口自注意力里的unfold:如何把特征图切成窗口

热词里提到的“wsa和跨窗口自注意力的网络结构”,落到底层实现时,第一步就是把整张特征图切成一个个不重叠的小窗口,这正是unfold的标准用法。假设输入是(B, C, H, W),窗口大小是M,步长是M,那么F.unfold之后得到(B, C*M*M, L),这里的L就是窗口个数。

然后你把第二维拆开成(B, C, M*M, L),再做一次permute变成(B, L, M*M, C),就得到了“每个窗口里所有token”的张量。接下来在这个张量上做自注意力,每个窗口之间互不干扰,计算效率很高。窗口内部做完注意力之后,再用F.fold把这些窗口拼回原图尺寸。如果做的是跨窗口的shifted window操作,就是把窗口划分的起点平移一下,然后再次走unfold、注意力、fold的流程。整个Swin Transformer就是在这个“切窗口、算注意力、拼回去、换一种切法再算一次”的循环里工作的。

2.4 unfold和fold是一对夫妻,别只认识一个

F.unfold负责把图切碎,F.fold负责把碎块拼回去。fold大致上是unfold的逆操作:给定一堆窗口的列向量和一个输出尺寸,它会按照每个值在原来窗口里的位置,把值填回输出特征图。如果窗口之间有重叠,同一个位置会被多个窗口填到,fold默认会把它们求和。

理解这层关系非常重要,因为很多自定义网络结构都要“先切再拼”。比如做图像修复,你可能想把图像切成patch分别处理,再拼回整图;做医学图像分割,也可能要在大patch上推理再融合结果。没有fold,你切完就回不去了。而“重叠区域求和”这个行为既可以是优点也可以是坑,后面我会讲怎么用计数矩阵实现平均而不是求和。

3. 手写一个unfold:维度、步长、填充背后的数学

3.1 unfold的输入输出形状和参数对照

在PyTorch里,F.unfold的签名是F.unfold(input, kernel_size, dilation=1, padding=0, stride=1)。输入张量的形状必须是(N, C, H, W),输出形状是(N, C * kernel_h * kernel_w, L),其中L的计算公式是:

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

这个公式和卷积输出尺寸的计算是同一个公式,因为unfold本来就在干卷积滑窗的事,只是它不做加权求和而已。我实测的时候踩过一个小小的坑:把dilation默认成1,然后忘了给padding,导致输出尺寸比预期小了一圈。每次写unfold之前,最好先在草稿纸上算一遍L,确认输出尺寸符合预期再喂给后面的层。

下面这张表总结了各参数对输出形状的影响,建议保存一份:

参数作用输出的影响
kernel_size窗口大小越大,C*K*K越大,L越小
stride滑窗步长越大,L越小;与窗口大小相同时不重叠
padding边缘补零越大,L越大,可保持尺寸不变
dilation空洞间隔越大,感受野越大,同尺寸下L变小

3.2 行顺序的排列规则:通道优先,然后是核内位置

unfold输出矩阵的每一行不是随机排列的,它遵循“通道优先”规则。假设输入有C个通道,窗口大小是K*K,那么输出矩阵的行索引顺序是:先枚举第一个通道上的(K, K)个像素位置,再枚举第二个通道上的(K, K)个像素位置,直到最后一个通道。

具体到代码里,第r行对应的含义是:先算c = r // (K*K)得到通道序号,再算offset = r % (K*K),接着根据offset反推行坐标kh = offset // K,列坐标kw = offset % K。这个排列方式意味着,如果你想把(B, C*K*K, L)还原成“窗口内所有的token向量”,正确的变换是view(B, C, K*K, L).permute(0, 3, 2, 1),也就是把通道维度放到最后一维。我见过不止一次有人漏掉这个permute,直接把维度当作(B, L, C*K*K)去算注意力,结果注意力权重错得离谱。

3.3 手写Python循环版unfold,验证对原理的理解

为了彻底搞懂底层逻辑,我当初给自己布置了一个小作业:不用PyTorch的unfold,纯用for循环实现同样的功能,然后用随机张量验证结果是否完全一致。这个过程很值得做一遍,尤其适合正在“深度学习入门”阶段挣扎的同学。

步骤是这样的:输入x形状(1, 1, 4, 4),窗口大小2,步长1。先算窗口位置总数,横纵各3个,共9个。然后两层for循环遍历所有窗口位置,每次取出x[0, 0, i:i+2, j:j+2],用reshape(-1)拉平成4个元素,填到输出矩阵的第col列。写完后和F.unfold的结果一对比,如果数值一致、形状一致,说明你对“展开”的理解已经到位了。

我用代码来演示一下这个循环版本的核心部分:

import torch x = torch.arange(1, 17, dtype=torch.float32).reshape(1, 1, 4, 4) kh, kw = 2, 2 stride = 1 H, W = 4, 4 Lh = (H - kh) // stride + 1 Lw = (W - kw) // stride + 1 out = torch.zeros(1, kh * kw, Lh * Lw) col = 0 for i in range(Lh): for j in range(Lw): patch = x[0, 0, i:i+kh, j:j+kw].reshape(-1) out[0, :, col] = patch col += 1 # 和PyTorch自带unfold对照 ref = torch.nn.functional.unfold(x, kernel_size=(kh, kw), stride=stride) print(out) print(ref) print(torch.allclose(out, ref))

跑完之后你会看到两个张量完全一致。能写出这个循环,说明你对“窗口位置如何编号、每个窗口内部如何排列”已经有了直观认识,后面做任何基于窗口的自定义结构都不会再心虚。

4. 用unfold做窗口自注意力的完整实操

4.1 一个可直接运行的窗口切分加注意力示例

现在我们来做一个真正能用的东西:给一张形状为(B, C, H, W)的特征图,切成不重叠的窗口,在窗口内部做标准自注意力,再把结果拼回原图。这是Swin Transformer窗口注意力的最小实现,去掉shift和相对位置偏置这些细节,剩下的骨架就是这样。

import torch import torch.nn.functional as F def window_attention(x, window_size, embed_dim, num_heads): B, C, H, W = x.shape M = window_size # 1. unfold切窗口 windows = F.unfold(x, kernel_size=M, stride=M) # (B, C*M*M, L) # 2. 重新排列成 (B, L, M*M, C) windows = windows.view(B, C, M * M, -1).permute(0, 3, 2, 1) # 3. 线性投影得到QKV,这里为了演示直接用x当Q、K、V q = windows k = windows v = windows # 4. 缩放点积注意力 d_k = q.shape[-1] attn = (q @ k.transpose(-2, -1)) / (d_k ** 0.5) attn = F.softmax(attn, dim=-1) out = attn @ v # (B, L, M*M, C) # 5. fold拼回原图 out = out.permute(0, 3, 2, 1).contiguous() # (B, C, M*M, L) out = F.fold(out, output_size=(H, W), kernel_size=M, stride=M) return out x = torch.randn(2, 8, 16, 16) y = window_attention(x, window_size=4, embed_dim=8, num_heads=1) print(y.shape)

这个代码跑起来输出(2, 8, 16, 16),和输入形状一样。实际使用中你会把q、k、v换成由nn.Linear投影得到的向量,并且引入多头机制。

4.2 重叠窗口的fold还原:别让重叠区域直接相加

上面的例子用的是不重叠窗口,所以fold拼回去没有任何歧义。但如果你在做重叠的滑窗处理,比如步长小于窗口大小,那么同一个位置会被多个窗口覆盖,fold默认会把值累加,这通常不是你想要的还原效果。更合理的做法是先做一个计数矩阵,然后用fold结果除以计数矩阵,得到重叠区域的平均值。

具体实现分三步:第一步,把unfold的输出变成全1矩阵;第二步,用fold把全1矩阵拼回去,得到一个“每个位置被多少个窗口覆盖”的计数图;第三步,把真正值的fold结果除以计数图。代码如下:

def fold_overlap_correct(unfold_out, output_size, kernel_size, stride): C_kk = unfold_out.shape[1] H, W = output_size values = F.fold(unfold_out, output_size=(H, W), kernel_size=kernel_size, stride=stride) count = F.fold(torch.ones_like(unfold_out), output_size=(H, W), kernel_size=kernel_size, stride=stride) count = torch.clamp(count, min=1.0) return values / count

这是我做分割任务时踩过的坑:当时直接fold回去,没有除以计数,结果图像中间出现亮带,因为中心区域被多个窗口重叠,值被加了好几遍。后来改成上面的归一化方式,问题立刻消失。凡是做重叠窗口处理的同学,建议直接把这个函数抄进自己的工具集。

4.3 窗口注意力的维度坑:为什么总是需要permute

我刚接触unfold做窗口注意力时,几乎每次都要在维度顺序上折腾半天。核心原因是unfold给的维度顺序是“通道优先”,而Transformer需要的token序列是“位置优先”。(B, C*M*M, L)这种形状对注意力计算很不友好,因为你期望的是“把每个窗口当作一个batch,窗口内每个patch当作一个token”,token特征维度得在最后一位。

标准解法就是我把前面写的那句:view(B, C, M*M, -1).permute(0, 3, 2, 1)。这里view先把通道和核内位置拆开,permute把维度重排成(B, L, M*M, C)。完成注意力后再用相反顺序的permute加contiguous回到(B, C, M*M, L),才能喂给fold。很多人失败就失败在忘了contiguous,导致permute之后的张量不是连续内存布局,直接报错。

5. 常见坑位排查与性能笔记

5.1 形状不对:unfold输出维度的自查清单

如果你在调试中遇到unfold相关错误,我建议按下面几项逐条排查:

  • 确认输入是4D张量(N, C, H, W),而不是3D。很多人把一个(N, C, L)的张量传进去,程序直接报错。
  • 确认kernel_size、stride、padding、dilation都是整数或成对整数,混合用tuple和int容易产生歧义。
  • 确认输出尺寸L的计算值和预期一致。如果L算出来是0,说明窗口比图还大,赶紧调参数。
  • 确认你处理第二维时没有忘记除以K*K。当C*K*K很大时,你看到的第一维可能不是通道数,容易产生错觉。

我把这个清单写在我自己的调试模板里。每次unfold相关代码报错,先过一遍再查别的,节省大量时间。

5.2 内存爆掉:unfold的显存放大效应

unfold有个明显的代价:显存开销可能比原图大很多倍。假设输入特征图是(B, C, H, W),窗口大小是K,步长是S,展开后的数据量大约是:

原数据量 × K * K / (S * S)

当K=7, S=1时,放大倍数接近49倍。如果输入是batch为8、通道256、分辨率128x128的特征图,展开后数据量大约是8乘256乘128乘128乘49,这个数字瞬间就能把消费级显卡显存打满。所以我的经验是:能用小窗口别用大窗口,能用大步长别用小步长,实在没办法要做重叠滑窗时,考虑用循环切patch而不是一次性unfold。

另外,unfold返回的张量在内存里是重新拷贝的一份数据,不是原图的视图。这意味着即使你只是测试一下,内存也会实打实地被占用。如果数据量实在太大,建议先裁剪一个小区域调试通逻辑,再扩展到全图。

5.3 梯度能不能回传?unfold的反向传播是怎么回事

很多自己搭网络的人会担心,unfold这种“先复制数据再计算”的操作,梯度还能不能正确回来。答案是能,因为unfold本质是一个线性操作,可以写成一个稀疏矩阵乘法,反向传播就是把梯度按照每个窗口的位置,累加回原图对应的位置。这个累加过程在数学上正好等价于fold。

我实际测试过,在unfold后续接一个线性层再fold回来,用随机梯度检验和数值梯度对比,误差在1e-5量级,说明框架的反向传播实现是可靠的。有个值得注意的细节:如果同一位置被多个窗口共享,反向传播时梯度会累加,这是正确行为,因为前向展开时该位置被复制了多次,梯度自然要按贡献累加回去。

5.4 不同框架里的同类操作:TensorFlow、JAX怎么实现

如果你需要在PyTorch之外实现类似功能,也别慌。TensorFlow里对应的API是tf.image.extract_patches,输入输出形状和PyTorch的unfold基本一致,只是维度顺序略有差异,需要手动调整通道维的位置。JAX里虽然没有直接的pytorch式unfold,但你可以用lax.conv_general_dilated配合全1卷积核来模拟,或者直接用jax.lax里的窗口操作。如果你用einops库,用rearrange配合unfold的语义也能写出可读性很强的版本。

我这里给一个TensorFlow的简单对照:

import tensorflow as tf x = tf.reshape(tf.range(1, 17, dtype=tf.float32), [1, 4, 4, 1]) patches = tf.image.extract_patches(x, sizes=[1, 2, 2, 1], strides=[1, 1, 1, 1], rates=[1, 1, 1, 1], padding='VALID') print(patches.shape) # (1, 3, 3, 4)

这个输出的含义是9个窗口位置,每个位置拉出4个像素,和PyTorch的unfold本质相同,只是维度排布不同。跨框架理解这些操作时,抓住“窗口、滑窗、拉平”三个关键词,就不会乱。

5.5 调试技巧:写一个assert来验证unfold的排列顺序

最后分享一个我用了很久的调试技巧。每次写基于unfold的自定义层,我都会先构造一张“坐标图”:把输入像素值设成它自身的坐标编码,比如x[0, 0, i, j] = i * 100 + j,然后执行unfold,并手动检查展开后的前几列是不是符合预期。如果坐标对得上,说明窗口划分逻辑没问题;如果对不上,再去看permute和view的顺序。

这个技巧在我调试跨窗口自注意力结构时救过我很多次。因为一旦窗口划分错了,后面所有注意力结果都是垃圾,而这个“坐标图”方法能在10秒内把问题锁定在unfold层面,而不是让你在整个网络里大海捞针。建议你把这段验证代码记在笔记本里:

x = torch.arange(16, dtype=torch.float32).reshape(1, 1, 4, 4) u = F.unfold(x, kernel_size=2, stride=1) assert u.shape == (1, 4, 9) assert torch.equal(u[0, :, 0], torch.tensor([0., 1., 4., 5.])) assert torch.equal(u[0, :, 1], torch.tensor([1., 2., 5., 6.]))

这类断言能保证你后续大量基于unfold的代码不会因为最底层的数据排布错误而全部白算。

6. 用unfold扩展你的网络结构:几个可落地的进阶方向

6.1 把unfold用在时间序列和点云上

很多人以为unfold只能处理图像,其实它是通用的张量操作,只要你的数据支持“窗口滑过”这个结构,就能用。时间序列数据可以先把形状整成(N, C, L),然后按时间维度做窗口展开,用来构造局部时序特征;点云数据如果体素化成3D网格,也可以用3D版本的unfold扩展局部邻域。这个思路在最近的一些局部建模网络里经常出现,原理都是一样的。

我做时间序列预测时,就喜欢先用unfold把长为L的序列切成一个个长度为K的窗口,再把窗口内数据作为特征输入一个小MLP,这样比直接用RNN或Transformer更直观,而且能灵活控制上下文长度。窗口之间的重叠比例直接决定了特征之间的冗余度,这也是一个可以调节的超参数。

6.2 配合可变形卷积:让unfold从固定窗口变成可学习偏移

固定滑窗最大的问题是感受野完全由参数决定,不能根据内容自适应。可变形卷积的思路是:先用一个普通卷积预测每个位置的偏移量,再根据偏移量去采样。虽然PyTorch有deform_conv2d这种封装好的算子,但如果你要自己实现,关键步骤就是把普通的unfold变成一个“带偏移的采样”。

具体做法是:先用unfold得到标准窗口数据,然后根据预测的偏移量对窗口内像素做双线性插值。这样你能保留整个操作的可微性,同时让网络学会在哪些位置采样。这个玩法非常锻炼对底层张量操作的理解,也是研究型工作里常见的手写实现路径。

6.3 和einops配合,写出可读性更高的切patch代码

unfold虽然强大,但代码可读性确实一般,尤其是一连串view加permute,让后来者看得头皮发麻。如果你喜欢更清晰的写法,可以试试einops库的rearrange,比如:

from einops import rearrange windows = rearrange(x, 'b c (h m1) (w m2) -> b (h w) (m1 m2) c', m1=window_size, m2=window_size)

这行代码直接把(B, C, H, W)的张量切成窗口并重排成(B, L, window_size^2, C),效果和unfold加permute一样,但一眼就能看出在做什么。不过rearrange背后也是在干切块和重排的事,理解unfold能帮你更快掌握它。

我现在的习惯是:正式实验代码里用unfold加fold,保证性能和框架原生支持;写快速原型或分享代码时,用einops让读者更容易读懂。

7. 实战复盘:我从unfold踩坑里学到的三件事

第一,理解维度顺序比记住API更重要。unfold文档看十遍,不如自己打印一次张量的形状和内容,亲手验证“每个窗口拉成列”这个行为。你一旦在脑子里建立起“图到矩阵”的画面,后面所有相关代码都顺了。

第二,性能和内存是实际工程里躲不过的两个坎。unfold不是银弹,它会把数据复制多份,显存开销很大。我做过一个比较极端的分割实验,当unfold展开后数据量比原图大50倍时,训练速度肉眼可见地下降,最后只能改成小窗口加混合精度训练才扛下来。所以设计网络时一定要估算内存放大倍数,别等OOM报错才后悔。

第三,unfold和fold是对双胞胎,别只用一个。很多自定义结构需要“切了再拼”,只学会切不会拼等于断了一条腿。尤其是重叠区域的处理,不除以计数矩阵会出大问题,这个坑我见太多人踩过了,包括我自己。

如果你正在研究Swin这类带窗口注意力的结构,或者想自己写一个更高效的局部建模模块,unfold是你绝对绕不过去的一个基础操作。花一晚上把它的形状变化、排列顺序、反向传播规则全部验证一遍,之后你会发现自己看各种视觉模型源码都轻松了很多。

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

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

立即咨询