Python实现N体问题求解时N=1结果异常、多粒子报错求助
N体问题Python复现问题排查
问题概述
尝试用Python实现N体问题模拟,自定义NBody类存储粒子参数,每个粒子按「x坐标、y坐标、x方向速度、y方向速度」的顺序拼接存储在同一列表中。初始版本存在两个问题:
- 单粒子测试(质量为1、初始x方向速度为1)运动结果不符合预期
- 2粒子无初速度模拟抛出
TypeError错误
初始版本代码
import numpy as np import scipy.integrate as integrate class NBody(): G = 1 t = [] rv= [] m = [] def __init__(self, G, t): self.G = G self.t = t def addMass(self, mass, initPos, initVel): self.m = mass for i in initPos: self.rv.append(i) for i in initVel: self.rv.append(i) def FNBody(self, t, rv): rx = rv[0::4].copy() ry = rv[1::4].copy() vx = rv[2::4].copy() vy = rv[3::4].copy() n = len(rx) F = np.zeros(n*4, dtype=np.float) for i in range(n): addx = 0 addy = 0 for j in range(n): if i != j: addx+=self.G*self.m[i]*self.m[j]*(rx[j]-rx[i])/((rx[i]-rx[j])**2+(ry[i]-ry[j])**2)**1.5 addy+=self.G*self.m[i]*self.m[j]*(ry[j]-ry[i])/((rx[i]-rx[j])**2+(ry[i]-ry[j])**2)**1.5 F[4*i]=vx[i] F[4*i+1]=vy[i] F[4*i+2]=addx F[4*i+3]=addy return F def solveODE(self): return integrate.solve_ivp(self.FNBody, (self.t[0],self.t[len(self.t)-1]),self.rv,t_eval=self.t)
初始版本报错信息
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) <ipython-input-56-9680340d6ca4> in <module> 3 myN2Body.addMass(1,[0,0],[0,0]) 4 myN2Body.addMass(1,[1,1],[0,0]) ----> 5 sol = myN2Body.solveODE() 6 7 plt.figure(1) <ipython-input-54-0e1c24c59824> in solveODE(self) 63 64 ---> 65 return integrate.solve_ivp(self.FNBody, (self.t[0],self.t[len(self.t)-1]),self.rv,t_eval=self.t) 66 /software/anaconda3/lib/python3.7/site-packages/scipy/integrate/_ivp/ivp.py in solve_ivp(fun, t_span, y0, method, t_eval, dense_output, events, vectorized, args, **options) 541 method = METHODS[method] 542 ---> 543 solver = method(fun, t0, y0, tf, vectorized=vectorized, **options) 544 545 if t_eval is None: /software/anaconda3/lib/python3.7/site-packages/scipy/integrate/_ivp/rk.py in __init__(self, fun, t0, y0, t_bound, max_step, rtol, atol, vectorized, first_step, **extraneous) 93 self.max_step = validate_max_step(max_step) 94 self.rtol, self.atol = validate_tol(rtol, atol, self.n) --- 95 self.f = self.fun(self.t, self.y) 96 if first_step is None: 97 self.h_abs = select_initial_step( /software/anaconda3/lib/python3.7/site-packages/scipy/integrate/_ivp/base.py in fun(t, y) 137 def fun(t, y): 138 self.nfev += 1 -- 139 return self.fun_single(t, y) 140 141 self.fun = fun /software/anaconda3/lib/python3.7/site-packages/scipy/integrate/_ivp/base.py in fun_wrapped(t, y) 19 20 def fun_wrapped(t, y): --- 21 return np.asarray(fun(t, y), dtype=dtype) 22 23 return fun_wrapped, y0 <ipython-input-54-0e1c24c59824> in FNBody(self, t, rv) 49 for j in range(n): 50 if i != j: --- 51 addx+=self.G*self.m[i]*self.m[j]*(rx[j]-rx[i])/((rx[i]-rx[j])**2+(ry[i]-ry[j])**2)**1.5 52 addy+=self.G*self.m[i]*self.m[j]*(ry[j]-ry[i])/((rx[i]-rx[j])**2+(ry[i]-ry[j])**2)**1.5 53 F[4*i]=vx[i] TypeError: 'int' object is not subscriptable
第一次修正后代码
修正了质量存储的问题,新增了clearmass方法清空历史数据:
import numpy as np import scipy.integrate as integrate class NBody(): G = 1 t = [] rv= [] m = [] def __init__(self, G, t): self.G = G self.t = t def clearmass(self): self.m = [] self.rv = [] def addMass(self, mass, initPos, initVel): self.m.append(mass) for i in initPos: self.rv.append(i) for i in initVel: self.rv.append(i) def FNBody(self, t, rv): rx = rv[0::4].copy() ry = rv[1::4].copy() vx = rv[2::4].copy() vy = rv[3::4].copy() n = len(rx) F = np.zeros(n*4, dtype=np.float) for i in range(n): addx = 0 addy = 0 for j in range(n): if i != j: addx+=self.G*self.m[j]*(rx[j]-rx[i])/((rx[i]-rx[j])**2+(ry[i]-ry[j])**2)**1.5 addy+=self.G*self.m[j]*(ry[j]-ry[i])/((rx[i]-rx[j])**2+(ry[i]-ry[j])**2)**1.5 F[4*i]=vx[i] F[4*i+1]=vy[i] F[4*i+2]=addx F[4*i+3]=addy return F def solveODE(self): return integrate.solve_ivp(self.FNBody, (self.t[0],self.t[len(self.t)-1]),self.rv,t_eval=self.t)
修正后遗留问题
单粒子测试代码
import matplotlib.pylab as plt time = np.linspace(0, 1, 100) myN1Body = NBody(G=1, t=time) myN1Body.clearmass() myN1Body.addMass(1,[0,0],[1,0]) sol = myN1Body.solveODE() plt.figure(1) plt.plot(sol.t, sol.y[0, :]) plt.show()
运行后得到的x-t图为斜线,被误判为不符合预期;如果绘制x-y图则为水平直线,符合单粒子匀速直线运动的预期。
2粒子测试代码
time = np.linspace(0, 1, 100) myN2Body = NBody(G=1, t=time) myN2Body.clearmass() myN2Body.addMass(1,[0,0],[0,0]) myN2Body.addMass(1,[1,1],[0,0]) sol = myN2Body.solveODE() plt.figure(1) plt.plot(sol.t, sol.y[0, :]) plt.show()
运行后得到的曲线不符合两粒子相互吸引的运动规律。
问题根因
- 类属性误用:初始定义的
G、t、rv、m是类级别的属性,所有实例共享,即使添加了clearmass方法,也容易出现多实例数据干扰的问题。 - 求解器精度不足:
scipy.integrate.solve_ivp默认的相对精度rtol=1e-3、绝对精度atol=1e-6对于引力这种长程、非线性问题来说精度不够,容易出现数值偏差。 - 绘图逻辑误解:单粒子x-t图本身就是斜率为初速度的斜线,属于正确结果;如果要验证粒子运动轨迹,应该绘制x-y坐标图。
最终修复代码
import numpy as np import scipy.integrate as integrate import matplotlib.pylab as plt class NBody(): def __init__(self, G, t): # 所有属性改为实例级别,避免共享 self.G = G self.t = t self.rv = [] self.m = [] def addMass(self, mass, initPos, initVel): self.m.append(mass) self.rv.extend(initPos) self.rv.extend(initVel) def FNBody(self, t, rv): rx = rv[0::4] ry = rv[1::4] vx = rv[2::4] vy = rv[3::4] n = len(rx) F = np.zeros(n*4, dtype=np.float64) for i in range(n): ax = 0.0 ay = 0.0 for j in range(n): if i != j: dx = rx[j] - rx[i] dy = ry[j] - ry[i] r3 = (dx**2 + dy**2)**1.5 ax += self.G * self.m[j] * dx / r3 ay += self.G * self.m[j] * dy / r3 # 导数定义:dx/dt=vx, dy/dt=vy, dvx/dt=ax, dvy/dt=ay F[4*i] = vx[i] F[4*i+1] = vy[i] F[4*i+2] = ax F[4*i+3] = ay return F def solveODE(self): # 提升求解精度,更换为更稳定的RK45求解器 return integrate.solve_ivp( self.FNBody, (self.t[0], self.t[-1]), np.array(self.rv, dtype=np.float64), t_eval=self.t, method='RK45', rtol=1e-8, atol=1e-10 ) # 单粒子测试 time = np.linspace(0, 1, 100) myN1Body = NBody(G=1, t=time) myN1Body.addMass(1,[0,0],[1,0]) sol1 = myN1Body.solveODE() plt.figure(1) plt.title("单粒子运动x-y轨迹") plt.plot(sol1.y[0, :], sol1.y[1, :]) plt.xlabel("x") plt.ylabel("y") plt.show() # 2粒子测试 myN2Body = NBody(G=1, t=time) myN2Body.addMass(1,[0,0],[0,0]) myN2Body.addMass(1,[1,1],[0,0]) sol2 = myN2Body.solveODE() plt.figure(2) plt.title("两粒子运动x轨迹") plt.plot(sol2.t, sol2.y[0, :], label="粒子1 x坐标") plt.plot(sol2.t, sol2.y[4, :], label="粒子2 x坐标") plt.xlabel("时间t") plt.ylabel("x坐标") plt.legend() plt.show()
内容的提问来源于stack exchange,提问作者PileOfCheese
相关产品推荐
相关产品推荐

