尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Numba并行计算官方文档解读:从njit到prange的实践

Numba并行计算官方文档解读:从njit到prange的实践 简介Numba官方文档中文翻译包定位为Python科学计算与并行计算学习资料面向想用JIT编译加速NumPy代码、借助prange进行多核并行、或通过CUDA调用GPU算力的开发者。包内将官方文档整理为80个markdown文件另有1个HTML导读页和1张Logo图压缩包共82个文件、整体约296KB纯文本为主体积轻量便于随时查阅与检索。内容包括安装指南、五分钟快速上手、jit装饰器详解、vectorize与generated_jit高级用法、prange并行循环、CUDA编程及常见局限说明等基本覆盖Numba核心功能。文档按序号组织目录结构清晰方便按需跳转。已有554人学习下载适合正在优化数值计算代码、希望系统掌握Numba并行加速技巧的Python开发者。通过学习这套翻译文档可减少理解官方英文资料的时间成本快速将并行计算与GPU加速方法用于实际项目。1. 为什么读 Numba 官方文档的人最后都卡在并行计算这一章Numba 是 Python 生态里离“直接跑出 C 性能”最近的一条路用 LLVM 把 Python 函数编译成机器码并在编译阶段完成类型推导。对数值计算来说这个机制的收益比多进程更直接因为它不涉及序列化和进程同步。而官方文档的英文版更像是一个参考手册它把类型、编译模式、并行机制讲得很细但初次接触的人很容易在 jit、njit、vectorize、prange 这几个入口里迷路。中文翻译的意义不是把句子逐字翻译一遍而是把“什么场景该读哪一节、哪个参数决定性能”这件事理顺。下面按“装起来、跑起来、开并行、查问题”的顺序把 Numba 官方文档里和并行计算相关的部分拆开讲。2. 安装与第一个可运行的 Numba 函数把 Quickstart 跑通再谈并行2.1 用 conda 建一个干净的 Python 环境规避 llvmlite 版本坑Numba 绑定 LLVM 和 numpy 版本所以它不是一个“pip install 万事大吉”的库。常见做法是先确认 Python 版本。我一般建议用 conda 环境因为 conda 会同时把 llvmlite 的二进制包装好用 pip 安装也完全可以但版本匹配要自己注意。先把环境建好conda create -n numba-env python3.11 conda activate numba-env conda install numba -c numba如果用 pip下面这段也能用python -m venv numba-env source numba-env/bin/activate pip install numba安装完成后验证一下版本顺带看 numpy 和 llvmlite 的匹配情况python -c import numba; print(numba.__version__) python -c import llvmlite; print(llvmlite.__version__) python -c import numpy; print(numpy.__version__)这里有一个隐藏约束Numba 的版本会和 numpy 的 major 版本强绑定。比如新版本往往要求 numpy1.22如果环境里 numpy 太老import 时就会报ImportError。遇到这种情况不要直接降 numpy而是把整个环境的 python 和 numba 一起升。官方文档的 Installation 页面里有一张版本兼容表这张表值得先截图存好。另外Python 3.13 后的版本往往要等 Numba 两三个迭代才能支持生产环境别急着升级 Python 小版本。这种环境隔离问题在中文社区特别常见因为教程里有时直接用pip install numba没先建虚拟环境。如果pip和python指向的不是同一份解释器import numba 会报 ModuleNotFoundError。可以用python -m pip代替裸pip这样 pip 一定会安装到当前python对应的环境里。另外切换 conda 环境后先跑一次python -c import numba确认没有报错再进入正式代码这一步能省掉大部分环境类排障。建完环境先不要急着写算法还要确认 CPU 指令集。如果 import numba 时报Illegal instruction (core dumped)通常是编译产物里包含当前 CPU 不支持的指令。这时候设置环境变量把 CPU 特性限制到当前机器支持的范围内export NUMBA_CPU_NAMEx86-64 export NUMBA_CPU_FEATURESsse,sse2这个设置在旧 CPU 或云服务器上很常见。官方文档的 Troubleshooting 里有一小节专门讲但经常被跳过。并行计算之前先确认 CPU 特性是因为 AVX-512 和 AVX2 的代码路径在 Numba 编译结果里差异巨大跑错平台可能出现几倍的性能差距。2.2 最小函数理解 jit 和 njit 的差别官方 Quickstart 的第一个例子是双重循环求和。我第一次照着跑的时候最大的困惑是装饰器到底该写jit还是njit。这里先把结论说透import numpy as np from numba import njit njit(cacheTrue) def sum2d(a): s 0.0 for i in range(a.shape[0]): for j in range(a.shape[1]): s a[i, j] return s arr np.random.rand(1000, 1000) print(sum2d(arr)) # 第一次调用触发编译之后才是真实性能代码逻辑a.shape[0]和a.shape[1]是编译期可推断的维度信息s被推断为float64整个双重循环不会创建任何中间数组。参数说明里最值得写的是cacheTrue它把编译产物写到磁盘缓存下次运行同一进程或另一个进程时跳过编译。njit等价于jit(nopythonTrue)官方文档已经明确建议除非你确实需要 object 模式否则一律用njit。内置的四个常用装饰器可以这样区分装饰器使用场景并行支持典型返回值jit快速试错允许 object 模式回退配合 parallel任意njit生产环境首选强制 nopython配合 parallel任意vectorize将标量函数变成 ufunc自动 SIMD数组guvectorize将任意形状的块运算变成 ufunc线程/GPU数组其中vectorize和guvectorize的官方文档篇幅很大但并行计算场景里最常见的还是njit(parallelTrue)下面所有例子都基于它。2.3 第一次计时要看的三个数字跑通之后不要急着做并行。先采集三个数字纯 Python 时间、第一次调用的编译时间、第二次及以后的执行时间。直接用 timeit 会默认跳过编译阶段容易误判。我一般这样测import time import numpy as np from numba import njit njit(cacheTrue) def compute_histogram(data, bins): hist np.zeros(len(bins) - 1, dtypenp.int64) for i in range(data.shape[0]): for j in range(bins.shape[0] - 1): if bins[j] data[i] bins[j 1]: hist[j] 1 return hist data np.random.rand(100_000) bins np.linspace(0, 1, 11) t0 time.perf_counter() compute_histogram(data[:10], bins) # 触发编译 t1 time.perf_counter() compute_histogram(data, bins) # 真实执行 t2 time.perf_counter() print(f编译时间: {t1 - t0:.4f}s, 执行时间: {t2 - t1:.6f}s)这里的两个参数说明第一次只传data[:10]是为了让 Numba 用最小数据量完成函数编译避免把编译时间和执行时间混在一起dtypenp.int64是提前固定累加器的整数类型防止类型推断在int32与int64之间做额外判断。常见误用是直接拿第一次计时当作性能数据这样 Numba 看起来比纯 Python 慢十倍实际慢的是编译不是执行。正确做法是永远用第二次调用以后的耗时对比性能并且在代码注释里写明“编译需要先跑一次”。3. 理解 nopython 模式再决定要不要开 parallel3.1 类型推断链路决定了 nopython 模式能编译到什么程度Numba 的编译过程大致是Python 字节码 → Numba IR → 类型推断 → LLVM IR → 机器码。nopython 模式要求整个 IR 里不出现 Python 对象所有变量在编译期都要确定类型。这带来一个直接后果函数签名在编译期固定之后运行时传入不同类型参数不会自动生成新版本除非重新触发编译。这个模型解释了中文社区经常出现的几个“玄学报错”。比如from numba import njit njit def bad_union(x): if x 0: return 1.0 else: return 1 # 类型无法统一 bad_union(1)会报Cannot unify float64 and int64。解决办法是让所有返回路径的类型一致要么都返回浮点要么显式写np.float64(1)。官方文档在 Type unification 一节里把这条规则讲得很清楚中文翻译文档中常常把它简单译成“类型无法统一”真正要理解的是 Numba 不是运行时动态派发它是编译期静态推导。另一个容易踩的点是容器类型。Numba 对容器使用 typed.List、typed.Dict而不是 CPython 原生 list、dict。原因很简单原生的 Python 容器是堆上的对象图无法用机器码直接操作。官方文档有一节专门讲 Differences between CPython and Numba。实际影响是如果在njit函数里构造d {}较新版本的 Numba 会推断成 typed.Dict但如果往里面插入不同类型值就会报Cannot unify。常见做法是用numba.typed.Dict.empty()显式声明键值类型from numba import njit from numba.typed import Dict from numba.core import types njit def fill_dict(keys, values): d Dict.empty(key_typetypes.unicode_type, value_typetypes.float64) for i in range(keys.shape[0]): d[keys[i]] values[i] return d这个函数在字典和数组之间做了一次索引映射。显式声明 key_type 和 value_type 的目的是让类型推断在函数入口就完成而不是在插入第一个元素时才临时决定。签名形式是另一个常见的可优化点。提前写死签名可以避免运行时重复推断也更容易在 CI 里做接口回归签名写法意思适用场景float64(float64[:])一维数组 → 标量归约、求和、范数float64[:, :](float64[:, :])二维数组 → 二维数组矩阵变换、距离矩阵void(float64[:], float64[:])两个一维数组无返回值原地逻辑配合 out 参数3.2 开了 parallel 反而变慢的两个原因内存带宽和线程超额订阅官方文档 Parallel Performance 一节里有一个经常被忽略的说明当数组小到一定程度时并行版本比串行版本慢。原因有两个。第一内存带宽。prange把循环分给多个线程后如果每个迭代都在读一个大数组的不同部分线程会争抢内存控制器。DDR4 单通道带宽约 20 GB/s双通道约 40 GB/s一旦数据总量超过 CPU 缓存加速比就会受带宽约束。比如一个 500MB 的数组8 线程读一遍需要约 12ms而 2 线程读一遍约 25ms实际加速比只有 2远低于线程数比例。第二线程超额订阅。parallelTrue默认使用numba.config.NUMBA_NUM_THREADS初始值取自机器逻辑核心数。开发机上可能是 16 线程容器里如果只分配到 4 核却依然开 16 线程就会频繁切换。加上multiprocessing后问题更明显每个进程都开一套 Numba 线程池进程数 × 线程数会把 CPU 打爆。我一般这样处理import os from numba import config, njit, prange # 在 import numba 之后、任何 njit 函数编译之前设置 cpu_count os.cpu_count() config.NUMBA_NUM_THREADS max(1, min(cpu_count, 8)) njit(parallelTrue) def padded_sum(x): total 0.0 for i in prange(x.shape[0]): total x[i] return totalconfig.NUMBA_NUM_THREADS也可以在命令行用NUMBA_NUM_THREADS环境变量设置适合不同环境切换。如果程序已经编译过函数再改这个值Numba 会警告不会全部生效所以最好在模块加载阶段就设好。这里max(1, min(cpu_count, 8))的意思很直白CPU 再少也保证至少有 1 个线程最多不超 8 个防止在共享服务器上把邻居的核占满。3.3 parallelTrue 配合 prange 的最小正确姿势有了前面的铺垫就可以写第一个带并行的实际函数了。拿 3.2 的padded_sum来说total 0.0是函数体内的局部变量total x[i]在 prange 下会被识别为可归约模式。Numba 会为每个线程生成私有 total循环结束后再统一求和。这是 Numba 支持最稳的并行模式。如果像 K-Means 一样要算点到多个中心点的距离并行结构也一样from numba import njit, prange import numpy as np njit(parallelTrue, cacheTrue) def compute_sse(data, centers): n, k data.shape[0], centers.shape[0] result np.empty((n, k), dtypenp.float64) for i in prange(n): for j in range(k): diff data[i] - centers[j] result[i, j] diff diff return result data np.random.rand(200_000, 64) centers np.random.rand(32, 64) sse compute_sse(data, centers)代码逻辑外层prange(n)让 Numba 把行迭代分到多个线程每个线程独立写result的不同行内层用普通range(k)因为centers是只读访问保持顺序执行反而更利于缓存。参数要点prange必须配合parallelTrue否则它会退化成普通rangediff diff是向量点积Numba 能直接编译成 BLAS 调用但前提是启用了 nopython 模式。这里有一个性能边界result是 200000×32 的float64数组大约 51 MB。多线程并行写时会竞争内存带宽所以线程超过 8 个之后加速比可能不再是线性的。如果centers的数量上升到几千内层计算量变大这时才值得考虑把内层也并发化但嵌套 prange 容易踩到编译器的归约解析限制需要先做内层数据切块再合并。对这类需求官方文档的 guvectorize 页面是更好的参考。3.4 线程池、GIL 与 Python 多进程的配合方式一个常见误解是 Numba 并行计算会绕开 GIL。更准确的说法是Numba 编译后的机器码在执行时不持有 GIL因此多线程能真正并行的部分就是 njit 函数体本身。但函数之外的 Python 代码仍然受 GIL 制约。所以如果你在njit函数外面做数据处理那段 Python 代码不会因为 parallelTrue 而加速。当任务规模超过单机内存带宽时多进程通常是更好的选择。Numba 与 multiprocessing 配合有两点要注意。第一每个子进程在首次调用 njit 函数时都会重新编译共享内存里放的是数据不是编译产物cacheTrue 可以跳过编译时间但不会减少子进程占用的内存总量。第二不要在一个进程里同时开 Numba 的 64 线程和 16 个 multiprocessing 进程建议进程数不超过物理核心数且每个进程内线程数设为 1 或 2。from multiprocessing import Pool from numba import njit, prange, config config.NUMBA_NUM_THREADS 1 njit(parallelTrue) def worker_chunk(x): return x.sum() def run_parallel_over_processes(chunks, n_workers4): with Pool(n_workers) as pool: return pool.map(worker_chunk, chunks)这里把线程数限制为 1主要靠 4 个进程分摊计算。Pool.map会把数据序列化后传给子进程数组大的时候序列化本身也是成本所以更推荐用multiprocessing.shared_memory共享一个底层缓冲区子进程只接收起始索引和长度。官方文档没有专门讲这层配合属于实际项目里才会遇到的边界条件。4. 从 prange 到归约与向量化并行计算的四种落地写法4.1 先用一张表区分四种并行机制Numba 官方文档里 Parallel 相关的内容分散在几个页面如果没有指引很容易把 prange、vectorize、guvectorize、multiprocessing 混在一起。先看一张对比表方式适用场景并行单位数据共享方式最容易踩的坑prangeparallelTrue循环体重复计算的大数组遍历进程内线程局部变量自动私有标量归约自动合并线程数默认取逻辑核心数容器里需要手动设置vectorize逐元素计算类似 numpy ufunc底层 SIMD/线程输入数组只读不适合滑动窗口和有跨点依赖的算法guvectorize任意形状的逐块操作线程或 GPU显式传入输出数组签名维度符号需要仔细写手动multiprocessing有 IO 阻塞或第三方库调用多进程依赖进程间通信每个子进程重新 import numba内存翻倍vectorize和guvectorize适合纯逐元素计算比如对数、三角函数、滤波核。它们的优点是代码短把一个标量函数装饰一下之后就能像np.exp一样对整个数组调用。缺点是无法表达循环内部的复杂控制流。prange则保留函数内部的控制流更适合数据并行算法。4.2 线程数量不是越多越好容器场景要手动限制上一节提到config.NUMBA_NUM_THREADS默认取逻辑核心数。这在开发机上没问题但部署到 Kubernetes 或 Docker 容器时宿主机往往是 64 核而容器只分到 8 核。Numba 不知道容器的配额照样开 64 个线程结果就是线程间切换严重加速比不升反降。常见的检查方式numba -s | grep -i Number of threadss参数会输出本机线程数、CPU 型号、SIMD 指令支持情况。在容器里看到宿主机核心数而不是配额时立刻就该设置环境变量export NUMBA_NUM_THREADS8也可以在 Python 里设置。要注意的是config.NUMBA_NUM_THREADS必须在任何njit函数编译之前赋值否则可能不会生效。开发时想验证不同线程数的效果可以这样from numba import set_num_threads for n in [1, 2, 4, 8]: set_num_threads(n) # 在这里调用待测函数set_num_threads是运行时接口灵活但文档写得很低调。它和config.NUMBA_NUM_THREADS的区别在于前者可以动态调用后者只能在编译前设置。多线程并行还有一个容易被忽略的系统层面问题NUMA 体系下的内存亲和性。当数据分布在高位 NUMA 节点时线程被调度到另一个节点上访问内存要走跨节点互联。Numba 本身不会处理 NUMA实际表现是set_num_threads(8)之后性能抖动明显。常见做法是用taskset固定 CPU 范围或者用numactl --interleaveall启动 Python 进程减少跨节点访问。4.3 归约写法决定能不能并行标量归约与线程私有数组标量归约是 Numba 支持最好的并行模式njit(parallelTrue) def reduce_sum(x): total 0.0 for i in prange(x.shape[0]): total x[i] return total但连续数组的原地累加则未必。尤其是索引值依赖数据本身时比如直方图统计多个线程可能同时写同一个桶。官方文档没有直接给出答案我建议用线程私有缓冲import numpy as np import numba from numba import njit, prange njit(parallelTrue, cacheTrue) def group_by_index(values, index, n_groups): n_threads numba.get_num_threads() sums np.zeros((n_threads, n_groups), dtypenp.float64) cnt np.zeros((n_threads, n_groups), dtypenp.int64) for i in prange(values.shape[0]): t numba.get_thread_id() g index[i] sums[t, g] values[i] cnt[t, g] 1 total_sum np.zeros(n_groups, dtypenp.float64) total_cnt np.zeros(n_groups, dtypenp.int64) for g in range(n_groups): for t in range(n_threads): total_sum[g] sums[t, g] total_cnt[g] cnt[t, g] return total_sum / np.maximum(total_cnt, 1) values np.random.rand(100_000) index np.random.randint(0, 10, 100_000) print(group_by_index(values, index, 10))逻辑说明sums[t, g]的 t 是当前线程编号g 是分组索引。每个线程只写自己编号对应的那一行不会出现两个线程同时写同一个位置的情况。合并阶段用普通 range 串行完成数据量是 n_threads × n_groups规模很小。numba.get_thread_id()只能在 prange 循环内使用返回的是当前线程在线程池里的编号。注意不要写成sums[g] values[i]这种形式。表面上看是对的但在并行循环里不同线程可能同时命中同一个 g产生数据竞争。结果在数据量大时可能差得很小很难察觉但偶尔会突然出现一个错误值。4.4 向量化与 prange 的取舍按数据规模决定当计算可以表示成逐元素函数时vectorize往往更简单from numba import vectorize, float64 import numpy as np vectorize([float64(float64, float64)]) def add_weighted(a, b): return a * 0.7 b * 0.3 x np.random.rand(1_000_000) y np.random.rand(1_000_000) z add_weighted(x, y)这段代码会被 Numba 转化成 ufunc底层的调度会根据数组大小动态决定用 SIMD 还是多线程。但它的限制也很明确函数体不允许有依赖数组位置的逻辑。比如卷积、相邻点差分这种算法必须写成 prange 循环。取舍原则很简单单点映射用 vectorize区域映射用 prange。5. 翻译文档查不到的坑缓存失效、禁用 JIT 与加速比验证5.1 cacheTrue 的失效规则改代码后总觉得没生效怎么办cacheTrue会把编译产物写到__pycache__下的.nbc文件。缓存 key 由函数源码、依赖版本、CPU 特性共同决定正常情况下改函数体应该触发重新编译。但如果 .pyc 文件没有刷新或者函数被包在另一个模块里 import编译结果可能不会按预期更新。我一般用三个手段定位find . -name *.nbc -delete删除所有 Numba 缓存强制下一轮重新编译。第二个手段是临时关缓存from numba import njit njit(cacheFalse) def debug_func(x): return x 1第三个手段是用环境变量NUMBA_DISABLE_JIT1完整关闭 JIT这时候所有装饰器退化成普通 Python 函数。这个变量也是排查并行结果异常的方式打开 parallelTrue 后结果和纯 Python 不一致时先跑一遍禁用 JIT 的版本。如果禁用后结果正确再看代码里的数据竞争如果禁用后结果就不正确说明算法本身写错了与 Numba 无关。确认这一步后再把数据量缩小到几千条直接打印中间数组肉眼检查通常能很快定位到是哪一行写坏了数据。开发早期还建议把 Numba 警告全开export NUMBA_WARNINGS1这样能看到函数里哪些操作触发了不必要的类型转换生产环境再关掉。5.2 用加速比曲线判断要不要继续加线程性能验证最终要落到数字上。我常用的做法是固定函数与数据只改变线程数观察执行时间from time import perf_counter import numpy as np from numba import njit, prange, set_num_threads njit(parallelTrue, cacheTrue) def dot_all(data): total 0.0 for i in prange(data.shape[0]): for j in range(data.shape[1]): total data[i, j] * data[i, j] return total data np.random.rand(500_000, 32) dot_all(data[:1000]) # 先编译 for n_threads in [1, 2, 4, 8, 16]: set_num_threads(n_threads) t0 perf_counter() dot_all(data) t1 perf_counter() print(f线程数 {n_threads}: {(t1 - t0) * 1000:.2f} ms)set_num_threads能在运行时调整线程数但要在第一次编译之后调用否则部分代码路径可能用了默认线程数。五个数字摆在一起就能判断当前数据量下是否有明显的可扩展性。如果 4 线程到 8 线程的耗时几乎没有变化说明已经碰到内存带宽上限再用prange优化意义不大。如果从 1 线程到 2 线程出现了超线性加速反而要怀疑是不是缓存命中率变化引起的这也算常见。在 CI 里跑这种线性扩展测试时要注意numba -s输出的 CPU 特性一致否则同一套代码在不同机器上会因为 SIMD 指令集差异出现明显速度差。把这行命令的输出保存下来和本机对比一次能省掉很多因指令集差异引起的性能排查。本文还有配套的精品资源点击获取
返回列表