在PyTorch中如何将0维张量拼接至3维张量中?
报错原因
torch.cat 要求参与拼接的所有张量除了拼接维度外,其余维度的尺寸必须完全匹配,你当前使用的两个张量维度差异过大:
X形状为(30, 1, 2),是3维张量t形状为(1,),是1维张量
而torch.stack要求所有输入张量形状完全一致,拼接时会新增一个维度,也不符合你的需求。
可行实现方案
你只需要先把t的形状调整为和X前两维匹配、最后一维尺寸为1的形状,再沿最后一维拼接即可,有两种常用实现方式:
方法1:调整t形状后拼接
先通过expand把t广播到匹配的形状,再用torch.cat拼接,适合t的值不是固定0的场景:
# 先把t调整为 [30, 1, 1] 的形状,和X前两维完全匹配 t_reshape = t.reshape(1, 1, 1).expand(30, 1, 1) result = torch.cat((X, t_reshape), dim=-1) # 输出result形状为 torch.Size([30, 1, 3])
如果t的值是固定的0,你也可以直接生成对应形状的张量简化代码:
t = torch.zeros(30, 1, 1) result = torch.cat((X, t), dim=-1)
方法2:直接对X做填充
如果你只是需要在X的最后一维末尾补1个固定值,用pad操作更简洁:
import torch.nn.functional as F # pad参数格式为 (最后一维左侧补的数量, 最后一维右侧补的数量),这里右侧补1个0 result = F.pad(X, (0, 1), value=0) # 输出result形状为 torch.Size([30, 1, 3])
内容的提问来源于stack exchange,提问作者Dong Le
相关产品推荐
相关产品推荐

