ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch Tensor.flatten 详解:start_dim 参数与实战避坑指南

PyTorch Tensor.flatten 详解:start_dim 参数与实战避坑指南 最近调一个 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_dim0, end_dim-1)默认值就是start_dim0, 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_dim1 print(y.shape) # torch.Size([2, 12])这里start_dim1表示从第 1 维开始把第 1 维和第 2 维3 和 4合并成 12而第 0 维2没有参与重排所以输出形状是(2, 12)。如果写成x.flatten()或x.flatten(0)结果会是torch.Size([24])。注意start_dim0就是从第 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-1end_dim默认也是-1两个参数指向同一个维度展平一个维度本身不改变形状。这种情况看似“没操作”但在某些动态代码里可以用它来确保张量在该维度上是连续存储的避免后续view报错。2.2 最容易错的地方别把 start_dim 当成“保留到第几维”我见过不少同事包括以前的我看到start_dim1第一反应是“保留前 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、6flatten(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_dim1}, ..., d_end_dim), d_{end_dim1}, ..., dn)也就是说start_dim之前的维度都保留start_dim到end_dim之间的所有维度乘起来成为一个新的维度end_dim之后的维度也保留。以x torch.rand(2, 3, 4, 5)为例做一张速查表flatten 参数参与合并的维度输出 shapeflatten()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_dim2所以第 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) # 保留 batch8 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), keepdimTrue) # 归一化成概率分布 # 展平 H*W变成 (B, H*W) flat_probs probs.flatten(1) # 对每个 batch 样本采样一个位置索引 sampled_indices torch.multinomial(flat_probs, num_samples1).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_size3), nn.ReLU(), nn.Flatten(start_dim1), nn.Linear(64 * 6 * 6, 10) # 假设输入为 8x8卷积后 6x6 )nn.Flatten()的默认行为也是start_dim1即保留 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。这个小习惯帮我省下了不少调试时间希望你也能用上。
返回列表