使用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
相关产品推荐
相关产品推荐

