最近调一个 Transformer 的相关代码时,被一句x.flatten(2)卡了半天。我心里一直默认flatten()就是把张量整个拉成一维,怎么这里还带个数字?后来翻文档才知道,PyTorch 的Tensor.flatten(start_dim)并没有那么“无脑”,start_dim决定了从哪个维度开始压平,也决定了哪些维度会被保留。这篇文章就专门把Tensor.flatten的用法讲透,重点是start_dim的实际语义,再结合几个我常遇到的真实场景,给新手和有经验的读者都做一份可以直接“抄作业”的参考。
1. Tensor.flatten 本质是什么:先看它和 view、reshape 的关系
1.1 默认的 flatten():就是把所有维度拉成一维
先说最简单的情况。一个任意形状的张量,直接调.flatten(),等价于.reshape(-1),返回一个一维张量。比如:
import torch t = torch.tensor([[1, 2, 3], [4, 5, 6]]) print(t.flatten()) # tensor([1, 2, 3, 4, 5, 6])这个操作把“行”和“列”合并成一条线,本质上是按内存中的行优先顺序把元素逐个取出。你可以把多维张量想象成一个多层文件柜:flatten 就是按从上到下、从左到右的顺序,把每一层抽屉里的文件全部倒出来,整齐排成一列。文件数量不变,排列顺序不变,只是“摆放形式”变了。
这种默认用法在验证模型输入输出、打印中间特征时很实用,但当你想保留某个维度(比如 batch 维度)时,就需要用到start_dim了。
1.2 start_dim 到底是干什么的:不是保留维度,而是“从第几个维度开始压”
Tensor.flatten的完整签名是:
Tensor.flatten(start_dim=0, end_dim=-1)默认值就是start_dim=0, end_dim=-1,意思是从第 0 维开始,一直压到最后一个维度,最终得到一维张量。start_dim真正的作用是:把从start_dim到end_dim之间的所有维度合并成一个维度,而start_dim之前的维度保持原样。
举一个最常见的例子:
x = torch.zeros(2, 3, 4) # shape: (2, 3, 4) y = x.flatten(1) # start_dim=1 print(y.shape) # torch.Size([2, 12])这里start_dim=1,表示从第 1 维开始,把第 1 维和第 2 维(3 和 4)合并成 12,而第 0 维(2)没有参与重排,所以输出形状是(2, 12)。
如果写成x.flatten()或x.flatten(0),结果会是torch.Size([24])。注意,start_dim=0就是从第 0 维开始压,也就是所有维度全部参与。
1.3 为什么需要这个参数:三个典型场景
- 全连接层输入:CNN 卷积层输出的特征图是
(B, C, H, W),全连接层接收的是二维矩阵(batch, features)。如果不保留 batch 维度,直接flatten(),batch 信息就会混进特征里,模型根本没法训练。 - 多模态特征拼接:把文本向量和图像特征展平到同一维度时,通常只展平特征相关维度,不触碰 batch 维度。
- 概率分布采样:想从
(B, H, W)的概率图中按像素采样位置,需要先展平成(B, H*W),然后用torch.multinomial。
所以start_dim出现的意义,就是为了让你能精准控制“哪些维度合并、哪些维度保留”。
2. start_dim 的细节与踩坑指南
2.1 维度索引规则:从 0 开始,也支持负数
PyTorch 的维度索引和 Python 列表一样:第 0 维是第一个维度,也可以使用负索引,-1指最后一个维度,-2指倒数第二个维度,以此类推。
x = torch.zeros(2, 3, 4) # 从倒数第 2 维开始压,合并最后两个维度 3、4 -> 12 y = x.flatten(-2) print(y.shape) # torch.Size([2, 12]) # 从最后一个维度开始压,只展平一个维度 z = x.flatten(-1) print(z.shape) # torch.Size([2, 3, 4])这里可能有人会觉得奇怪:flatten(-1)为什么没有变化?因为start_dim=-1,end_dim默认也是-1,两个参数指向同一个维度,展平一个维度本身不改变形状。这种情况看似“没操作”,但在某些动态代码里,可以用它来确保张量在该维度上是连续存储的,避免后续view报错。
2.2 最容易错的地方:别把 start_dim 当成“保留到第几维”
我见过不少同事(包括以前的我)看到start_dim=1,第一反应是“保留前 1 维”,也就是以为输出的第一维是原来的第 0 维和第 1 维拼接。这是完全错误的。
start_dim的含义是“压平的起点”,而不是“保留的终点”。想判断到底保留哪些维度,要看start_dim之前的维度。
x = torch.rand(4, 5, 6) print(x.flatten(1).shape) # torch.Size([4, 30]) print(x.flatten(2).shape) # torch.Size([4, 5, 6]) print(x.flatten(0).shape) # torch.Size([120])从左到右看:flatten(1)保留第 0 维,合并后面 5、6;flatten(2)保留第 0、1 维,只合并第 2 维(单个维度合并等于没变);flatten(0)不保留任何维度,全部合并。
2.3 形状推导公式:所有情况直接手算
其实不需要死记,给一个通用公式。输入形状为(d0, d1, ..., dn),如果调用flatten(start_dim, end_dim),输出形状为:
(d0, d1, ..., d_{start_dim-1}, prod(d_start_dim, d_{start_dim+1}, ..., d_end_dim), d_{end_dim+1}, ..., dn)也就是说,start_dim之前的维度都保留,start_dim到end_dim之间的所有维度乘起来成为一个新的维度,end_dim之后的维度也保留。
以x = torch.rand(2, 3, 4, 5)为例,做一张速查表:
| flatten 参数 | 参与合并的维度 | 输出 shape |
|---|---|---|
flatten() | 0,1,2,3 | (120,) |
flatten(0) | 0,1,2,3 | (120,) |
flatten(1) | 1,2,3 | (2, 60) |
flatten(0, 2) | 0,1,2 | (20, 5) |
flatten(2, 3) | 2,3 | (2, 3, 20) |
flatten(-2) | 2,3 | (2, 3, 20) |
注意,flatten(0, 2)显式指定了end_dim=2,所以第 3 维保留,输出(2*3*4, 5)=(20,5)。用这个公式,任何组合都能手算。
3. 实操过程:从场景到代码
3.1 场景一:CNN 特征图展平后接全连接层
这是flatten(1)最经典的使用场景。假设你的卷积网络输出一个特征图,形状为(B, 64, 7, 7),接下来要接nn.Linear(64 * 7 * 7, 10),你需要把每个样本的 64 个通道、7x7 的空间位置全部展开成 3136 个特征。但 batch 维度要保留,否则多个样本全混在一起。
import torch import torch.nn as nn x = torch.randn(8, 64, 7, 7) # 模拟 CNN 输出 x_flat = x.flatten(1) # 保留 batch=8 print(x_flat.shape) # torch.Size([8, 3136]) fc = nn.Linear(64 * 7 * 7, 10) out = fc(x_flat) print(out.shape) # torch.Size([8, 10])如果错误地用了x.view(-1),得到的形状是(25088,),nn.Linear第一维和它完全不匹配。而且从语义上讲,这等于把所有样本的特征全部拼在一起,模型无法区分样本边界,等于直接废掉了 batch 结构。
3.2 场景二:Transformer 中序列维度合并
Transformer 里更常见的操作是(B, S, D),分别表示 batch、序列长度、特征维度。有时候你想把序列长度和特征维度合并,比如做一些全局池化前的特征融合,就可以用flatten(1, 2):
x = torch.randn(2, 10, 32) # (batch, seq_len, d_model) y = x.flatten(1, 2) # 合并 seq_len 和 d_model print(y.shape) # torch.Size([2, 320])另外多头注意力中常用(B, num_heads, S, head_dim)。如果想把 head 维度和 head_dim 合并,可以写成x.flatten(2, 3),保留 batch 和 heads。如果想把 batch 和 heads 合并,则写成x.flatten(0, 1)。理解这一点后,你会发现用 flatten 操作 Transformer 中间张量,比手算 reshape 数字要安全得多。
3.3 场景三:根据某个 tensor 来采样一个值:flatten 与 multinomial 配合
“根据某个 tensor 来 sample 一个值”这个问题经常出现在强化学习、目标检测、多模态模型里,比如从概率图中采样一个坐标,从分类分布中采样一个类别索引。torch.multinomial是一个很好用的函数,但它要求输入是二维矩阵:第一维是 batch 或独立分布,第二维是每个类别的概率。
假设有一个概率矩阵probs,形状是(B, H, W),表示每个 batch 样本中,每个空间位置的概率。直接对三维张量调用multinomial会报错,所以先展平空间维度,再采样,最后还原坐标。
import torch B, H, W = 2, 3, 4 probs = torch.rand(B, H, W) probs = probs / probs.sum(dim=(1, 2), keepdim=True) # 归一化成概率分布 # 展平 H*W,变成 (B, H*W) flat_probs = probs.flatten(1) # 对每个 batch 样本采样一个位置索引 sampled_indices = torch.multinomial(flat_probs, num_samples=1).squeeze(-1) print(sampled_indices.shape) # torch.Size([2]) # 由展平后的索引还原 H、W 坐标 h_coords = sampled_indices // W w_coords = sampled_indices % W print(h_coords, w_coords)这里flatten(1)保证了每个 batch 的概率分布是独立的一行,然后multinomial对每一行采样一个位置。展平顺序是行优先,所以sampled_index // W就是行坐标,sampled_index % W就是列坐标。如果想从(B, C, H, W)特征图采样一个通道位置,可以先flatten(1)变成(B, C*H*W),采样索引再用连续的除法取模恢复三个坐标。这套方法我在做目标检测的稀疏采样时经常用,比一层层for循环不知道快多少。
3.4 nn.Flatten 层:模型定义时直接用模块
PyTorch 还提供了torch.nn.Flatten模块,里面可以设置start_dim和end_dim。如果你喜欢用nn.Sequential搭网络,可以直接嵌入:
import torch.nn as nn model = nn.Sequential( nn.Conv2d(3, 64, kernel_size=3), nn.ReLU(), nn.Flatten(start_dim=1), nn.Linear(64 * 6 * 6, 10) # 假设输入为 8x8,卷积后 6x6 )nn.Flatten()的默认行为也是start_dim=1,即保留 batch 维,后面全部展平,跟Tensor.flatten(1)一致。从使用习惯上看,函数式x.flatten(1)更灵活,适合在自定义forward里动态处理;模块式nn.Flatten在搭建网络时更直观,方便别人一眼看到“这里做了展平”。
4. 常见问题与排查技巧实录
4.1 为什么 flatten(0) 和 flatten(1) 效果差别这么大
这是新手最容易迷惑的问题。flatten(0)表示从第 0 维开始展平,整个张量变成一维;flatten(1)表示从第 1 维开始展平,第 0 维保留。
x = torch.randn(2, 3, 4) print(x.flatten(0).shape) # torch.Size([24]) print(x.flatten(1).shape) # torch.Size([2, 12])判断标准很简单:你的数据里,哪个维度代表“独立样本”?如果是 batch,那就是第 0 维,通常要用flatten(1)而不是flatten(0)。如果你是在处理单样本特征,或者是最后一次输出,可能确实需要flatten(0),但那时要确认模型结构真的不依赖 batch 维度。
4.2 start_dim 超出维数范围会怎样
会直接报IndexError。比如三维张量调用flatten(3):
IndexError: Dimension out of range (expected to be in range of [-3, 2], but got 3)意思是当前张量只有 3 个维度,合法索引范围是-3到2。排查时先打印x.dim()看维度总数,再检查start_dim是否在合法范围。另外还要注意start_dim必须小于等于end_dim,否则会报类似“start_dim cannot come after end_dim”的错误。
4.3 非连续张量:为什么 flatten 比 view 更省心
这是一个很隐蔽的坑。view要求张量在内存中是连续的,而转置、切片等操作会产生非连续张量。直接对转置张量调用view(-1)会报错:
x = torch.randn(3, 4) x_t = x.t() # 非连续 print(x_t.is_contiguous()) # False # 报错 # x_t.view(-1) # RuntimeError: view size is not compatible with input tensor...但flatten()不会报错,因为它内部会按需复制数据,保证展平结果可用:
y = x_t.flatten() print(y.shape) # torch.Size([12])这也是我推荐在不确定是否连续时使用flatten而不是view的原因。但要注意:如果输入非连续,flatten返回的可能是重新拷贝后的张量,与原始张量不共享内存。修改展平结果不会影响原张量,这一点和view的“视图”语义不同。如果你需要严格的内存共享,还是要先调用.contiguous(),再使用view。
4.4 问题排查速查表
| 现象 | 可能原因 | 解决方式 |
|---|---|---|
| 全连接层输入维度对不上 | 忘记保留 batch 维度,用了flatten(0) | 改用flatten(1) |
| 输出 shape 始终不变 | start_dim指向了最后一个维度且 end 相同 | 检查参数是否传了-1, -1 |
| 展平后数据顺序和预期不一致 | 输入张量非连续,展平按拷贝后的连续顺序 | 先打印is_contiguous(),必要时先contiguous() |
调用multinomial报错 | 概率张量不是二维 | 先用flatten(1)变成(B, dim) |
| 手工还原坐标错误 | 没有考虑行优先展平顺序 | 用//和%分别取行、列 |
4.5 一个额外的小技巧:用 unflatten 逆操作
如果你在flatten之后需要根据原始形状还原,不需要自己写divmod或者手工reshape。PyTorch 提供了unflatten方法,它和flatten是一对:
x = torch.randn(2, 3, 4) y = x.flatten(1) # (2, 12) z = y.unflatten(1, (3, 4)) # 恢复 (2, 3, 4) print(z.shape) # torch.Size([2, 3, 4])这在写可复用的数据变换、做 batch 维保持的解码时非常方便。尤其是在采样坐标需要还原时,unflatten能让你少写很多容易出错的坐标换算公式。
我在实际项目里最后总结出一条经验:碰到任何涉及形状变换的代码,先在手边写一下原始 shape 和目标 shape,标清楚哪一维是 batch、哪一维是特征,再决定start_dim取值。特别是flatten和multinomial配合时,展平顺序决定了还原索引的算法,必须先想明白物理意义再动手。另外,如果你不确定张量是否连续,用flatten而不是view,至少不会突然给你抛一个 runtime error。这个小习惯帮我省下了不少调试时间,希望你也能用上。