Python numba JIT 编译实战:为什么它有时比 numpy 慢 7%,有时却快 41 倍(2026)

📝 460 字 · ☕ 2 分钟阅读

先说个我自己翻车的事。上周我在给一个量化回测脚本提速,里面有一段期权定价的循环,跑了 2.8 秒。我第一反应是上 numba,@njit 一加,编译一过,你猜怎么着——居然比纯 numpy 向量化还慢了 7%

我当时有点懵。网上一堆文章把 numba 吹成”一行加速 100 倍”的神器,怎么到我这儿就失效了?后来我把它和 numpy 在两种完全不同类型的计算上做了个对比,才彻底搞明白 numba 的脾气:它只在特定场景下快,用错地方就是纯纯的负优化。这篇文章就把这两组真实基准数据摊开给你看。

numba 到底干了什么

numba 是一个 JIT(Just-In-Time,即时编译)编译器。普通 Python 代码交给解释器,逐行翻译成字节码执行;numba 的做法是:第一次调用时,把你标记了 @njit 的函数用 LLVM 直接编译成机器码,之后就用编译好的机器码跑,不再经过解释器。

这里的关键差异有两个:

  • 机器码 vs 解释执行:CPU 直接跑编译后的指令,绕开了 Python 的对象装箱、类型检查、垃圾回收这些运行时开销。
  • 循环融合:numba 会把整个循环体当成一个整体优化,不产生中间数组;而 numpy 每做一步向量化运算,都会在内存里新建一个临时数组。

这第二点,正是 numba 和 numpy 分道扬镳的根源。下面用两组基准看清楚。

场景一:Black-Scholes 期权定价(可以向量化)

Black-Scholes 公式是期权定价的经典模型,正好也是我们做红利策略复盘时经常碰到的东西。输入是 100 万份期权的标的价 S、行权价 K、到期时间 T,输出是每份期权的理论价格。

这段计算是标准的”逐元素独立运算”——第 i 份期权的价格只依赖 S[i]、K[i]、T[i],互相之间没有依赖。这种结构天然适合向量化。四套实现我都跑了一遍,全部用同一个标准正态分布 CDF 近似(Abramowitz & Stegun 7.1.26),保证对比公平:

import numpy as np
from numba import njit, prange

# 纯 Python 循环版(baseline)
def bs_call_python(S, K, T, r, sigma):
    out = np.empty(len(S))
    for i in range(len(S)):
        s, k, t = S[i], K[i], T[i]
        d1 = (np.log(s/k) + (r + 0.5*sigma**2)*t) / (sigma*np.sqrt(t))
        d2 = d1 - sigma*np.sqrt(t)
        out[i] = s * norm_cdf(d1) - k * np.exp(-r*t) * norm_cdf(d2)
    return out

# numpy 向量化版
def bs_call_numpy(S, K, T, r, sigma):
    d1 = (np.log(S/K) + (r + 0.5*sigma**2)*T) / (sigma*np.sqrt(T))
    d2 = d1 - sigma*np.sqrt(T)
    return S * norm_cdf(d1) - K * np.exp(-r*T) * norm_cdf(d2)

# numba JIT 串行版
@njit(cache=True)
def bs_call_numba(S, K, T, r, sigma):
    out = np.empty(len(S))
    for i in range(len(S)):
        s, k, t = S[i], K[i], T[i]
        d1 = (np.log(s/k) + (r + 0.5*sigma**2)*t) / (sigma*np.sqrt(t))
        d2 = d1 - sigma*np.sqrt(t)
        out[i] = s * norm_cdf(d1) - k * np.exp(-r*t) * norm_cdf(d2)
    return out

结果(100 万份,取 5 次最优):

实现方式 耗时 相对纯 Python 加速
纯 Python 循环 2855.7 ms 1.0×
numpy 向量化 42.8 ms 66.7×
numba @njit 串行 45.9 ms 62.2×
numba 并行(prange) 31.8 ms 89.9×

看到没?numba 串行版比 numpy 还慢 7%。原因就是那句话:这种逐元素独立、没有分支的运算,numpy 的 C 底层循环已经优化到极致了,numba 编译出来的机器码跟它半斤八两,反而还多了一点调用开销。只有当 numba 开了 parallel=True 用多核并行时,才靠着 CPU 核心数把 45.9ms 压到 31.8ms,勉强赢了 numpy。

所以结论很清楚:能向量化的东西,直接用 numpy 就够了,numba 帮不上多少忙

场景二:Mandelbrot 分形(向量化不动的循环)

真正让 numba 大放异彩的,是 numpy 搞不定的循环。最典型的例子就是 Mandelbrot 分形——生成下面这张图:

用 numba 并行生成的 Mandelbrot 分形图
用 numba 并行生成的 Mandelbrot 分形(1600×1600,inferno 配色)

Mandelbrot 的循环长这样:每个像素点要反复迭代 z = z² + c,直到逃逸半径超过 2 或者达到最大迭代次数才停。问题的关键在“直到……才停”——每个像素的迭代次数都不一样,有的点 3 次就逃逸了,有的点要算满 256 次。

@njit(cache=True, parallel=True)
def mandelbrot(w, h, max_iter):
    out = np.zeros((h, w), dtype=np.int32)
    dx = (xmax - xmin) / w
    dy = (ymax - ymin) / h
    for row in prange(h):
        cy = ymin + row * dy
        for col in range(w):
            cx = xmin + col * dx
            x, y, it = 0.0, 0.0, 0
            while x*x + y*y <= 4.0 and it < max_iter:
                x, y = x*x - y*y + cx, 2.0*x*y + cy
                it += 1
            out[row, col] = it
    return out

