Pytorch 多维Tensor从lookup_table查询对应值索引的向量化实现咨询
解决方案
方案1:值映射法(性能最优,推荐)
该方案通过提前构建值到索引的映射表,实现O(1)时间复杂度的查询,性能远高于循环实现,支持任意维度的data输入:
import torch # 你的lookup_table和data定义在这里 lookup_table = torch.tensor([266, 103, 84, 12, 32, 34, 1, 523, 22, 136, 268, 432, 53, 63, 201, 51, 164, 69, 31, 42, 122, 131, 119, 36, 245, 60, 28, 81, 9, 114, 105, 3, 41, 86, 150, 79, 104, 120, 74, 420, 39, 427, 40, 59, 24, 126, 202, 222, 145, 429, 43, 30, 38, 55, 10, 141, 85, 121, 203, 240, 96, 7, 64, 89, 127, 236, 117, 99, 54, 90, 57, 11, 21, 62, 82, 25, 267, 75, 111, 518, 76, 56, 20, 2, 61, 516, 80, 78, 555, 246, 133, 497, 33, 421, 58, 107, 92, 68, 13, 113, 235, 875, 35, 98, 102, 27, 14, 15, 72, 37, 16, 50, 517, 134, 223, 163, 91, 44, 17, 412, 18, 48, 23, 4, 29, 77, 6, 110, 67, 45, 161, 254, 112, 8, 106, 19, 498, 101, 5, 157, 83, 350, 154, 238, 115, 26, 142, 143]) data = torch.tensor([ [523, 114, 350, 246, 30, 222, 39, 517, 106, 2], [ 35, 235, 120, 99, 266, 63, 236, 133, 412, 38], [555, 104, 14, 81, 55, 497, 222, 64, 57, 131] ]) # 构建值到索引的映射 max_val = lookup_table.max().item() index_map = torch.empty(max_val + 1, dtype=torch.long, device=lookup_table.device) index_map[lookup_table] = torch.arange(len(lookup_table), device=lookup_table.device) # 直接索引得到结果,注意把data转成long类型避免索引报错 result = index_map[data.long()] print(result)
输出和你给出的期望结果完全一致。
适用场景:lookup_table的最大值不超过1e6级别,内存占用可忽略,整体时间复杂度为O(L + N),L为lookup_table长度,N为data元素总数。
方案2:广播比较法(适用大值域场景)
如果lookup_table的值域过大,无法构建映射表,可以使用全向量化的广播比较方案,不需要提前预处理:
# 支持任意维度的data输入,无需修改代码 result = (lookup_table.view(*([1]*data.ndim), -1) == data.unsqueeze(-1)).argmax(dim=-1) print(result)
适用场景:lookup_table值域过大,但长度较小(一般不超过1000),整体时间复杂度为O(N*L),性能依然远高于手动循环实现。
注意事项
如果你的场景中存在data元素不在lookup_table里的情况,可以先用torch.isin(data, lookup_table)做合法性校验,避免返回异常值。
内容的提问来源于stack exchange,提问作者Lupos
相关产品推荐
相关产品推荐

