如何用张量R高效排序多维张量C以得到Sorted_C?
基于张量R对张量C排序的高效实现方法
我需要基于张量R对张量C进行排序,得到目标张量sorted_c。相关张量定义如下:
import torch c = torch.tensor([[[0, 1, 0, 0, 0], [1, 0, 1, 0, 0], [0, 0, 0, 1, 0], [0, 0, 0, 0, 1]], [[0, 0, 1, 1, 0], [1, 1, 0, 0, 0], [0, 0, 1, 0, 0], [0, 1, 0, 1, 0]]]) r = torch.tensor([[[0, 0, 0, 7.2, 0], [0, 25.4, 0, 0, 0], [0, 0, 43.6, 0, 0], [61.8, 0, 0, 0, 0], [0, 0, 0, 0, 80]], [[0, 0, 0, 0, 98.2], [116.4, 0, 0, 0, 0], [0, 134.6, 0, 0, 0], [0, 0, 152.8, 0, 0], [0, 0, 0, 169.2, 0]]]) # 目标结果 sorted_c = torch.tensor([[[0, 1, 0, 0, 0], [0, 0, 1, 1, 0], [1, 0, 0, 0, 0], [0, 0, 0, 0, 1]], [[0, 0, 0, 1, 1], [0, 1, 1, 0, 0], [0, 0, 0, 1, 0], [0, 0, 1, 0, 1]]])
修正说明:P6至P10的预期排序应为:P10 -> P6 -> P7 -> P8 -> P9
高效实现方案
根据修正说明,排序规则为:每个样本中元素的优先级由R每行的非零值升序决定(如第二个样本中P10(98.2) < P6(116.4) < P7(134.6) < P8(152.8) < P9(169.2))。以下是两种常见排序场景的实现:
场景1:对C的列按元素优先级重排
# 提取每个元素的排序键,shape: [2,5] sort_keys = torch.max(r, dim=2)[0] # 生成元素的排序索引(升序),shape: [2,5] sorted_element_indices = torch.argsort(sort_keys, dim=1) # 重排C的每一行的列 sorted_c = torch.gather(c, dim=2, index=sorted_element_indices.unsqueeze(1).repeat(1, c.shape[1], 1))
场景2:对C的行按行内元素的最高优先级排序
# 提取每个元素的排序键,shape: [2,5] sort_keys = torch.max(r, dim=2)[0] # 计算每行的排序键(取行内元素的最大排序键),shape: [2,4] row_sort_keys = torch.max(c * sort_keys.unsqueeze(1), dim=2)[0] # 生成行的排序索引(升序),shape: [2,4] sorted_row_indices = torch.argsort(row_sort_keys, dim=1) # 重排C的行 sorted_c = torch.gather(c, dim=1, index=sorted_row_indices.unsqueeze(2).repeat(1, 1, c.shape[2]))
关键说明
- 使用
torch.max快速提取R每行的非零值,避免循环遍历 torch.argsort生成排序索引,保证操作的向量化高效性torch.gather是PyTorch原生的高效张量重排操作,适合大规模数据处理
内容的提问来源于stack exchange,提问作者Sam
相关产品推荐
相关产品推荐

