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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 17:40:18