numpy 想向量化它,只能”全像素同步迭代”——每个像素都硬算满 max_iter 次,用掩码把已经逃逸的点标出来。问题是大部分像素早就逃逸了,你还在白白给它们算,纯属浪费。结果(1600×1600,取 3 次最优):

实现方式 耗时 相对纯 Python 加速
纯 Python 循环 15.21 s 1.0×
numpy 朴素向量化 2.79 s 5.5×
numba @njit 串行 0.37 s 41.3×
numba 并行(prange) 0.22 s 70.1×

这组数据就漂亮多了:numba 串行版比 numpy 快 7.5 倍,并行版再翻 1.7 倍,相比纯 Python 直接 70 倍。两组场景放在一起看最直观:

numba 在可向量化与不可向量化两种场景下的加速对比柱状图
两种场景的耗时对比(对数刻度):可向量化时 numpy 与 numba 打平,不可向量化时 numba 领先一个数量级

为什么差距差这么多

根本原因就一条:numpy 的向量化,本质是把循环搬进 C 语言里跑,但它没法”提前退出”

  • 可向量化场景(场景一):每个元素独立、无分支、无提前退出。numpy 的 C 循环和 numba 的 LLVM 机器码能力相当,numba 没有优势,还多出函数调用开销。
  • 不可向量化场景(场景二):循环里有 while 条件、数据依赖的分支。numpy 被迫全量计算,numba 却能逐元素提前退出,这一进一出就是 7 倍的差距。

另外还有个 numba 隐藏的加分项:循环融合不产生中间数组。numpy 写一句 S/K 就会分配一个新数组,复杂表达式层层套下来,内存峰值和 cache 命中率都会变差。numba 把整段循环编译成一个函数,中间变量都留在寄存器或 L1 cache 里。数据量越大,这个差异越明显。

什么时候用 numba:一个简单的判断

别一上来就 numba,先问自己三个问题:

  1. 这段循环能用 numpy 向量化吗? 能,就先向量化——大概率已经够快,代码还更好维护。
  2. 循环里有没有分支、提前退出、状态依赖? 有,numba 才是正解。比如逐元素的数值迭代、蒙特卡洛路径模拟、图遍历、分形生成。
  3. 是不是在热路径上,且数据量够大? numba 有编译开销(首次调用几百毫秒到几秒),数据量小、只跑一两次的话,编译时间比省下的时间还多,得不偿失。

顺带一提,如果你在做量化或金融计算,numba 有个亲兄弟叫 @guvectorize,可以给 numpy 数组写自定义的通用函数(ufunc),两者配合能覆盖绝大多数数值密集型场景。这块跟之前聊过的 Python 并发编程的 GIL 选型是一体两面:numba 解决的是”单核内把循环跑快”,多进程/多线程解决的是”怎么用满多核”。

numba 的几个坑,提前帮你踩了

  1. 编译开销:第一次调用会卡一下,@njit(cache=True) 可以把编译结果缓存到磁盘,第二次启动直接加载。生产环境务必开 cache。
  2. 类型不稳定会拖慢:numba 靠类型推断,如果一个函数里同一个变量一会儿是 int 一会儿是 float,它会退回到 object 模式,速度暴跌甚至比纯 Python 还慢。用 numba.core.dispatcher@njit 后看 .inspect_types() 能查出来。
  3. 不是所有 Python 语法都支持:字典、集合、字符串处理的支持很有限,标准库函数也只支持一部分。写 numba 代码时脑子里要切到”C 语言模式”,只用标量运算、数组和简单循环。
  4. prange 不是免费午餐:并行有线程启动和调度开销,循环体太轻的话,开了 parallel=True 反而更慢。我这个 Mandelbrot 例子每行计算量够重,才吃到 1.7× 的并行收益。

常见问题(FAQ)

Q: numba 和 numpy 到底该先学哪个?

先 numpy。80% 的性能问题靠向量化就能解决,numba 是向量化解决不了时的后手。把 numpy 的广播、切片、ufunc 用熟练,再碰 numba,否则容易本末倒置,写出又难维护又没快多少的代码。

Q: numba 编译后的代码能快到 C 的水平吗?

接近,但通常差一点。numba 用 LLVM 生成机器码,理论上和手写 C 同档,但受限于 Python 的调用约定和边界处的对象转换,实际性能一般是手写 C 的 80%~95%。对绝大多数场景,这个差距可以忽略。

Q: 我的循环里有 if 判断,numpy 就真的没辙了吗?

可以用 np.where、布尔掩码、np.select 把很多分支转成向量化表达,这些能覆盖相当一部分场景。真正搞不定的,是”迭代次数依赖数据、需要提前退出”这种结构(比如 Mandelbrot、蒙特卡洛路径),那才是 numba 的主场。

Q: 怎么判断 numba 有没有真正生效?

看两个信号:一是函数第一次调用有明显的编译延迟,二是 func.inspect_types() 里没有出现 object 模式(出现说明类型推断失败,速度会很差)。另外可以和 numpy 向量化版跑个基准对比,如果没比 numpy 快,说明这活儿可能根本不适合 numba。

总结

回到开头那个翻车现场:我给期权定价加 numba,结果比 numpy 还慢 7%,不是 numba 不行,是我用错了地方。numba 不是”加一行就 100 倍”的银弹,它是 numpy 向量化够不着时的那把趁手工具——循环里有分支、有提前退出、有状态依赖,才是它的主场。

判断标准记住一句话:先向量化,向量化不动再上 numba,热路径加 cache,重计算才开 prange。关于 numpy 侧的更多提速技巧,可以接着看这篇 Pandas 性能优化实战,以及这篇讲”你以为很快其实很慢”的 Python 性能翻车现场,搭配着读效果更好。

📤 分享这篇文章