如何将PyTorch中(2,3,4)维度张量转为(3,2,4)并满足指定值对应关系
解决方法
你的问题核心是交换张量的前两个维度,而torch.reshape和torch.view只是基于内存连续性重新排列元素形状,不会改变轴的顺序,所以无法满足需求。正确的做法是使用维度转置类操作:torch.permute 或 torch.transpose。
具体实现示例
import torch # 创建测试张量,shape=(2,3,4) x = torch.arange(2*3*4).reshape(2,3,4) # 方法1:使用permute交换第0和第1维度,得到shape=(3,2,4) y_permute = x.permute(1, 0, 2) # 验证:转换后的[:,0,:]与原张量[0,:,:]相等 print(torch.allclose(y_permute[:,0,:], x[0,:,:])) # 输出True # 方法2:使用transpose交换两个维度(仅适用于交换两个轴的场景) y_transpose = x.transpose(0, 1) # 结果与permute完全一致 print(torch.allclose(y_transpose[:,0,:], x[0,:,:])) # 输出True
关键说明
permute:支持一次性交换任意数量的维度,参数是新维度顺序的索引(比如原维度是(0,1,2),传入(1,0,2)就表示把原第1维度放到新第0位,原第0维度放到新第1位,第2维度保持不变)。transpose:只能交换两个维度,参数是要交换的两个维度索引,适合仅需交换两个轴的简单场景。- 如果后续需要对转换后的张量使用
reshape或view,可以添加.contiguous()方法让张量内存连续:y = x.permute(1,0,2).contiguous()。
内容的提问来源于stack exchange,提问作者Sirojbek Safarov
相关产品推荐
相关产品推荐

