大数跨境

Tensor变形记:快速掌握Tensor切分、变形等方法

Tensor变形记:快速掌握Tensor切分、变形等方法 知识代码AI
2026-08-21
1
导读:Tensor变形记:快速掌握Tensor切分、变形等方法一、Tensor的连接操作在深度学习项目开发中,某一层

Tensor变形记:快速掌握Tensor切分、变形等方法


一、Tensor的连接操作

在深度学习项目开发中,某一层神经元的数据可能有多个不同的来源,需要将数据进行组合,这个操作称为连接

1. torch.cat —— 在已有维度上拼接

torch.cat(tensors, dim=0, out=None)

参数说明:

  • tensors:若干个准备进行拼接的 Tensor
  • dim:约定拼接的方向(维度)

示例(两个 3×3 矩阵):

A = torch.ones(33)
B = 2 * torch.ones(33)

dim=0(按"行"方向拼接): 得到 6×3 矩阵dim=1(按"列"方向拼接): 得到 3×6 矩阵

核心要点:dim 的数值是多少,两个矩阵就按照相应维度的方向连接。cat 不会改变 Tensor 的维度(秩)。[pdf_5]


2. torch.stack —— 增加新维度进行拼接

torch.stack(inputs, dim=0)

与 cat 的核心区别:

函数
操作方式
维度变化
torch.cat
在已有维度上连接
维度不变
torch.stack 新增一个维度
再拼接
升维

示例(两个 4 元素向量):

A = torch.arange(04)    # tensor([0, 1, 2, 3])
B = torch.arange(59)    # tensor([5, 6, 7, 8])

dim=0(在"行"方向新建维度): 得到 2×4 矩阵dim=1(在"列"方向新建维度): 得到 4×2 矩阵[pdf_5]


二、Tensor的切分操作

切分是连接的逆操作,分为三种类型:chunk、split、unbind。

1. chunk —— 按份数尽可能平均切分

torch.chunk(input, chunks, dim=0)

参数说明:

  • input:待切分的 Tensor
  • chunks:将要被划分的块的数量(整型),注意不是每组数量
  • dim:按哪个维度切分

chunk 的分配规律:先做除法再向上取整得到每组数量。例如 17 个元素切 4 份 → 17/4=4.25 → 向上取整为 5 → 先逐个生成若干个长度为 5 的向量,最后不够的放一块作为最后一个向量。[pdf_5]

示例:

A = torch.tensor([12345678910])
B = torch.chunk(A, 20)    # 10位切2份 → 每份5个
C = torch.chunk(A, 30)    # 10位切3份 → 4个、4个、2个

注意:当 chunks 大于可切分长度时,结果是每个元素单独一份,多余的会返回空。[pdf_5]


2. split —— 按每份指定大小切分

torch.split(tensor, split_size_or_sections, dim=0)

参数说明:

  • split_size_or_sections为整数时,表示每块大小为该整数;为列表时,表示切成与列表中元素大小一样的块
  • dim:按哪个维度切分

示例一(整数参数,4×4 按每份 2 行切分): 得到 2 个 2×4 矩阵示例二(不能整除时,4×4 按每份 3 行切分): 得到 3×4 和 1×4 矩阵示例三(列表参数,5×4 按 (2, 3) 切分): 得到 2×4 和 3×4 矩阵[pdf_5]

规律:PyTorch 会尽可能凑够每一个结果,使对应 dim 的数据大小等于 split_size_or_sections,最后剩下的不够就作为最后一个结果。


3. unbind —— 逐条去除某个维度

torch.unbind(input, dim=0)

示例(4×4 矩阵):

A = torch.arange(016).view(44)

沿第 0 维("行"方向)切分: 得到 4 个长度为 4 的向量沿第 1 维("列"方向)切分: 得到 4 个长度为 4 的向量(每一列)

核心要点:unbind 是降维切分方式,相当于删除一个维度之后的结果。[pdf_5]


