PyTorch中对应tf.transpose(y, perm=[2,0,1])的等效实现是什么?
PyTorch 中与
tf.transpose(y, perm=[2, 0, 1]) 功能等价的实现方法 TensorFlow 中tf.transpose的perm参数用于直接指定调整后的维度顺序,PyTorch 中完全对应这个功能的 API 是torch.permute,和只能两两交换维度的torch.transpose不同,permute可以直接传入自定义的维度索引序列完成重排。
针对你给出的示例,张量y的形状为(10, 480, 640),三个维度的索引分别对应0:10、1:480、2:640,要实现perm=[2, 0, 1]的效果,直接按顺序传入维度索引即可:
import torch y = torch.randn(10, 480, 640) x = y.permute(2, 0, 1) # 验证输出形状:x.shape 为 torch.Size([640, 10, 480]),和 TensorFlow 版本输出完全一致
如果需要用torch.transpose间接实现也可以,需要多次交换维度,可读性比permute差,不推荐使用:
# 等价的间接实现,功能完全一致 x = y.transpose(0, 2).transpose(1, 2)
补充说明:PyTorch 的permute和transpose返回的都是原张量的视图,不会额外复制数据,和 TensorFlow 中tf.transpose的底层行为一致。
内容的提问来源于stack exchange,提问作者Dalek
相关产品推荐
相关产品推荐

