PyTorch中如何通过二维张量高效索引另一二维张量
PyTorch 二维张量分组高效索引实现
这个场景不需要写Python循环分组,全用PyTorch内置向量化算子就能实现,CPU/GPU都能高效运行,步骤如下:
- 先拆分目标张量和索引张量的组号、值、组内偏移字段
- 利用
torch.unique_consecutive快速定位每个分组在A中的起始行位置(因为同组行连续排列,这个算子比普通unique快很多) - 把索引张量里的组号匹配到对应分组的起始位置,加上组内偏移得到要取值的全局行索引
- 直接用全局索引取值,和原组号拼接得到最终结果
完整可运行代码:
import torch A = torch.tensor([ [0, 0], [0, 2], [0, 3], [0, 4], [0, 5], [0, 6], [1, 0], [1, 1], [1, 4], [1, 5], [1, 6] ]) b = torch.tensor([[0, 2], [1, 2]]) # 拆分字段 A_groups = A[:, 0] A_values = A[:, 1] b_groups = b[:, 0] b_offsets = b[:, 1] # 提取每个连续分组的起始行索引 unique_groups, group_starts = torch.unique_consecutive(A_groups, return_index=True) # 匹配查询组对应的起始位置 start_pos = group_starts[torch.searchsorted(unique_groups, b_groups)] # 计算全局索引、取值、拼接结果 global_idx = start_pos + b_offsets result = torch.stack([b_groups, A_values[global_idx]], dim=1)
运行后得到的result和预期完全一致:
tensor([[0, 3], [1, 4]])
方案说明
- 全程无Python层循环,所有操作都是PyTorch底层优化过的张量算子,大规模数据下性能远高于手写循环分组的实现
- 不强制要求组号是从0开始的连续整数,只要A中同组的行连续排列即可正常运行
- 如果输入的A中同组行不连续,提前加一行
A = A[A[:, 0].argsort()]按组号排序即可,排序算子同样支持GPU加速,开销很低
内容的提问来源于stack exchange,提问作者Kriss
相关产品推荐
相关产品推荐

