如何用Numpy风格方法计算不等长带权重数组的交集规模
用Numpy风格方法计算带权重不等长数组的交集加权和
问题描述
现有两个整数类型的Numpy数组,格式如下:
arr1=[[item11,count11],[item12,count12],[item13,count13],...] arr2=[[item21,count21],[item22,count22],[item23,count23],...]
这类数组用于记录购物清单汇总,每个元素[itemX, countY]表示某人购买了countY件itemX。两个数组长度不同且未排序,未购买的物品不会出现在清单中。
需求是统计同时出现在两个数组中的物品,将每个共同物品的两个计数取最小值后求和。例如:若arr1中有[item1,count1]、arr2中有[item1,count2],则将min(count1,count2)加入总和。
非Numpy实现代码
count = 0 for i in range(len(arr1)): for j in range(len(arr2)): if arr1[i][0] == arr2[j][0]: count += min(arr1[i][1], arr2[j][1]) return count
示例
arr1 = [[1,10],[2,100],[3,1000],[4,10000]] arr2 = [[1,10],[3,100],[4,1000],[5,10000],[6,99]]
该示例应返回1110,因为物品1取10、物品3取100、物品4取1000,三者求和为10+100+1000=1110。
Numpy实现方案
方案一:直观向量化实现(适合小数据量)
利用Numpy的数组操作替代嵌套循环,步骤清晰:
import numpy as np # 示例数组 arr1 = np.array([[1,10],[2,100],[3,1000],[4,10000]]) arr2 = np.array([[1,10],[3,100],[4,1000],[5,10000],[6,99]]) # 拆分物品ID与对应计数 items1, counts1 = arr1[:, 0], arr1[:, 1] items2, counts2 = arr2[:, 0], arr2[:, 1] # 获取两个数组的共同物品ID common_items = np.intersect1d(items1, items2) # 计算最小值总和 total = 0 for item in common_items: # 提取当前物品在两个数组中的计数 cnt1 = counts1[items1 == item][0] cnt2 = counts2[items2 == item][0] total += min(cnt1, cnt2) print(total) # 输出:1110
方案二:完全向量化实现(适合大数据量)
通过结构化数组和批量操作进一步提升效率,避免Python层面的循环:
import numpy as np # 示例数组 arr1 = np.array([[1,10],[2,100],[3,1000],[4,10000]]) arr2 = np.array([[1,10],[3,100],[4,1000],[5,10000],[6,99]]) # 转换为结构化数组,按物品ID和计数分组 struct1 = np.array(list(zip(arr1[:,0], arr1[:,1])), dtype=[('item', int), ('count', int)]) struct2 = np.array(list(zip(arr2[:,0], arr2[:,1])), dtype=[('item', int), ('count', int)]) # 筛选同时存在于两个数组的物品 common_mask = np.isin(struct1['item'], struct2['item']) common_struct1 = struct1[common_mask] common_struct2 = struct2[np.isin(struct2['item'], common_struct1['item'])] # 按物品ID排序,确保计数位置对应 common_struct1_sorted = np.sort(common_struct1, order='item') common_struct2_sorted = np.sort(common_struct2, order='item') # 批量取最小值并求和 total = np.sum(np.minimum(common_struct1_sorted['count'], common_struct2_sorted['count'])) print(total) # 输出:1110
内容的提问来源于stack exchange,提问作者aellab
相关产品推荐
相关产品推荐

