如何在PyTorch中拼接tensor?转列表操作是否会影响性能?
第一个问题解答
- 把tensor转列表的操作确实会导致性能下降,尤其当你处理的是大尺寸tensor、或者这块逻辑需要高频执行的时候问题更明显。
list(a)的操作本质是把输入tensor在第0维的每个切片都生成独立的tensor对象,不仅会产生额外的内存开销,如果你的tensor存放在GPU上,还会涉及不必要的CPU-GPU数据同步,性能比原生tensor操作差很多。
第二个问题解答
- 有纯tensor的实现方式,而且效率更高。你当前的代码逻辑是把形状为
(2, 3, 4, 5)的a在第0维拆成2个(3,4,5)的子tensor,再在第2维拼接得到(3,4,10)的结果,下面两种方法都可以实现同等效果:
# 方法1:用unbind直接拆分tensor,不需要转列表 b = torch.cat(a.unbind(0), dim=2) # 方法2:用维度重排+变形实现,无额外拷贝开销,性能最优 b = a.permute(1, 2, 0, 3).reshape(3, 4, -1)
你可以用torch.equal(b, 原代码生成的b)验证两者结果完全一致。
内容的提问来源于stack exchange,提问作者roger
相关产品推荐
相关产品推荐

