如何将扁平化张量的索引还原为原张量的位置索引?
还原扁平化索引为原张量位置
可以直接用PyTorch内置的torch.unravel_index()函数,它专门用于将扁平化后的一维索引转换为对应原张量形状的多维坐标,是最简便的解决方案。
代码示例
import torch x = torch.tensor([[[1, 2, 3, 4, 5], [2, 3, 4, 5, 6], [4, 5, 6, 7, 8]]]) v, i = torch.topk(x.flatten(), 3, largest=False) # 转换索引为原张量的多维坐标 coords = torch.unravel_index(i, x.shape) # 整理为直观的位置列表 position_list = list(zip(*coords)) print(position_list) # 输出:[(0, 0, 0), (0, 1, 0), (0, 0, 1)]
手动计算方式(理解原理)
如果想手动实现转换逻辑,基于原张量形状(1,3,5),拆分每个扁平化索引的步骤如下:
- 原张量各维度元素数:
d0=1(第一维)、d1=3(第二维)、d2=5(第三维) - 对每个扁平化索引
idx:- 第一维坐标:
idx // (d1*d2),即每个第一维元素包含3*5=15个元素,索引0、5、1除以15均得0 - 计算剩余索引:
rem = idx % (d1*d2),得到0、5、1 - 第二维坐标:
rem // d2,5//5=1,0//5=0,1//5=0 - 第三维坐标:
rem % d2,0%5=0,5%5=0,1%5=1
- 第一维坐标:
手动实现代码:
shape = x.shape d0, d1, d2 = shape positions = [] for idx in i: dim0 = idx // (d1*d2) rem = idx % (d1*d2) dim1 = rem // d2 dim2 = rem % d2 positions.append((dim0.item(), dim1.item(), dim2.item())) print(positions) # 输出:[(0, 0, 0), (0, 1, 0), (0, 0, 1)]
内容的提问来源于stack exchange,提问作者Hawkeye
相关产品推荐
相关产品推荐

