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
相关产品推荐
相关产品推荐

