PyTorch如何将(4,4,4,4)张量扩展为(4,4,4,5)且新增元素为1?
给Tensor新增全1维度的正确实现方式
给定形状为(4,4,4,4)的PyTorch Tensor,要扩展为(4,4,4,5)且新增的最后一维元素全为1,你提供的代码思路是对的,但需要注意设备和数据类型匹配的问题,以下是几种可靠的实现方式:
方式一:修正原代码(适配设备与类型)
原代码的逻辑没问题,但默认生成的torch.ones是CPU上的float32类型,如果你的原Tensor在GPU上或是其他数据类型,会导致拼接失败。修正后代码如下:
import torch # 示例:创建一个(4,4,4,4)的Tensor points = torch.randn(4, 4, 4, 4) # 可替换为你的实际Tensor # 获取原形状并修改最后一维为1 pshape = list(points.size()) pshape[-1] = 1 # 生成与原Tensor同设备、同数据类型的全1 Tensor z = torch.ones(pshape, device=points.device, dtype=points.dtype) # 在最后一维拼接 result = torch.cat((points, z), dim=-1) print(result.shape) # 输出: torch.Size([4, 4, 4, 5])
方式二:更简洁的写法
利用ones_like直接生成与原Tensor某一子维度匹配的全1 Tensor,省去手动构造形状的步骤:
z = torch.ones_like(points[..., :1]) # 取原Tensor最后一维的第一个元素的形状,生成全1 Tensor result = torch.cat((points, z), dim=-1)
方式三:直接构造目标形状再赋值
先创建符合最终形状的全1 Tensor,再把原Tensor的值赋值进去:
# 直接生成(4,4,4,5)的全1 Tensor result = torch.ones((*points.shape[:-1], 5), device=points.device, dtype=points.dtype) # 将原Tensor的值赋值到前4个维度 result[..., :-1] = points
以上三种方式都能实现需求,其中方式一和方式二更适合保留原Tensor的所有属性,方式三在需要直接生成最终Tensor时更直观。
内容的提问来源于stack exchange,提问作者enryuxbt
相关产品推荐
相关产品推荐

