如何用Numpy高效获取满足容量限制的作业组合?
基于Numpy的高效作业组合筛选方案
核心思路是利用二进制掩码枚举所有非空组合,通过Numpy的向量化计算快速筛选出总大小符合容量限制的组合,彻底避开Python循环或itertools迭代的低效问题。
具体实现步骤
以你的示例数据为例:
import numpy as np # 定义输入参数 jobs = np.array([1, 3, 5]) job_sizes = np.array([1, 3, 2]) capacity = 5
- 生成所有非空组合的二进制掩码
每个作业有「选」或「不选」两种状态,对应二进制位的1或0。我们生成从1到2^n - 1的整数(跳过全0的空组合),再将这些整数转换为二进制矩阵,每行代表一个组合的选择状态:
n = len(jobs) # 生成所有非空组合的索引(1到2^n -1) mask_indices = np.arange(1, 2 ** n) # 转换为二进制掩码矩阵,每行对应一个组合的选择状态 masks = np.unpackbits(mask_indices.view('uint8'), axis=1)[:, -n:]
- 向量化计算组合总大小
用矩阵点积一次性算出所有组合的总大小,这一步是Numpy效率的核心——完全是C级别的批量计算:
total_sizes = masks @ job_sizes
- 筛选符合容量限制的组合
通过布尔索引快速过滤出总大小不超过容量的掩码:
valid_masks = masks[total_sizes <= capacity]
- 映射为作业编号组合
根据筛选后的掩码,提取对应的作业编号,生成最终的组合列表:
valid_combinations = [tuple(jobs[mask == 1]) for mask in valid_masks]
运行后得到的结果完全符合你的预期:
[(1,), (3,), (5,), (1, 3), (1, 5), (3, 5)]
效率优势说明
- 对比
itertools.combinations:后者需要遍历所有长度的组合(1个元素、2个元素...n个元素),每个组合都要单独计算总大小,本质是Python层面的循环;而Numpy的向量化操作直接在底层批量处理所有组合,速度提升非常明显,尤其是当作业数量n≥15时,差距会呈指数级拉大。 - 内存注意点:当n超过25时,
2^25约3300万条数据,会占用较多内存。如果遇到这种场景,可以分块生成掩码(比如每次处理100万条),避免内存溢出。
内容的提问来源于stack exchange,提问作者Paulo Nascimento
相关产品推荐
相关产品推荐

