Python中基于Numpy/PyTorch实现列表元素在数组行中的定位
解决方案:基于Numpy/PyTorch的向量化实现
Numpy 版本
利用广播机制实现向量化比较,避免循环,效率更高:
import numpy as np a = np.array([1, 2, 3]) b = np.array([[2, 4, 6],[3, 2, 5],[4, 1, 3]]) # 广播生成每行与对应a元素的匹配矩阵 matches = b == a[:, np.newaxis] # 对每行:有匹配则取第一个匹配的索引,无匹配则设为-1 c = np.where(matches.any(axis=1), np.argmax(matches, axis=1), -1) print(c.tolist()) # 输出: [-1, 1, 2]
原理说明:
- 将
a转为列向量(a[:, np.newaxis]),和二维数组b广播比较,得到一个形状与b一致的布尔矩阵,每个位置表示b[i,j]是否等于a[i] matches.any(axis=1)检查每行是否存在匹配项np.argmax(matches, axis=1)返回每行第一个True的索引(因为argmax会优先取首个最大值位置,布尔值中True等价于1)- 最后用
np.where完成条件赋值,无匹配的位置设为-1
PyTorch 版本
如果用PyTorch处理,思路和Numpy一致,利用张量的广播与向量化操作:
import torch a = torch.tensor([1, 2, 3]) b = torch.tensor([[2, 4, 6],[3, 2, 5],[4, 1, 3]]) # 广播生成匹配张量 matches = b == a.unsqueeze(1) # 取每行首个匹配的索引 c = torch.argmax(matches.int(), dim=1) # 将无匹配的位置替换为-1 c = torch.where(matches.any(dim=1), c, torch.tensor(-1, dtype=c.dtype)) print(c.tolist()) # 输出: [-1, 1, 2]
原理说明:
a.unsqueeze(1)将一维张量转为列张量,实现和b的广播比较matches.int()把布尔张量转为整数张量(True→1,False→0),才能用torch.argmax取首个匹配位置torch.where根据每行是否有匹配的条件,保留索引或替换为-1
这两种向量化方案的时间复杂度都是O(n*m)(n为行数,m为列数),但底层是C/CUDA实现,比Python循环快得多,尤其适合大规模数据场景。
内容的提问来源于stack exchange,提问作者vendrick17
相关产品推荐
相关产品推荐

