numpy数组循环优化方案选型:numba与dask哪种效果更佳?
问题背景
当前需要重构以下代码,尽可能降低运行时长和内存占用:
for i in range(gbl.NumStoreRows): cal_effects[i,:,:len(orig_cols)] = cal_effects_vals # 这一行占用约1GB内存 priors[i,:len(orig_cols)] = orig_prior_coeffs priors_SE[i,:len(orig_cols)] = orig_prior_SE
循环中仅第一个操作耗时和内存占用较高,尝试将耗资源的行和另外两行拆分为两个独立循环,反而慢了1秒,内存占用也无改善。后续尝试编写JIT函数优化时运行报错,停在LoadFunctions()步骤。
已尝试的JIT变体
变体1
@jit def populate_cal_effects(cal_effects_vals): for i in range(gbl.NumStoreRows): cal_effects[i,:,:len(orig_cols)] = cal_effects_vals populate_cal_effects(cal_effects_vals) for i in range(gbl.NumStoreRows): priors[i,:len(orig_cols)] = orig_prior_coeffs priors_SE[i,:len(orig_cols)] = orig_prior_SE
变体2:添加返回语句
@jit def populate_cal_effects(cal_effects_vals): for i in range(gbl.NumStoreRows): cal_effects[i,:,:len(orig_cols)] = cal_effects_vals return cal_effects[i,:,:len(orig_cols)]
变体3:合并所有操作+并行
预期该方案速度最快且不影响输出,但无法正常运行:
@jit(parallel=True) def populate_cal_effects(cal_effects_vals): for i in prange(gbl.NumStoreRows): cal_effects[i,:,:len(orig_cols)] = cal_effects_vals priors[i,:len(orig_cols)] = orig_prior_coeffs priors_SE[i,:len(orig_cols)] = orig_prior_SE
上下文与复现条件
- JIT函数目前定义在主加载函数内部,后续计划移出Load函数后重试
- 若JIT方案无效,考虑尝试用Dask实现单机多核并行处理
- 固定参数:
gbl.NumstoreRows = 866(门店数量) - 所有数据均为numpy数组:
cal_effects = np.zeros((gbl.NumStoreRows, n_days, n_cal_effects), dtype=np.float64) priors = np.zeros((gbl.NumStoreRows, n_cal_effects), dtype=np.float64) priors_SE = np.zeros((gbl.NumStoreRows, n_cal_effects), dtype=np.float64)
优化建议
- 优先使用NumPy原生广播,完全消除循环
这是性能最高、内存占用最低的方案,不需要循环或额外工具。直接对整个数组切片赋值,NumPy会自动沿第一个维度广播,避免循环过程中产生的大量中间切片临时对象,内存占用仅为原循环方案的1/3不到,运行速度快数个量级。
n_orig = len(orig_cols) # 直接广播赋值,无需循环 cal_effects[:, :, :n_orig] = cal_effects_vals priors[:, :n_orig] = orig_prior_coeffs priors_SE[:, :n_orig] = orig_prior_SE
你之前拆分两个循环变慢的核心原因是两次遍历数组行维度,CPU缓存命中率降低,额外增加了内存访问开销。
- 如果一定要用Numba JIT,修复如下问题即可正常运行
之前JIT报错的核心原因是函数内部调用了大量全局变量(gbl、orig_cols、cal_effects、priors等),Numba的nopython模式对全局可变对象支持很差,尤其是JIT函数定义在其他函数内部时,作用域解析会出错。并行模式需要明确开启nopython,且prange需要从numba导入,所有数组要作为参数传入函数,不要用全局变量,正确写法参考:
from numba import jit, prange # 所有依赖变量都作为参数传入,不要用全局值 @jit(nopython=True, parallel=True) def populate_arrays(cal_effects, priors, priors_SE, cal_effects_vals, orig_prior_coeffs, orig_prior_SE, n_orig, num_rows): for i in prange(num_rows): cal_effects[i, :, :n_orig] = cal_effects_vals priors[i, :n_orig] = orig_prior_coeffs priors_SE[i, :n_orig] = orig_prior_SE # 调用时传入所有参数 n_orig = len(orig_cols) populate_arrays(cal_effects, priors, priors_SE, cal_effects_vals, orig_prior_coeffs, orig_prior_SE, n_orig, gbl.NumStoreRows)
额外内存优化
如果业务精度允许,将数组的dtype从float64改为float32,内存占用直接减半,运行速度也会提升30%以上。不需要使用Dask
当前数据量极小(仅866个门店维度),NumPy原生操作已经能在几毫秒内完成,Dask的调度开销远大于执行收益,完全没必要使用。
内容的提问来源于stack exchange,提问作者Paul Russell
相关产品推荐
相关产品推荐

