You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

指定随机种子下Numba并行化与pytest测试异常排查

解决Numba并行化下随机种子不固定的测试问题

问题背景

在Jupyter Lab中用pytest测试Numba并行函数时,即使设置了np.random.seed(0),开启并行化后函数输出无法复现,导致测试失败;关闭并行化则测试正常通过。

复现代码

np.random.seed(0)

@njit(parallel=True)
def sum_shape(repeats,samples):
    
    sum = 0
    
    for i in prange(0,repeats):
        x = np.random.uniform(0,100,int(samples))
        sum += np.sum(x)
        
    return sum

%%ipytest
def test():
    assert sum_shape(10000,10) == 4995000.330114928 # expected output given seed

错误表现

开启并行化时测试失败:

def test():
       assert sum_shape(10000,10) == 4995000.330114928
E       assert 4996808.331482845 == 4995000.330114928
E        +  where 4996808.331482845 = sum_shape(10000, 10)

关闭并行化时测试通过:

.                                                                                            [100%]
1 passed in 0.05s

问题原因

  1. 并行线程随机状态独立:全局np.random.seed(0)仅初始化主线程的随机状态,并行子线程会使用各自默认的随机状态,不受全局种子控制。
  2. prange执行顺序不确定:并行循环的迭代顺序不固定,即使每个线程种子固定,不同执行顺序也会导致最终求和结果不同。
  3. Numba对numpy随机函数的处理:并行模式下,Numba不会同步主线程的numpy随机状态到子线程,导致每个线程生成的随机序列不可控。

解决方案

使用Numba官方提供的线程安全随机数生成API,替代numpy的随机函数,确保每个线程的随机序列可复现。

修改后的代码

from numba import njit, prange
from numba.random import create_xoroshiro128p_states, xoroshiro128p_uniform_float64
import numpy as np

@njit(parallel=True)
def sum_shape(repeats, samples, seed=0):
    # 为每个并行线程创建独立且可复现的随机状态
    rng_states = create_xoroshiro128p_states(prange.num_threads(), seed=seed)
    total_sum = 0.0
    
    for i in prange(repeats):
        # 绑定当前迭代到对应线程的随机状态
        thread_idx = i % prange.num_threads()
        state = rng_states[thread_idx]
        # 生成0-1的均匀分布,再缩放至0-100
        x = xoroshiro128p_uniform_float64(state, samples) * 100
        total_sum += np.sum(x)
    
    return total_sum

%%ipytest
def test():
    # 先运行一次并行版本获取固定的预期值
    expected = sum_shape(10000, 10)
    # 使用np.isclose处理浮点数精度问题
    assert np.isclose(sum_shape(10000, 10), expected)

关键说明

  • create_xoroshiro128p_states:基于指定的全局种子,为每个并行线程生成独立的随机状态,确保每次运行时每个线程的随机序列一致。
  • xoroshiro128p_uniform_float64:Numba提供的线程安全随机数生成函数,替代np.random.uniform。
  • np.isclose:避免浮点数精确比对的误差,即使结果有微小精度差异也能通过测试。

替代方案(不推荐)

如果坚持使用numpy随机函数,需要在每个线程内手动设置种子,但会损失并行效率,且无法解决prange执行顺序问题:

@njit(parallel=True)
def sum_shape(repeats, samples, seed=0):
    total_sum = 0.0
    for i in prange(repeats):
        # 每个迭代设置独立种子
        np.random.seed(seed + i)
        x = np.random.uniform(0, 100, int(samples))
        total_sum += np.sum(x)
    return total_sum

内容的提问来源于stack exchange,提问作者Luke S

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 05:40:32