如何让torch.argsort返回二维索引?
如何让torch.argsort返回二维索引?
torch.argsort会根据dim参数返回行内或列内的索引,如何让它返回形如
[[r1,r2,r3,...], [c1,c2,c3,.....]]的二维索引?
可以通过实现一个unravel_indices函数,将一维扁平化索引转换为对应形状的多维坐标,结合torch.argsort的结果就能得到目标格式的二维索引。
实现代码
import torch def unravel_indices(indices, shape): coord = [] for dim in reversed(shape): coord.append(torch.fmod(indices, dim)) indices = torch.div(indices, dim, rounding_mode='floor') coord = torch.stack(coord[::-1], dim=-1) return coord # 将函数绑定到torch模块,方便后续调用 torch.unravel_indices = unravel_indices
使用示例
假设我们有一个2×3的张量,想要获取排序后的二维索引:
tensor = torch.tensor([[3, 1, 2], [6, 4, 5]]) # 先将张量扁平化,再获取排序后的一维索引 sorted_flat_idx = torch.argsort(tensor.flatten()) # 转换为二维坐标索引(形状为[N, 2],N为元素总数) two_d_coords = torch.unravel_indices(sorted_flat_idx, tensor.shape) # 转换为要求的[[行索引], [列索引]]格式 row_idx, col_idx = two_d_coords[:, 0], two_d_coords[:, 1] final_indices = torch.stack([row_idx, col_idx]) print(final_indices) # 输出结果: # tensor([[0, 0, 1, 1, 0, 1], # [1, 2, 1, 2, 0, 0]])
内容的提问来源于stack exchange,提问作者sten
相关产品推荐
相关产品推荐

