如何为PyTorch Tensor新增维度并重复原维度对应值
PyTorch张量扩维并重复元素的实现方法
要将形状为[a,b]的PyTorch张量转换为[a,b,c],且新维度上的元素重复原张量对应位置的值,有以下几种简洁高效的实现方式:
方法1:unsqueeze + repeat
先通过unsqueeze在目标位置新增一个维度(使张量变为[a,b,1]),再用repeat在新维度上重复指定次数:
import torch # 原张量示例 x = torch.tensor([[1,2,3], [4,5,6]]) c = 4 # 扩维并重复 result = x.unsqueeze(2).repeat(1, 1, c)
unsqueeze(2):在第三个维度(索引从0开始)新增维度,将原形状[2,3]转为[2,3,1]repeat(1,1,c):前两个维度保持原大小(重复1次),第三个维度重复c次,最终得到[2,3,4]的张量
方法2:unsqueeze + expand
expand是更高效的方式,它不会复制数据,仅返回张量的扩展视图(前提是原张量连续):
result = x.unsqueeze(-1).expand(-1, -1, c)
unsqueeze(-1):等价于unsqueeze(2),表示在最后一维新增维度expand(-1,-1,c):-1表示保持对应维度的原有大小,无需手动指定a和b的值,适配性更强
方法3:利用广播机制
借助PyTorch的广播特性,将原张量与形状为[1,1,c]的全1张量运算,实现元素重复:
result = x.unsqueeze(-1) * torch.ones(1, 1, c, dtype=x.dtype, device=x.device)
这种方式无需显式调用重复方法,但会额外创建全1张量,效率略低于前两种方法。
验证结果
以上方法执行后,result的输出均为:
tensor([[[1, 1, 1, 1], [2, 2, 2, 2], [3, 3, 3, 3]], [[4, 4, 4, 4], [5, 5, 5, 5], [6, 6, 6, 6]]])
形状为torch.Size([2, 3, 4]),完全符合需求。
内容的提问来源于stack exchange,提问作者Lrom
相关产品推荐
相关产品推荐

