如何在NumPy计算标量序列与标量组合时实现高效过滤
问题背景
我需要生成a、b、c、x、y、z六个参数的所有可能组合,其中x、y、z支持数组或float类型输入。目前我通过以下代码实现了组合生成能力,兼容不同类型的入参:
from typing import Union, Sequence import numpy as np from numbers import Real def cartesian_product(*arrays: np.ndarray) -> np.ndarray: la = len(arrays) dtype = np.result_type(*arrays) arr = np.empty([len(a) for a in arrays] + [la], dtype=dtype) for i, a in enumerate(np.ix_(*arrays)): arr[..., i] = a return arr.reshape(-1, la) def iter_func( *args: Union[Real, Sequence[Real], np.ndarray], ) -> np.ndarray: return cartesian_product(*( np.atleast_1d(a) for a in args ))
运行iter_func(5,[2,3],3,[3,6,9],2,[1,2,4])会生成完整的笛卡尔积数组。
后续我需要对每个组合做运算,追加运算结果后过滤出符合规则的组合,目前的实现逻辑如下:
# 示例运算函数 def Operation(args): args=args.tolist() a,b,c,*args = args return (a+b+c)/sum(args) def Operation2(args): args=args.tolist() a,b,c,*args = args return (a/b/c) # 全量生成组合后遍历校验、拼接 new_list = [np.append(element,[Operation(element),Operation2(element)]) for element in iter_func(5,[2,3],3,[3,6,9],2,[1,2,4]) if 0.7<Operation(element)<1.2 and 0.55<Operation2(element)<0.85]
上述逻辑可以正常得到结果,但存在两个问题:一是需要先生成全量组合数组,参数多的时候内存占用极高;二是每个组合需要重复计算两次运算值(一次用于条件判断,一次用于结果拼接),运算效率低。
我希望实现:生成单个组合的同时就完成运算、校验,符合条件的才留存,不需要缓存全量组合,提升运行效率、降低内存占用。
解决方案
可以用Python内置的itertools.product实现惰性生成组合,它不会一次性生成所有组合存入内存,每次迭代只生成一个组合,刚好适配你的需求,同时可以优化运算逻辑,每个组合仅做一次运算。
完整实现代码如下:
from typing import Union, Sequence, Callable import numpy as np from numbers import Real import itertools # 优化运算函数,无需转list,直接用numpy索引取值,运算更快 def Operation(args: np.ndarray) -> float: a,b,c = args[:3] return (a + b + c) / args[3:].sum() def Operation2(args: np.ndarray) -> float: a,b,c = args[:3] return a / b / c def lazy_filtered_combinations( *args: Union[Real, Sequence[Real], np.ndarray], filter_func: Callable[[float, float], bool] ) -> np.ndarray: # 统一入参格式,转成可迭代的列表 iterable_args = [np.atleast_1d(arg).tolist() for arg in args] # 惰性迭代生成组合,不会一次性生成全量数据 for combo in itertools.product(*iterable_args): combo_arr = np.array(combo) # 每个组合仅计算一次两个运算值 op1_res = Operation(combo_arr) op2_res = Operation2(combo_arr) # 校验过滤 if filter_func(op1_res, op2_res): # 返回拼接后的结果 yield np.append(combo_arr, [op1_res, op2_res]) # 定义过滤规则 def my_filter(op1: float, op2: float) -> bool: return 0.7 < op1 < 1.2 and 0.55 < op2 < 0.85 # 调用示例 if __name__ == "__main__": # 迭代获取结果,全程不会生成全量组合数组 for res in lazy_filtered_combinations(5,[2,3],3,[3,6,9],2,[1,2,4], filter_func=my_filter): print(res) # 如果需要最终结果转成numpy数组,直接转为list再封装即可 # result_arr = np.array(list(lazy_filtered_combinations(5,[2,3],3,[3,6,9],2,[1,2,4], filter_func=my_filter)))
方案优势
- 内存占用极低:仅在迭代时生成单个组合,参数取值再多也不会出现内存溢出问题
- 运算效率更高:每个组合仅做一次运算,相比原实现减少了一半的重复运算量
- 灵活性更强:如果不需要获取所有符合条件的结果,可以随时终止迭代,避免无效计算
内容的提问来源于stack exchange,提问作者nuwe
相关产品推荐
相关产品推荐

