关于torch.stack()不同dim参数下张量堆叠规则的疑问
PyTorch torch.stack() 接口dim参数逻辑详解
核心本质
torch.stack() 和常用的 torch.cat() 核心差异是:torch.cat() 是在输入张量已有的维度上做拼接,输出张量维度数和输入一致;torch.stack() 是新增一个独立维度用来堆叠多个张量,输出维度数 = 输入张量维度数 + 1。
基础规则
- 所有待堆叠的输入张量必须形状完全一致
dim参数的合法取值范围是0 ≤ dim ≤ 输入张量的维度数,代表新增维度插入的位置- 输出张量新增维度的长度等于待堆叠的张量总数量
示例拆解
你给出的示例中,三个输入张量 t1/t2/t3 都是1维张量,形状为[3],待堆叠数量为3:
- 若设置
dim=0,代表新增维度插入到第0位,输出形状为[3, 3],相当于把三个1维张量按顺序整体摞起来,结果为:
tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
- 你使用的
dim=1,代表新增维度插入到第1位,输出形状同样为[3, 3],此时会将三个输入张量的同索引元素配对,组成新维度的元素,最终得到你看到的输出:
tensor([[1, 4, 7], [2, 5, 8], [3, 6, 9]])
高维张量通用逻辑
不管输入张量是2维、3维还是更高维度,都可以用统一逻辑判断:固定其他维度的索引,仅改变输出张量dim对应维度的索引,得到的序列就是各个输入张量对应位置的元素。
我们以2维输入为例验证,输入两个形状为[2, 3]的张量,堆叠数量为2:
import torch a = torch.tensor([[1,2,3],[4,5,6]]) b = torch.tensor([[7,8,9],[10,11,12]]) # dim=0堆叠,新增维度在最前 out0 = torch.stack([a,b], dim=0) print(out0.shape) # 输出 torch.Size([2, 2, 3]) print(out0[0].equal(a)) # True,dim=0取0对应第一个输入张量 print(out0[1].equal(b)) # True,dim=0取1对应第二个输入张量 # dim=1堆叠,新增维度在第1位 out1 = torch.stack([a,b], dim=1) print(out1.shape) # 输出 torch.Size([2, 2, 3]) print(out1[:,0,:].equal(a)) # True,dim=1取0对应第一个输入张量 print(out1[:,1,:].equal(b)) # True,dim=1取1对应第二个输入张量 # dim=2堆叠,新增维度在最后 out2 = torch.stack([a,b], dim=2) print(out2.shape) # 输出 torch.Size([2, 3, 2]) print(out2[:,:,0].equal(a)) # True,dim=2取0对应第一个输入张量 print(out2[:,:,1].equal(b)) # True,dim=2取1对应第二个输入张量
快速验证技巧
拿到任意堆叠需求,可以先快速推算输出形状:假设输入张量形状为S = (s0, s1, ..., sn-1),待堆叠张量数量为N,指定dim=k,则输出形状为将N插入到S的第k位后的元组,比如输入形状为(3,4,5),堆叠数量为2,dim=2,输出形状就是(3,4,2,5)。
内容的提问来源于stack exchange,提问作者Dhruv Vashist
相关产品推荐
相关产品推荐

