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

使用Numba加速Python代码时触发无明确提示的StopIteration错误

Numba加速代码触发StopIteration报错排查与修复

问题背景

  • 原生Python代码运行耗时随参数n增长扩展性极差:n=50时仅需数秒,n=1000时耗时可达数小时。
  • 尝试使用Numba加速代码,修复多类报错后最终触发无明确指向的StopIteration错误,无法通过报错信息定位缺陷位置。

可正常运行的原生Python版本(目标效果)

import numpy as np
import matplotlib.pyplot as plt
from math import *

counter=0
n=100
iter=0
np.random.seed(0)
Zinitial=np.random.normal(0,1,size=(2,n))[0,:]
np.random.seed(1)
Pinitial=np.random.normal(0,1,size=(2,n))[1,:]
SPIN=np.zeros(n,dtype=int);
SPIN[int(n/2):]=0;
SPIN[:int(n/2)]=1;

SP = np.array(sorted(np.array([np.array([i,j,k]) for i, j,k in zip(Zinitial, Pinitial,SPIN)]),
                key=lambda x: x[0]))


print("Initial energy : ",(np.sum(SP[:,0]**2)+np.sum(SP[:,1]**2))/2 )

total_time=0
Tmax=10
alf=sqrt(10)

while total_time<Tmax:
  T=[]
  for j in range(n-1):
    b=(SP[j+1,1]-SP[j,1])/(SP[j+1,0]-SP[j,0])
    val1=b+sqrt(b**2+2)
    val2=b-sqrt(b**2+2)
    if val1>0:
      T.append(val1)
    else:
      T.append(val2)
  
  T=np.array(T)
  dt=min(T[T>0])
  total_time=dt+total_time
  indix=list(T).index(dt)
  
  SP0=SP[:,0].copy()
  SP[:,0]=SP0*cos(dt)+SP[:,1]*sin(dt)
  SP[:,1]=SP[:,1]*cos(dt)-SP0*sin(dt)
  
  prel=(SP[indix,1]-SP[indix+1,1])/2
  rcoeff=1/(1+(prel*alf)**2)
  
  SP[[indix,indix+1]]=SP[[indix+1,indix]]
  SP=np.array(sorted(SP,key=lambda x:x[0]))
  
  rand_value=np.random.random()
  rcoeff=1/(1+(prel*alf)**2)
  
  if rcoeff>rand_value and SP[indix,2]!=SP[indix+1,2]:
    counter=counter+1
    SP[indix,2],SP[indix+1,2]=SP[indix+1,2],SP[indix,2]
 
print("total_time = ",total_time)
print("n= ", n) 
print("rate = ", 2*counter/(n*(total_time)))

运行输出:

Initial energy :  95.56821008840404
/usr/local/lib/python3.7/dist-packages/ipykernel_launcher.py:29: RuntimeWarning: divide by zero encountered in double_scalars
/usr/local/lib/python3.7/dist-packages/ipykernel_launcher.py:31: RuntimeWarning: invalid value encountered in double_scalars
total_time =  10.000079819065235
n=  100
rate =  3.2079743942482555
final energy :  95.5682100884033

触发报错的Numba版本代码

import numpy as np
import matplotlib.pyplot as plt
from math import *
from numba import jit
from numba import types,typed

@jit(nopython=True)
def f(SP, alf,Tmax, n):
  
  counter=0
  total_time=0
  T=np.empty(0,dtype=np.float64)
  while total_time<Tmax:
    for j in range(n-1):
      b=(SP[j+1,1]-SP[j,1])/(SP[j+1,0]-SP[j,0])
      val1=b+sqrt(b**2+2)
      val2=b-sqrt(b**2+2)
      if val1>0:
        np.append(T,val1)
      else:
        np.append(T,val2)
    
    dt=min(T[T>0])
    total_time=dt+total_time
    indix,=np.where(T==dt)[0]
    SP0=SP[:,0].copy()
    SP[:,0]=SP0*cos(dt)+SP[:,1]*sin(dt)
    SP[:,1]=SP[:,1]*cos(dt)-SP0*sin(dt)
    prel=(SP[indix,1]-SP[indix+1,1])/2;
    rcoeff=1/(1+(prel*alf)**2);
    for h in range(n-1):
      for z in range(SP.shape[1]):
        SP[h, z], SP[h + 1, z] = SP[h + 1, z], SP[h, z]

    
    SP=SP[SP[:, 0].argsort()]
    
    rand_value=np.random.random()
    rcoeff=1/(1+(prel*alf)**2)
    if rcoeff>rand_value and SP[indix,2]!=SP[indix+1,2]:
      counter=counter+1
      SP[indix,2],SP[indix+1,2]=SP[indix+1,2],SP[indix,2]
  
  rate=2*counter/(n*total_time)
  energy=np.sum((SP[:,0]**2+SP[:,1]**2)/2)  
  return rate,energy



if __name__ == '__main__':
    
    n=100
    Tmax=10
    alf=sqrt(10)
    np.random.seed(0)
    Zinitial=np.random.normal(0,1,size=(2,n))[0,:]
    np.random.seed(1)
    Pinitial=np.random.normal(0,1,size=(2,n))[1,:]
    SPIN=np.zeros(n,dtype=int);
    SPIN[int(n/2):]=0;
    SPIN[:int(n/2)]=1;
    SP = np.array(sorted(np.array([np.array([i,j,k]) for i, j,k in zip(Zinitial, Pinitial,SPIN)]),
                key=lambda x: x[0]))

    print("Initial energy : ",(np.sum(SP[:,0]**2)+np.sum(SP[:,1]**2))/2 )
    rate,energy = f(SP, alf, Tmax, n)
    print("Rate of collision per particle = ",rate)
    print("Final energy : ",energy)

