如何以Pythonic简洁写法实现PyTorch张量unsqueeze维度调整
PyTorch 多次unsqueeze的简洁替代写法
连续重复调用unsqueeze确实冗余,下面两种原生PyTorch写法都可以完全替代重复调用,更符合Pythonic规范,且没有额外性能开销:
用None做新维度占位符(最简洁)
PyTorch张量和NumPy数组一致,支持在索引中用None插入长度为1的新维度,不需要反复调用函数,写法非常清爽:
- 对应你第一种「在最后一个维度连续扩展3次」的场景,用省略号代指前面所有原有维度,后面直接跟3个None即可:
inps = torch.FloatTensor(data[0]) tgts = torch.FloatTensor(data[1]) # 效果完全等价于连续3次调用unsqueeze(dim=-1) tgts = tgts[..., None, None, None] inps = inps[..., None, None, None]
- 对应你第二种「在dim=1位置连续扩展3次」的场景,先保留第0维,插入3个新维度后再用省略号代指后面所有原有维度即可:
inps = torch.FloatTensor(data[0]) tgts = torch.FloatTensor(data[1]) # 效果完全等价于连续3次调用unsqueeze(dim=1) tgts = tgts[:, None, None, None, ...] inps = inps[:, None, None, None, ...]
显式构造目标形状(可读性最强)
如果希望维度逻辑更直白,不需要记忆索引规则,可以直接基于原张量的shape拼接目标形状,用view(张量内存连续时使用)或reshape(无内存连续要求)实现:
- 末尾加3个维度的场景:
# 原形状后拼接3个长度为1的维度 inps = inps.view(*inps.shape, 1, 1, 1) tgts = tgts.view(*tgts.shape, 1, 1, 1)
- dim=1位置插入3个维度的场景:
# 第0维保持不变,插入3个长度为1的维度,后续维度按原顺序保留 inps = inps.view(inps.shape[0], 1, 1, 1, *inps.shape[1:]) tgts = tgts.view(tgts.shape[0], 1, 1, 1, *tgts.shape[1:])
注意:上面两种方法本质都是返回原张量的视图,不会复制底层数据,和重复调用
unsqueeze的计算结果、性能表现完全一致。不建议用循环写法批量调用unsqueeze,反而会增加不必要的函数调用开销,可读性也没有优势。
内容的提问来源于stack exchange,提问作者Totoro
相关产品推荐
相关产品推荐

