Numba仅使N体问题代码提速4倍,如何进一步优化?
N体问题Numba优化疑问与代码优化建议
我正在研究N体问题(给定空间中N个天体的位置,计算它们之间的相互作用)。当N=10000个粒子时,未使用即时编译(non-jitted)的函数耗时约84秒,经Numba即时编译(jitted)的函数耗时约22秒。但根据相关文章和视频,Numba宣传可使代码提速1-2个数量级!因此我分享代码,询问是否还有优化空间,同时怀疑要达到1-2个数量级的提速只能使用多线程或GPU,这种想法是否正确?
测试代码如下:
import numpy as np from numba import jit, njit import time import timeit def compute_acc( pos, mass, G, softening ): """ Computes the acceleration of N bodies Args: pos (type=np.array, size= Nx3): x, y, z positions of the N particles mass (type=np.array, size= Nx1): mass of the particles G (float): Newton's Gravitational constant softening (float): softening parameter Returns: acc (type=np.array, size= Nx3): ax, ay, az accelerations of the N particles """ # positions r = [x,y,z] for all particles x = pos[:,0:1] y = pos[:,1:2] z = pos[:,2:3] # matrix that stores all pairwise particle separations: r_j - r_i dx = x.T - x dy = y.T - y dz = z.T - z # matrix that stores 1/r^3 for all particle pairwise particle separations inv_r3 = (dx**2 + dy**2 + dz**2 + softening**2)**(-1.5) ax = G * (dx * inv_r3) @ mass ay = G * (dy * inv_r3) @ mass az = G * (dz * inv_r3) @ mass # pack together the acceleration components acc = np.hstack((ax,ay,az)) return acc #Define the jitted version of compute_acc compute_acc_jit= njit(cache=True,fastmath=True) (compute_acc) #Initialize the parameters to test the functions np.random.seed(123) N=10000 pos=np.random.uniform(low=-10, high=10, size=(N,3)) # Random uniform positions mass=np.random.uniform(low=1, high=20, size=(N,1)) # Random uniform masses G=1.0 softening=0.1 # Compute Non-Jitted time: T1= min(timeit.repeat(stmt='compute_acc(pos, mass, G, softening)',\ timer=time.perf_counter,repeat=3, number=1,globals=globals()) ) print("Non-JIT time=",T1," ") # Compute Jitted time: T2= min(timeit.repeat(stmt='compute_acc_jit(pos, mass, G, softening)',\ timer=time.perf_counter,repeat=3, number=1,globals=globals()) ) print("JIT time=",T2," ")
代码优化方向
1. 削减不必要的数组维度开销
原代码中x = pos[:,0:1]会生成(N,1)维度的数组,后续转置、运算会产生额外的内存和计算消耗。直接使用一维数组可以简化逻辑并提升效率:
x = pos[:, 0] y = pos[:, 1] z = pos[:, 2]
此时dx = x[None, :] - x[:, None]与原代码的x.T - x逻辑完全一致,但内存占用更低。
2. 避免存储全量成对距离数组,降低内存瓶颈
原代码创建了多个N×N规模的大数组(N=10000时,单个float64类型的N×N数组约763MB,三个坐标数组加inv_r3总占用超3GB),远超CPU缓存容量,导致频繁的内存读写拖慢速度。改用循环计算加速度,可彻底避免存储全量成对数据:
@njit(cache=True, fastmath=True, parallel=True) def compute_acc_opt(pos, mass, G, softening): N = pos.shape[0] acc = np.zeros_like(pos) # 启用并行循环,利用CPU多核 for i in range(N): dx = pos[:, 0] - pos[i, 0] dy = pos[:, 1] - pos[i, 1] dz = pos[:, 2] - pos[i, 2] r_sq = dx**2 + dy**2 + dz**2 + softening**2 inv_r3 = r_sq ** (-1.5) # 直接计算当前粒子的加速度,无需存储中间大数组 acc[i, 0] = G * np.sum(dx * inv_r3 * mass) acc[i, 1] = G * np.sum(dy * inv_r3 * mass) acc[i, 2] = G * np.sum(dz * inv_r3 * mass) return acc
同时将mass转为一维数组mass = mass.reshape(-1),减少矩阵运算时的维度转换开销。
3. 开启Numba多线程并行
在njit装饰器中添加parallel=True,让Numba自动对循环进行多线程拆分,充分利用CPU多核资源,进一步提升运算速度。
关于提速量级的疑问
你的部分想法是对的,但并非只能依赖多线程或GPU:
- 原代码的核心瓶颈是内存带宽:N=10000时,N×N的大数组远超CPU缓存,导致大量数据在内存与缓存间来回搬运,即使Numba优化了运算逻辑,也会被内存瓶颈限制速度。优化内存使用后,单线程下就能获得明显提速(比如从22秒降至几秒)。
- 若要达到1-2个数量级的提速(比如从84秒降至0.8-8秒),单线程确实难以实现。此时开启Numba的多线程并行模式,可利用多核CPU进一步提速;而GPU凭借极高的并行计算能力,能处理更大规模的N体问题,提速效果会更显著。
内容的提问来源于stack exchange,提问作者Rafid Bendimerad
相关产品推荐
相关产品推荐

