You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

关于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.04 01:39:03