使用Numba @njit加速类外N体模拟函数时出现TypingError的解决方法
问题:Numba @njit(parallel=True) 加速N体模拟时类外函数报错
将N体模拟中原本的类内加速度计算函数移到类外,使用Numba的@njit(nopython=True, parallel=True)装饰器加速,通过类内方法调用该外部函数时出现TypingError。
原类内加速度计算函数
def _calculate_acceleration(self, mass, pos, rsoft): """ Calculate the acceleration. """ # TODO: N = self.particles rsoft = self.rsoft posx = pos[:,0] posy = pos[:,1] posz = pos[:,2] G = self.G npts = self.nparticles acc = np.zeros((npts, 3)) for i in prange(npts): for j in prange(npts): if (j>i): x = (posx[i]-posx[j]) y = (posy[i]-posy[j]) z = (posz[i]-posz[j]) rsq = x**2 + y**2 + z**2 req = np.sqrt(x**2 + y**2) f = -G*mass[i,0]*mass[j,0]/rsq theta = np.arctan2(y, x) phi = np.arctan2(z, req) fx = f*np.cos(theta)*np.cos(phi) fy = f*np.sin(theta)*np.cos(phi) fz = f*np.sin(phi) acc[i,0] += fx/mass[i] acc[i,1] += fy/mass[i] acc[i,2] += fz/mass[i] acc[j,0] -= fx/mass[j] acc[j,1] -= fy/mass[j] acc[j,2] -= fz/mass[j] return acc
粒子初始化及模拟执行代码
def initialRandomParticles(N = 100, total_mass = 10): """ Initial particles """ particles = Particles(N) masses = particles.masses mass = total_mass/particles.nparticles particles.masses = (masses*mass) positions = np.random.randn(N,3) velocities = np.random.randn(N,3) accelerations = np.random.randn(N,3) particles.positions = positions particles.velocities = velocities particles.accelerations = accelerations return particles particles = initialRandomParticles(N = 10**5, total_mass = 20) sim = NbodySimulation(particles) sim.setup(G=G,method="RK4",io_freq=200,io_title=problem_name,io_screen=True,visualized=False, rsoft=0.01) sim.evolve(dt=0.01,tmax=10) # Particles and NbodySimulation are defined class.
移到类外后的新函数代码
@njit(nopython=True, parallel=True) def _calculate_acceleration(n, npts, G, mass, pos, rsoft): """ Calculate the acceleration. This function is out of the class. """ # TODO: posx = pos[:,0] posy = pos[:,1] posz = pos[:,2] acc = np.zeros((n, 3)) sqrt = np.sqrt for i in prange(npts): for j in prange(npts): if (j>i): x = (posx[i]-posx[j]) y = (posy[i]-posy[j]) z = (posz[i]-posz[j]) rsq = x**2 + y**2 + z**2 req = sqrt(x**2 + y**2 + z**2) f = -G*mass[i,0]*mass[j,0]/(req + rsoft)**2 fx = f*x**2/rsq fy = f*y**2/rsq fz = f*z**2/rsq acc[i,0] = fx/mass[i] + acc[i,0] acc[i,1] = fy/mass[i] + acc[i,1] acc[i,2] = fz/mass[i] + acc[i,2] acc[j,0] = fx/mass[j] - acc[j,0] acc[j,1] = fy/mass[j] - acc[j,1] acc[j,2] = fz/mass[j] - acc[j,2] return acc
错误信息
TypingError Traceback (most recent call last) :428, in NbodySimulation._update_particles_rk4(self, dt, particles) 426 position = particles.positions # y0[0] 427 velocity = particles.velocities # y0[1], k1[0] --> 428 acceleration = self._calculate_acceleration_inclass() # k1[1] 430 position2 = position + 0.5*velocity * dt # y1[0] 431 velocity2 = velocity + 0.5*acceleration * dt # y1[1], k2[0] :381, in NbodySimulation._calculate_acceleration_inclass(self) 377 def _calculate_acceleration_inclass(self): 378 """ 379 Calculate the acceleration. ... <source elided> acc[i,0] = fx/mass[i] + acc[i,0] ^
解决方法
1. 修正变量维度不匹配问题
报错核心是Numba无法推断变量类型:原代码中mass是二维数组(npts×1),但类外函数中用mass[i]获取标量,实际会得到一维数组,导致类型冲突。
- 把类外函数中的
mass[i]改为mass[i,0],mass[j]改为mass[j,0],和原代码保持一致。
2. 移除冗余参数并统一粒子数量变量
类外函数同时传入n和npts,两者都是粒子数量,且acc初始化用n、循环用npts,易引发维度不匹配。
- 删除
n参数,统一用npts初始化acc:acc = np.zeros((npts, 3))。
3. 修复力的计算逻辑错误
类外函数修改了原有的力分量计算逻辑,导致物理错误同时触发类型推断失败:
- 恢复正确的力分量计算:万有引力分量应为标量力乘以方向向量(
x/req、y/req、z/req),而非x²/rsq。 - 修正软处理方式:正确的软处理是给距离平方加
rsoft²,避免除以零,而非给距离加rsoft后平方。
4. 修正并行循环的线程竞争问题
嵌套prange会导致多个线程同时修改acc数组的同一位置,引发竞争和计算错误。
- 仅对**外层循环
i**使用prange,内层循环用普通range。
修正后的类外函数代码
from numba import njit, prange import numpy as np @njit(nopython=True, parallel=True) def _calculate_acceleration(npts, G, mass, pos, rsoft): """ Calculate the acceleration. This function is out of the class. """ posx = pos[:,0] posy = pos[:,1] posz = pos[:,2] acc = np.zeros((npts, 3)) sqrt = np.sqrt rsoft_sq = rsoft ** 2 for i in prange(npts): for j in range(npts): if j > i: x = posx[i] - posx[j] y = posy[i] - posy[j] z = posz[i] - posz[j] rsq = x**2 + y**2 + z**2 req = sqrt(rsq + rsoft_sq) # 万有引力标量值,软处理避免除以零 f_mag = -G * mass[i,0] * mass[j,0] / (req ** 2) # 计算力的分量 fx = f_mag * x / req fy = f_mag * y / req fz = f_mag * z / req # 更新加速度,遵循牛顿第三定律 acc[i,0] += fx / mass[i,0] acc[i,1] += fy / mass[i,0] acc[i,2] += fz / mass[i,0] acc[j,0] -= fx / mass[j,0] acc[j,1] -= fy / mass[j,0] acc[j,2] -= fz / mass[j,0] return acc
类内调用方法修正
确保传入参数与类外函数匹配:
def _calculate_acceleration_inclass(self): return _calculate_acceleration( self.nparticles, self.G, self.particles.masses, self.particles.positions, self.rsoft )
内容的提问来源于stack exchange,提问作者hozi602
相关产品推荐
相关产品推荐