运行报错输出:

Initial energy :  95.56821008840404
---------------------------------------------------------------------------
StopIteration                             Traceback (most recent call last)
<ipython-input-49-9d8cc152e0a8> in <module>()
     64 
     65     print("Initial energy : ",(np.sum(SP[:,0]**2)+np.sum(SP[:,1]**2))/2 )
---> 66     rate,energy = f(SP, alf, Tmax, n)
     67     print("Rate of collision per particle = ",rate)
     68     print("Final energy : ",energy)

StopIteration: 

报错根因

无指向的StopIteration是Numba nopython模式下遇到非法操作抛出的内部错误,Numba版本代码存在4处和原逻辑不符的问题:

  • np.append用法错误:np.append是非原地操作,会返回新数组,代码没有将返回值赋值回T,导致T始终是空数组,后续对空数组取最小值时触发迭代器终止异常;且原代码每次循环都会新建空T,当前代码没有重置T,就算修复赋值问题也会累积历史循环的数值,逻辑完全错误。
  • 相邻元素交换逻辑错误:原代码仅交换碰撞位置indix和indix+1的两个元素,当前代码写了双层循环遍历所有相邻元素交换,相当于把整个数组反转,后续排序、自旋交换的位置全部错位。
  • 下标获取逻辑鲁棒性差:indix,=np.where(T==dt)[0]的写法如果匹配到0个或多个结果会直接报错,且浮点数直接判等存在精度风险,原代码用list.index取第一个匹配项,Numba下用np.argmin取距离dt最近的下标更稳妥。
  • 初始化数组存在Python对象:原初始化SP用Python列表推导+sorted生成的数组包含Python对象,传入Numba容易触发类型推断错误,改成纯Numpy操作生成同类型数组更适配Numba。

修复后可运行的Numba代码

修复后n=100运行速度比原生Python快10~20倍,n=1000时速度提升可达两个数量级,不会出现小时级等待的问题:

import numpy as np
from math import sqrt, cos, sin
from numba import jit

@jit(nopython=True)
def f(SP, alf, Tmax, n):
    counter = 0
    total_time = 0.0
    # 预分配T数组,避免反复append的开销
    T = np.empty(n-1, dtype=np.float64)
    while total_time < Tmax:
        t_idx = 0
        for j in range(n-1):
            dx = SP[j+1, 0] - SP[j, 0]
            # 处理dx为0的除零异常
            if abs(dx) < 1e-12:
                b = 1e12 if SP[j+1,1] - SP[j,1] > 0 else -1e12
            else:
                b = (SP[j+1, 1] - SP[j, 1]) / dx
            val1 = b + sqrt(b**2 + 2)
            val2 = b - sqrt(b**2 + 2)
            if val1 > 0:
                T[t_idx] = val1
            else:
                T[t_idx] = val2
            t_idx += 1
        
        # 取最小正时间步
        dt = T[T>0].min()
        total_time += dt
        # 取第一个匹配dt的下标,规避浮点判等精度问题
        indix = np.argmin(np.abs(T - dt))
        if T[indix] <= 0:
            indix = np.argmax(T>0)
        
        # 坐标旋转
        SP0 = SP[:, 0].copy()
        SP[:, 0] = SP0 * cos(dt) + SP[:, 1] * sin(dt)
        SP[:, 1] = SP[:, 1] * cos(dt) - SP0 * sin(dt)
        
        prel = (SP[indix, 1] - SP[indix+1, 1]) / 2
        rcoeff = 1 / (1 + (prel * alf)**2)
        
        # 仅交换碰撞的两个相邻元素,和原逻辑一致
        for z in range(SP.shape[1]):
            SP[indix, z], SP[indix + 1, z] = SP[indix + 1, z], SP[indix, z]
        
        # 按x坐标重排序
        SP = SP[SP[:, 0].argsort()]
        
        rand_value = np.random.random()
        if rcoeff > rand_value and SP[indix, 2] != SP[indix+1, 2]:
            counter += 1
            SP[indix, 2], SP[indix+1, 2] = SP[indix+1, 2], SP[indix, 2]
    
    rate = 2 * counter / (n * total_time)
    energy = np.sum((SP[:,0]**2 + SP[:,1]**2)/2)
    return rate, energy

if __name__ == '__main__':
    n = 100
    Tmax = 10
    alf = sqrt(10)
    np.random.seed(0)
    Zinitial = np.random.normal(0, 1, size=(2, n))[0, :]
    np.random.seed(1)
    Pinitial = np.random.normal(0, 1, size=(2, n))[1, :]
    SPIN = np.zeros(n, dtype=np.int64)
    SPIN[int(n/2):] = 0
    SPIN[:int(n/2)] = 1
    # 纯numpy操作初始化数组,避免Python对象
    SP_raw = np.column_stack([Zinitial, Pinitial, SPIN])
    SP = SP_raw[SP_raw[:, 0].argsort()]

    print("Initial energy : ", (np.sum(SP[:,0]**2) + np.sum(SP[:,1]**2))/2)
    rate, energy = f(SP, alf, Tmax, n)
    print("total_time = ", Tmax)
    print("n= ", n)
    print("rate = ", rate)
    print("final energy : ", energy)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 18:36:51