如何通过交换行列索引重塑Torch张量
PyTorch 二维张量行列交换实现
你需要的效果本质是交换张量的行、列维度(即矩阵转置),将原张量的第i列转换为新张量的第i行,直接用PyTorch内置的维度交换接口即可实现,不需要手动遍历修改。
实现代码
首先定义原始张量:
import torch a = torch.tensor([[ 101, 101, 101, 101], [14812, 16890, 2586, 2586], [10337, 1830, 3842, 3842], [ 7257, 14541, 3293, 3297]])
两种等价的实现方式,任选其一即可:
# 写法1:二维张量转置的简写形式 a_target = a.T # 写法2:通用维度交换接口,明确指定交换第0维(行)和第1维(列) a_target = torch.transpose(a, dim0=0, dim1=1)
结果验证
打印输出a_target就能得到你需要的目标结构:
tensor([[ 101, 14812, 10337, 7257], [ 101, 16890, 1830, 14541], [ 101, 2586, 3842, 3293], [ 101, 2586, 3842, 3297]])
注意:上述两种操作返回的都是原张量的视图,不重新开辟内存拷贝数据;如果后续需要对张量做变形等要求内存连续的操作,可以在转置后调用
.contiguous()方法,例如a.T.contiguous()。
内容的提问来源于stack exchange,提问作者Dammio
相关产品推荐
相关产品推荐

