为何Numba JIT编译新函数会导致已有JIT函数性能下降?
Numba JIT函数性能莫名下降问题分析与解决
原函数及初始性能
我编写了一个经Numba JIT编译的高性能函数,用于生成三角形区域的坐标:
import numpy as np import numba as nb @nb.njit(cache=True, parallel=True, nogil=True) def triangle_half_UR_LL(size: int, swap: bool = False) -> tuple[np.ndarray, np.ndarray]: total = (size + 1) * size // 2 x_coords = np.full(total, 0, dtype=np.uint16) y_coords = np.full(total, 0, dtype=np.uint16) offset = 0 side = np.arange(size, dtype=np.uint16) for i in nb.prange(size): offset = i * size - (i - 1) * i // 2 end = offset + size - i x_coords[offset:end] = i y_coords[offset:end] = side[i:] return (x_coords, y_coords) if not swap else (y_coords, x_coords)
初始测试性能极佳,调用triangle_half_UR_LL(10)输出符合预期,且执行size=1000的性能稳定在166μs左右:
In [2]: triangle_half_UR_LL(10) Out[2]: (array([0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 6, 6, 6, 6, 7, 7, 7, 8, 8, 9], dtype=uint16), array([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 1, 2, 3, 4, 5, 6, 7, 8, 9, 2, 3, 4, 5, 6, 7, 8, 9, 3, 4, 5, 6, 7, 8, 9, 4, 5, 6, 7, 8, 9, 5, 6, 7, 8, 9, 6, 7, 8, 9, 7, 8, 9, 8, 9, 9], dtype=uint16)) In [3]: %timeit triangle_half_UR_LL(1000) 166 μs ± 489 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [4]: %timeit triangle_half_UR_LL(1000) 166 μs ± 270 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [5]: %timeit triangle_half_UR_LL(1000) 166 μs ± 506 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)
性能下降复现
当定义并调用另一个简单的Numba JIT函数后,原函数性能骤降,耗时从166μs飙升至970μs左右:
In [6]: @nb.njit(cache=True) ...: def dummy(): ...: pass In [7]: dummy() In [8]: %timeit triangle_half_UR_LL(1000) 980 μs ± 20 μs per loop (mean ± std. dev. of 7 runs, 1,000 loops each) In [9]: %timeit triangle_half_UR_LL(1000) 976 μs ± 9.9 μs per loop (mean ± std. dev. of 7 runs, 1,000 loops each) In [10]: %timeit triangle_half_UR_LL(1000) 974 μs ± 3.11 μs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
该问题可稳定复现:新启动解释器会话时原函数快速运行,调用任意其他JIT函数后性能立刻下降。
移除nogil=True后的表现
有趣的是,若移除原函数的nogil参数,性能下降问题完全消失:
In [1]: import numpy as np ...: import numba as nb ...: ...: ...: @nb.njit(cache=True, parallel=True) ...: def triangle_half_UR_LL(size: int, swap: bool = False) -> tuple[np.ndarray, np.ndarray]: ...: total = (size + 1) * size // 2 ...: x_coords = np.full(total, 0, dtype=np.uint16) ...: y_coords = np.full(total, 0, dtype=np.uint16) ...: offset = 0 ...: side = np.arange(size, dtype=np.uint16) ...: for i in nb.prange(size): ...: offset = i * size - (i - 1) * i // 2 ...: end = offset + size - i ...: x_coords[offset:end] = i ...: y_coords[offset:end] = side[i:] ...: ...: return (x_coords, y_coords) if not swap else (y_coords, x_coords) In [2]: %timeit triangle_half_UR_LL(1000) 186 μs ± 47.9 μs per loop (mean ± std. dev. of 7 runs, 1 loop each) In [3]: %timeit triangle_half_UR_LL(1000) 167 μs ± 1.61 μs per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [4]: %timeit triangle_half_UR_LL(1000) 166 μs ± 109 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [5]: @nb.njit(cache=True) ...: def dummy(): ...: pass In [6]: dummy() In [7]: %timeit triangle_half_UR_LL(1000) 167 μs ± 308 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [8]: %timeit triangle_half_UR_LL(1000) 166 μs ± 312 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [9]: %timeit triangle_half_UR_LL(1000) 167 μs ± 624 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each)
其他触发场景
除了调用简单的dummy函数,重新定义原函数,或调用其他JIT函数(如Farey_sequence)也会触发性能下降:
重新定义原函数触发
In [7]: dummy() In [8]: %timeit triangle_half_UR_LL(1000) 168 μs ± 750 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [13]: @nb.njit(cache=True, parallel=True) ...: def triangle_half_UR_LL(size: int, swap: bool = False) -> tuple[np.ndarray, np.ndarray]: ...: total = (size + 1) * size // 2 ...: x_coords = np.full(total, 0, dtype=np.uint16) ...: y_coords = np.full(total, 0, dtype=np.uint16) ...: offset = 0 ...: side = np.arange(size, dtype=np.uint16) ...: for i in nb.prange(size): ...: offset = i * size - (i - 1) * i // 2 ...: end = offset + size - i ...: x_coords[offset:end] = i ...: y_coords[offset:end] = side[i:] ...: ...: return (x_coords, y_coords) if not swap else (y_coords, x_coords) In [14]: %timeit triangle_half_UR_LL(1000) 1.01 ms ± 94.3 μs per loop (mean ± std. dev. of 7 runs, 1 loop each) In [15]: %timeit triangle_half_UR_LL(1000) 964 μs ± 2.02 μs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
调用Farey_sequence触发
Farey_sequence函数定义:
@nb.njit(cache=True) def Farey_sequence(n: int) -> np.ndarray: a, b, c, d = 0, 1, 1, n result = [(a, b)] while 0 <= c <= n: k = (n + b) // d a, b, c, d = c, d, k * c - a, k * d - b result.append((a, b)) return np.array(result, dtype=np.uint64)
调用后原函数性能下降:
In [6]: %timeit triangle_half_UR_LL(1000) 166 μs ± 296 ns per loop (mean ± std. dev. of 7 runs, 10,000 loops each) In [7]: %timeit Farey_sequence(16) The slowest run took 6.25 times longer than the fastest. This could mean that an intermediate result is being cached. 6.03 μs ± 5.72 μs per loop (mean ± std. dev. of 7 runs, 1 loop each) In [8]: %timeit Farey_sequence(16) 2.77 μs ± 50.8 ns per loop (mean ± std. dev. of 7 runs, 100,000 loops each) In [9]: %timeit triangle_half_UR_LL(1000) 966 μs ± 6.48 μs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
原因分析
这是Windows环境下Numba的一个类已知问题,根源在于parallel=True和nogil=True组合使用时,Numba的线程池管理机制存在冲突。当第一个并行+无GIL的函数运行时,Numba初始化了高效的线程池;但后续编译执行其他非并行JIT函数时,Numba的线程池状态被修改,导致原并行函数的线程调度开销大幅增加——原本复用线程池的逻辑变成了频繁创建/销毁线程,直接拉低了执行效率。
解决办法
- 移除
nogil=True参数:从测试结果看,移除该参数后,原函数性能几乎不受影响,同时彻底避免了后续性能下降的问题。这是最直接有效的方案。 - 预编译所有JIT函数:在程序启动阶段,提前编译并调用所有需要用到的Numba JIT函数,确保线程池状态在程序运行过程中不再发生变化。
- 显式指定线程数:使用
nb.set_num_threads()强制Numba使用固定数量的线程,减少线程池动态调整带来的开销,示例:import numba as nb nb.set_num_threads(nb.config.NUMBA_DEFAULT_NUM_THREADS)
内容的提问来源于stack exchange,提问作者Ξένη Γήινος
相关产品推荐
相关产品推荐

