news 2026/10/6 8:25:57

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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_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。这个小习惯帮我省下了不少调试时间,希望你也能用上。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/6 8:24:58

C#电动车租赁会员管理系统课设拆解:三层架构与计费实战

简介:一份面向高校计算机相关专业的电动车租赁会员管理系统完整项目包,覆盖会员管理、租赁流程、后台数据维护等典型业务,适合作为毕业设计、课程设计或C#项目实战练习。压缩包共381个文件,总大小52.61MB,核心为176个C…

作者头像 李华
网站建设 2026/10/6 8:24:56

电动车租赁会员管理系统设计与实现:从业务模型到数据库部署全解析

简介:这是一份面向计算机相关专业学生的电动车租赁会员管理系统项目包,适合用于毕业设计、课程设计或初学C#开发的进阶练习。系统包含完整源码、数据库设计、项目说明与用户手册,覆盖会员管理与租赁业务的主要流程;代码经测试可正…

作者头像 李华
网站建设 2026/10/6 8:23:59

Ventoy 1.1.11 详解:多系统引导、ISO 拷贝与避坑指南

简介:Ventoy 1.1.11 Windows 版是一款开源多重U盘启动盘制作工具,面向需要频繁维护系统、安装多套操作系统的用户与运维人员,解决传统U盘启动盘一次只能写入一个ISO、反复格式化的痛点。它支持将微PE、老毛桃、Ubuntu、Debian、Windows Serve…

作者头像 李华
网站建设 2026/10/6 8:23:56

iPaaS平台选型实战:四大集成工具深度评测与避坑指南

企业一到系统集成这个环节就头大。客户数据在CRM里,订单在ERP里,营销数据在数据库里,报表要拉到数仓,再给老板出一张实时看板,每个系统单独看都没问题,一搞集成就是一场接一场的“点对点扯皮”。写接口、排…

作者头像 李华
网站建设 2026/10/6 8:22:10

Spring Boot课程管理系统开发全攻略:从建表到部署答辩

前阵子帮一个学弟把Springboot的课程设计项目调通了,顺手把这类系统从选题、建表、写代码到部署调试、写论文的完整思路整理了一遍。如果你的课设题目刚好是课程管理系统,或者你刚拿到一份Springboot课程管理系统源码、数据库脚本和配套论文文档&#xf…

作者头像 李华
网站建设 2026/10/6 8:21:09

网页文字复制被禁?ALLOWCOPY插件原理、安装与使用全指南

平时查资料最烦遇到什么?说白了就是好不容易找到一段有用的文字,结果鼠标一选选不了,右键菜单弹出来的是“您已被禁止复制”,或者一复制就自动往你剪贴板后面塞一段推广语。这类页面见得多了以后,我干脆在自己的浏览器…

作者头像 李华