You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 12:55:37