PyTorch中用reshape/view实现张量指定顺序展平重塑的通用方法
原因说明
PyTorch中view和reshape默认按内存行优先(C顺序)读取元素:遍历张量时优先走最后一个维度,逐行拼接取值。你给出的2维示例张量形状为(3,2),直接展平时会依次取每一行的元素,因此得到[1 4 2 5 3 6],和你预期的先遍历第一列、再遍历第二列的顺序不符。
2维场景解法
要得到目标顺序[1 2 3 4 5 6],先转置交换两个维度,再展平即可:
import torch import numpy as np a = torch.tensor([[1, 4], [2, 5],[3, 6]]) # 写法1:用view,需要先把转置后的张量转为内存连续格式 b = a.t().contiguous().view(6).detach().numpy() print(b) # 输出 [1 2 3 4 5 6] # 写法2:用reshape可以省略contiguous,reshape会自动处理非连续内存 b = a.t().reshape(6).detach().numpy()
高维场景通用方案
对于3维及以上的张量,核心逻辑是先调整维度顺序,再展平,步骤如下:
- 明确你需要的元素遍历优先级:最先遍历的维度优先级最高,最后遍历的维度优先级最低
- 调用
permute()方法,按优先级从高到低传入维度索引,重排张量维度 - 调用
reshape(-1)即可得到对应顺序的一维展平结果
举个3维张量的示例:
# 形状为(2,3,2)的3维张量 a = torch.tensor([ [[1,7],[2,8],[3,9]], [[4,10],[5,11],[6,12]] ]) # 按「第2维→第0维→第1维」的优先级遍历展平,得到[1,2,3,4,5,6,7,8,9,10,11,12] b = a.permute(2, 0, 1).reshape(-1).numpy() print(b)
如果不需要保留张量的梯度传播能力,也可以直接转numpy后用列优先模式展平,仅2维场景下可以直接得到你要的结果:
b = a.detach().numpy().flatten(order='F')
内容的提问来源于stack exchange,提问作者Miraboreasu
相关产品推荐
相关产品推荐

