如何快速批量计算整数范围内自定义递归互调函数的结果?
问题
给定以下支持互相调用的Python递归函数:
QUX = 100 QUUX = 250 def foo(x): if x == 0: return QUX return foo(x - 1) - bar(x) def bar(x): if x == 0: return 0 return QUUX - baz(x) def baz(x): if x == 0: return 0 return foo(x - 1) * 0.01
需完成两项任务:
- 计算函数在
x ∈ [0, 1000]范围内的所有结果; - 修改常量
QUX、QUUX后,重复上述计算200000次。
已尝试lru_cache、map、numpy vectorize、Pybind11等方案,但性能提升有限,期望找到接近C/C++速度的高效评估方法。
高效解决方案
1. 数学推导化简(性能最优)
先将递归关系转化为递推公式,彻底消除递归调用开销:
从函数定义逐步代入推导:
- 当
x ≥ 1时:baz(x) = 0.01 * foo(x-1)bar(x) = QUUX - baz(x) = QUUX - 0.01*foo(x-1)foo(x) = foo(x-1) - bar(x) = foo(x-1)*(1 + 0.01) - QUUX
最终得到仅依赖前序值的递推公式:
foo(0) = QUX foo(x) = 1.01 * foo(x-1) - QUUX (x ≥ 1) bar(x) = QUUX - 0.01*foo(x-1) (x ≥1),bar(0)=0 baz(x) = 0.01*foo(x-1) (x≥1),baz(0)=0
基于此编写纯循环计算代码,无任何递归或额外函数调用,性能接近C/C++:
def compute_all(QUX, QUUX, max_x=1000): foo_results = [0.0] * (max_x + 1) foo_results[0] = QUX bar_results = [0.0] * (max_x + 1) baz_results = [0.0] * (max_x + 1) for x in range(1, max_x + 1): foo_prev = foo_results[x-1] foo_results[x] = 1.01 * foo_prev - QUUX baz_results[x] = 0.01 * foo_prev bar_results[x] = QUUX - baz_results[x] return foo_results, bar_results, baz_results # 单次计算示例 foo_res, bar_res, baz_res = compute_all(100, 250) # 重复200000次示例(可根据需求修改QUX/QUUX的生成逻辑) for _ in range(200000): qux = 100 + _ % 50 # 示例修改逻辑 quux = 250 + _ % 100 compute_all(qux, quux)
2. Numba即时编译(保留Python语法,接近C速度)
若不想手动推导公式,可使用Numba将Python代码编译为机器码,消除解释器开销:
from numba import njit @njit(fastmath=True) def compute_all_numba(QUX, QUUX, max_x=1000): foo_results = [0.0] * (max_x + 1) foo_results[0] = QUX bar_results = [0.0] * (max_x + 1) baz_results = [0.0] * (max_x + 1) for x in range(1, max_x + 1): foo_prev = foo_results[x-1] foo_results[x] = 1.01 * foo_prev - QUUX baz_results[x] = 0.01 * foo_prev bar_results[x] = QUUX - baz_results[x] return foo_results, bar_results, baz_results # 使用方式与纯Python版本一致 foo_res, bar_res, baz_res = compute_all_numba(100, 250)
Numba首次调用时完成编译,后续调用均为机器码执行,性能与C基本持平。
3. Cython静态编译(极致性能)
若需要绝对极致的速度,可将代码转为Cython并编译为C扩展:
编写compute.pyx文件:
def compute_all_cython(double QUX, double QUUX, int max_x=1000): cdef int x cdef double[:] foo_results = [0.0]*(max_x + 1) cdef double[:] bar_results = [0.0]*(max_x + 1) cdef double[:] baz_results = [0.0]*(max_x + 1) cdef double foo_prev foo_results[0] = QUX for x in range(1, max_x + 1): foo_prev = foo_results[x-1] foo_results[x] = 1.01 * foo_prev - QUUX baz_results[x] = 0.01 * foo_prev bar_results[x] = QUUX - baz_results[x] return list(foo_results), list(bar_results), list(baz_results)
通过setup.py编译为C扩展后,调用速度与原生C完全一致。
原方案性能受限原因
lru_cache:仅避免重复计算,但Python递归的栈帧创建、函数调用开销依然存在;numpy vectorize:本质仍是Python循环,未实现真正的底层加速;Pybind11:若基于递归实现,C递归仍有开销,且跨语言调用会增加额外成本;若改为递推实现,性能与数学推导后的Python+Numba接近,但需要编写C代码,成本更高。
内容的提问来源于stack exchange,提问作者zchmielewska
相关产品推荐
相关产品推荐

