如何按指定列值对批量PyTorch张量进行行维度排序?
PyTorch批量张量按指定列排序行的实现
针对形状为b×m×n的PyTorch张量,要对每个batch中的m×n子张量,按每行第k列的值对行进行排序(保持输入输出形状均为b×m×n),可以通过以下步骤实现:
核心实现代码
import torch def sort_rows_by_column(tensor, target_col): # 提取每个batch中所有行的目标列值,形状为(b, m) sort_keys = tensor[:, :, target_col] # 获取按目标列升序排列的行索引,dim=1对应行维度 sorted_indices = torch.argsort(sort_keys, dim=1) # 将索引扩展为(b, m, n)形状,匹配原张量维度以支持gather操作 expanded_indices = sorted_indices.unsqueeze(-1).repeat(1, 1, tensor.size(2)) # 根据索引重新排列行,得到排序后的张量 sorted_tensor = tensor.gather(1, expanded_indices) return sorted_tensor
代码解释
- 提取排序依据:
tensor[:, :, target_col]取出每个batch里所有行的第target_col列(索引从0开始),作为每行排序的依据。 - 获取排序索引:
torch.argsort(..., dim=1)对每个batch的目标列值按行排序,返回排序后的行索引,确保每个batch内的排序独立进行。 - 扩展索引维度:通过
unsqueeze(-1)和repeat将索引从(b, m)扩展为(b, m, n),保证和原张量的维度完全匹配,这样才能在gather操作中正确定位每一行的所有列。 - 重排行:
tensor.gather(1, expanded_indices)在第1维度(行维度)上根据索引重新排列行,最终输出形状保持b×m×n不变。
示例验证
用题目给出的张量测试:
# 原始张量 a = torch.as_tensor([[[1, 3, 7, 6], [9, 0, 6, 2], [3, 0, 5, 8]], [[1, 0, 1, 0], [2, 1, 0, 3], [0, 0, 6, 1]]]) # 指定按第3列(索引为2)排序 sorted_column = 2 sorted_a = sort_rows_by_column(a, sorted_column) print(sorted_a)
输出结果:
tensor([[[3, 0, 5, 8], [9, 0, 6, 2], [1, 3, 7, 6]], [[2, 1, 0, 3], [1, 0, 1, 0], [0, 0, 6, 1]]])
完全符合题目中的预期结果。
扩展:降序排序
如果需要按目标列降序排列行,只需在argsort中添加descending=True参数:
sorted_indices = torch.argsort(sort_keys, dim=1, descending=True)
内容的提问来源于stack exchange,提问作者BeginnersMindTruly
相关产品推荐
相关产品推荐