切分函数对比

函数
切分方式
维度变化
chunk
份数切分(尽可能平均)
维度不变
split
指定大小切分
维度不变
unbind
逐条去除某个维度
降维

三、Tensor的索引操作

当只需要部分数据时,使用索引操作。

1. index_select —— 按给定索引选择

torch.index_select(tensor, dim, index)

参数说明:

  • tensor:待处理的 Tensor
  • dim:选择数据的维度
  • index:从 dim 维度中的哪些位置选择数据(注意:index 是 torch.Tensor 类型)

示例(4×4 矩阵):

A = torch.arange(016).view(44)

选择第 0 维的第 1 行和第 3 行:

B = torch.index_select(A, 0, torch.tensor([13]))
# 结果:[[4,5,6,7], [12,13,14,15]]

选择第 1 维的第 0 列和第 3 列:

C = torch.index_select(A, 1, torch.tensor([03]))
# 结果:[[0,3], [4,7], [8,11], [12,15]]

2. masked_select —— 按掩码条件选择

torch.masked_select(input, mask, out=None)

参数说明:

  • input:待处理的 Tensor
  • mask:掩码张量,即满足条件的特征掩码。mask 须与 input 有相同数量的元素数目,但形状或维度不需要相同

示例:

A = torch.rand(5)
B = A > 0.3# 生成布尔掩码
C = torch.masked_select(A, B)  # 提取A中>0.3的元素

简化写法:

C = torch.masked_select(A, A > 0.3)

应用场景:常用于提取网络中某一层数值大于 0 的参数等。[pdf_5]


索引函数对比

函数
选择方式
输出维度
index_select
按给定的索引位置提取
与输入维度相同
masked_select
条件(掩码)提取
返回一维输出

四、关键注意事项与总结

  1. 边界数值要仔细:使用这些函数时,最需要关注的是维度和大小相关的参数,要提前仔细计算好,否则会产生错误结果。
  2. chunk vs split:chunk 按"切分成确定的份数"切分;split 按"每份确定的大小"切分。
  3. cat vs stack:cat 在已有维度上拼接、不升维;stack 增加新维度拼接、会升维
  4. unbind 是降维操作:相当于删除一个维度后的结果。
  5. index_select vs masked_select:前者基于给定的索引位置提取数据(输出与输入同维);后者基于判断条件(掩码)提取数据(输出为一维)。
  6. 补充知识点:index_select 返回的结果和输入是一个维度,而 masked_select 返回一维输出;split 获取的是原输入的视图,对 split 结果的操作会影响原数据。[pdf_5]

五、每课一练

问题torch.cat 和 torch.stack 有什么区别?

解答

  • torch.cat 是在已有维度上连接,不会改变 Tensor 的维度(秩)。
  • torch.stack 会新增一个维度再进行拼接,是升维操作。
  • 例如:将两个形状为 (3, 3) 的矩阵用 cat 拼接,dim=0 时得到 (6, 3);用 stack 拼接,dim=0 时得到 (2, 3, 3)。

六、小结

  1. 连接操作cat(不升维)和 stack(升维),注意 dim 参数的方向。
  2. 切分操作chunk(按份数)、split(按大小)、unbind(降维切分)。
  3. 索引操作index_select(按索引,同维输出)和 masked_select(按条件,一维输出)。
  4. 核心原则:理解 dim(维度)的含义是掌握所有操作的关键,一切操作都围绕维度展开。


【声明】内容源于网络
0
0
知识代码AI
技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
内容 403
粉丝 0
知识代码AI 技术基底 机器视觉全栈 × 光学成像 × 图像处理算法 编程栈 C++/C#工业开发 | Python智能建模 工具链 Halcon/VisionPro工业部署 | PyTorch/TensorFlow模型炼金术 | 模型压缩&嵌入式移植
总阅读6.6k
粉丝0
内容403