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

PyTorch中如何通过创建新维度拼接张量?有无更简便方法?

更简便的张量新维度拼接方法

嘿,你想要的这种“创建新维度来拼接”的操作,PyTorch里其实有好几个更简洁的实现方式,不用手动加None再cat,下面给你介绍两个最常用的:

1. 使用torch.stack() - 通用的新维度堆叠

torch.stack()就是专门为这种场景设计的!它会自动帮你在指定的位置插入一个新维度,然后把所有输入张量沿着这个新维度拼接起来,而且要求所有输入张量的形状完全一致(正好符合你重复x的需求)。

示例代码:

import torch
x = torch.randn(2, 3)

# 在dim=0的位置创建新维度并堆叠4次x
stack_result = torch.stack([x, x, x, x], dim=0)
# 或者用列表生成式更简洁:torch.stack([x]*4, dim=0)
print(stack_result.shape)  # 输出: torch.Size([4, 2, 3])

如果想把新维度放在中间或者最后,只需要调整dim参数就行,比如dim=1会得到(2,4,3),dim=-1会得到(2,3,4),非常灵活。

2. 使用torch.repeat() - 重复单个张量的快捷方式

如果你只是想把同一个张量重复多次并新增维度,torch.repeat()会更省事,不用写多个x的列表。它的参数是每个维度的重复次数,我们只需要在原张量的维度前面加一个重复次数,就能实现新增维度的效果:

示例代码:

repeat_result = x.repeat(4, 1, 1)
print(repeat_result.shape)  # 输出: torch.Size([4, 2, 3])

这里的(4,1,1)表示:在新增的第一个维度重复4次,原有的两个维度(2和3)各重复1次(也就是保持不变)。本质上相当于先给x新增一个维度变成(1,2,3),然后在第一个维度重复4次。

额外小技巧:内存友好的expand()

如果你的张量很大,不想复制太多内存,可以先用unsqueeze()新增维度,再用expand()进行广播式扩展(不会实际复制数据,只是逻辑上扩展):

expand_result = x.unsqueeze(0).expand(4, 2, 3)
print(expand_result.shape)  # 输出: torch.Size([4, 2, 3])

不过要注意,expand()得到的张量和原张量共享内存,如果后续对它进行修改,会影响原张量,适合只读场景使用。

内容的提问来源于stack exchange,提问作者Ben

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 07:10:45