Tensor变形记:快速掌握Tensor切分、变形等方法
一、Tensor的连接操作
在深度学习项目开发中,某一层神经元的数据可能有多个不同的来源,需要将数据进行组合,这个操作称为连接。
1. torch.cat —— 在已有维度上拼接
torch.cat(tensors, dim=0, out=None)
参数说明:
-
tensors:若干个准备进行拼接的 Tensor -
dim:约定拼接的方向(维度)
示例(两个 3×3 矩阵):
A = torch.ones(3, 3)
B = 2 * torch.ones(3, 3)
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(0, 4) # tensor([0, 1, 2, 3])
B = torch.arange(5, 9) # 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([1, 2, 3, 4, 5, 6, 7, 8, 9, 10])
B = torch.chunk(A, 2, 0) # 10位切2份 → 每份5个
C = torch.chunk(A, 3, 0) # 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(0, 16).view(4, 4)
沿第 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(0, 16).view(4, 4)
选择第 0 维的第 1 行和第 3 行:
B = torch.index_select(A, 0, torch.tensor([1, 3]))
# 结果:[[4,5,6,7], [12,13,14,15]]
选择第 1 维的第 0 列和第 3 列:
C = torch.index_select(A, 1, torch.tensor([0, 3]))
# 结果:[[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 |
|
|
四、关键注意事项与总结
-
边界数值要仔细:使用这些函数时,最需要关注的是维度和大小相关的参数,要提前仔细计算好,否则会产生错误结果。 -
chunk vs split:chunk 按"切分成确定的份数"切分;split 按"每份确定的大小"切分。 -
cat vs stack:cat 在已有维度上拼接、不升维;stack 增加新维度拼接、会升维。 -
unbind 是降维操作:相当于删除一个维度后的结果。 -
index_select vs masked_select:前者基于给定的索引位置提取数据(输出与输入同维);后者基于判断条件(掩码)提取数据(输出为一维)。 -
补充知识点: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)。
六、小结
-
连接操作: cat(不升维)和stack(升维),注意 dim 参数的方向。 -
切分操作: chunk(按份数)、split(按大小)、unbind(降维切分)。 -
索引操作: index_select(按索引,同维输出)和masked_select(按条件,一维输出)。 -
核心原则:理解 dim(维度)的含义是掌握所有操作的关键,一切操作都围绕维度展开。

