粒子模拟运行缓慢且高速粒子出现穿透碰撞问题求助
粒子模拟程序优化方案
问题概述
开发了一个粒子模拟程序,用于观测随机初始位置和速度的粒子间相互作用,但遇到两个问题:
- 模拟整体运行缓慢
- 粒子高速运动时互相穿透,无法触发正常碰撞
问题分析与解决
1. 运行缓慢的解决
根源
- 碰撞检测采用双层循环,时间复杂度为O(n²),当粒子数Np=100时,每次要执行约5000次距离计算,效率极低
- Matplotlib动画更新方式错误:每次通过
ax.clear()重新绘制散点,再加上plt.pause()强制暂停,大幅拖慢渲染速度
解决办法
- 碰撞检测优化:采用网格空间分区,将粒子按位置分配到不同网格中,仅检测同一网格和相邻网格内的粒子,减少不必要的距离计算
- 动画渲染优化:复用同一个scatter绘图对象,直接更新其位置、颜色等数据,避免重复创建绘图元素;移除
plt.pause(),让FuncAnimation自动控制帧率
2. 高速粒子穿透的解决
根源
- 固定时间步长
dt:高速粒子在一个dt周期内的移动距离超过粒子半径,直接穿过对方,碰撞检测无法捕捉到此次碰撞 - Euler积分精度低:先检测碰撞再更新位置的逻辑,误差积累会导致位置偏移
- 碰撞公式错误:当前速度更新未考虑粒子质量,且碰撞后未修正粒子重叠的位置
解决办法
- 动态时间步长:预测所有可能的碰撞时间,取最小的碰撞时间作为当前步进的
dt,确保粒子不会在一步内穿过对方 - 碰撞位置修正:碰撞时将粒子移动到刚好接触的位置,避免重叠
- 正确弹性碰撞公式:使用考虑粒子质量的弹性碰撞速度更新公式
修改后的完整代码
import numpy as np import matplotlib.pyplot as plt from matplotlib.animation import FuncAnimation class Particle: def __init__(self, id=0, charge=1.602E-19, r=np.zeros(2), v=np.zeros(2), rad=0.01, m=1): self.id = id self.r = r # 粒子的x、y坐标 self.v = v # 粒子的x、y速度分量 self.rad = rad # 粒子半径 self.m = m # 粒子质量 self.charge = charge * (np.random.randint(0, 2) * 2 - 1) # 随机正负电荷 self.color = "blue" if self.charge > 0 else "green" # 正电粒子蓝色,负电粒子绿色 class Sim: X = 2 # 环境尺寸 Y = 2 def __init__(self, dt=0.00005, Np=100): self.dt = dt # 初始时间步长 self.Np = Np # 粒子数量 self.particles = [Particle(i) for i in range(Np)] # 初始化网格参数,网格大小设为粒子最大直径的2倍 self.cell_size = 2 * max(p.rad for p in self.particles) self.grid_x = int(np.ceil(self.X / self.cell_size)) self.grid_y = int(np.ceil(self.Y / self.cell_size)) def _build_grid(self): # 构建粒子网格,将粒子分到对应网格中 grid = {(i, j): [] for i in range(self.grid_x) for j in range(self.grid_y)} for p in self.particles: # 计算粒子所在网格坐标(从- X/2, -Y/2转换到0,0起始) cell_x = int((p.r[0] + self.X/2) // self.cell_size) cell_y = int((p.r[1] + self.Y/2) // self.cell_size) # 确保网格坐标在范围内 cell_x = max(0, min(self.grid_x-1, cell_x)) cell_y = max(0, min(self.grid_y-1, cell_y)) grid[(cell_x, cell_y)].append(p) return grid def _predict_collision_time(self, p1, p2): # 预测两个粒子的碰撞时间(如果会碰撞) dr = p1.r - p2.r dv = p1.v - p2.v dist_sq = np.dot(dr, dr) rad_sum = p1.rad + p2.rad rad_sum_sq = rad_sum ** 2 # 相对速度点乘相对位置 dv_dot_dr = np.dot(dv, dr) if dv_dot_dr >= 0: # 粒子互相远离,不会碰撞 return np.inf # 计算判别式 disc = dv_dot_dr ** 2 - np.dot(dv, dv) * (dist_sq - rad_sum_sq) if disc < 0: # 没有实根,不会碰撞 return np.inf # 取较小的时间(最近的碰撞) t = (-dv_dot_dr - np.sqrt(disc)) / np.dot(dv, dv) return t if t > 0 else np.inf def _predict_wall_collision_time(self, p): # 预测粒子和墙壁的碰撞时间 t_list = [] # 左右墙 if p.v[0] != 0: if p.v[0] > 0: t = (self.X/2 - p.r[0] - p.rad) / p.v[0] else: t = (-self.X/2 - p.r[0] + p.rad) / p.v[0] if t > 0: t_list.append(t) # 上下墙 if p.v[1] != 0: if p.v[1] > 0: t = (self.Y/2 - p.r[1] - p.rad) / p.v[1] else: t = (-self.Y/2 - p.r[1] + p.rad) / p.v[1] if t > 0: t_list.append(t) return min(t_list) if t_list else np.inf def coll_det(self, dt): # 先更新位置到dt时间后 for p in self.particles: p.r += dt * p.v # 墙壁碰撞处理 for p in self.particles: # 左右墙 if p.r[0] - p.rad < -self.X/2: p.r[0] = -self.X/2 + p.rad p.v[0] *= -1 elif p.r[0] + p.rad > self.X/2: p.r[0] = self.X/2 - p.rad p.v[0] *= -1 # 上下墙 if p.r[1] - p.rad < -self.Y/2: p.r[1] = -self.Y/2 + p.rad p.v[1] *= -1 elif p.r[1] + p.rad > self.Y/2: p.r[1] = self.Y/2 - p.rad p.v[1] *= -1 # 粒子间碰撞处理(用网格优化) grid = self._build_grid() visited = set() for (cell_x, cell_y), particles in grid.items(): # 检查当前网格和相邻网格的粒子 for dx in [-1, 0, 1]: for dy in [-1, 0, 1]: neighbor_cell = (cell_x + dx, cell_y + dy) if neighbor_cell not in grid: continue neighbor_particles = grid[neighbor_cell] for i, p1 in enumerate(particles): for p2 in neighbor_particles: if p1.id >= p2.id or (p1.id, p2.id) in visited: continue visited.add((p1.id, p2.id)) visited.add((p2.id, p1.id)) dist = np.linalg.norm(p1.r - p2.r) rad_sum = p1.rad + p2.rad if dist <= rad_sum + 1e-8: # 允许微小误差 # 修正位置到刚好接触 overlap = rad_sum - dist if dist < 1e-10: # 粒子完全重合,随机偏移 dir_vec = np.array([np.random.randn(), np.random.randn()]) dir_vec /= np.linalg.norm(dir_vec) else: dir_vec = (p1.r - p2.r) / dist p1.r += dir_vec * overlap * 0.5 p2.r -= dir_vec * overlap * 0.5 # 正确的弹性碰撞速度公式 m1, m2 = p1.m, p2.m r1, r2 = p1.r, p2.r v1, v2 = p1.v, p2.v n = (r1 - r2) / np.linalg.norm(r1 - r2) v1_new = v1 - (2 * m2 / (m1 + m2)) * np.dot(v1 - v2, n) * n v2_new = v2 - (2 * m1 / (m1 + m2)) * np.dot(v2 - v1, n) * n p1.v = v1_new p2.v = v2_new def increment(self): # 预测所有可能的碰撞时间,取最小的作为当前dt min_collision_time = np.inf # 粒子间碰撞时间 grid = self._build_grid() visited = set() for (cell_x, cell_y), particles in grid.items(): for dx in [-1,0,1]: for dy in [-1,0,1]: neighbor_cell = (cell_x+dx, cell_y+dy) if neighbor_cell not in grid: continue neighbor_particles = grid[neighbor_cell] for p1 in particles: for p2 in neighbor_particles: if p1.id >= p2.id or (p1.id,p2.id) in visited: continue visited.add((p1.id,p2.id)) t = self._predict_collision_time(p1,p2) if t < min_collision_time: min_collision_time = t # 墙壁碰撞时间 for p in self.particles: t = self._predict_wall_collision_time(p) if t < min_collision_time: min_collision_time = t # 取预测碰撞时间和初始dt的较小值,避免步进过大 dt = min(min_collision_time, self.dt) if min_collision_time != np.inf else self.dt self.coll_det(dt) def get_particle_data(self): positions = np.array([p.r for p in self.particles]) colors = [p.color for p in self.particles] sizes = [p.rad * 10000 for p in self.particles] # 放大半径方便显示 return positions, colors, sizes # 初始化模拟 sim = Sim(Np=100) for particle in sim.particles: particle.r = np.random.uniform([-sim.X/2, -sim.Y/2], [sim.X/2, sim.Y/2], size=2) particle.v = np.array([np.random.uniform(-50,50), np.random.uniform(-50,50)]) sim.particles[0].color = "red" # 标记第一个粒子为红色 # 初始化绘图 fig, ax = plt.subplots(figsize=(8,8)) ax.set_xlim(-sim.X/2, sim.X/2) ax.set_ylim(-sim.Y/2, sim.Y/2) ax.set_aspect('equal') scatter = ax.scatter([], [], s=[], alpha=0.7) def init(): scatter.set_offsets(np.empty((0,2))) scatter.set_sizes([]) scatter.set_color([]) return scatter, def update(frame): sim.increment() positions, colors, sizes = sim.get_particle_data() scatter.set_offsets(positions) scatter.set_sizes(sizes) scatter.set_color(colors) return scatter, # 创建动画,帧率设为30 animation = FuncAnimation(fig, update, frames=12000, init_func=init, interval=1000/30, blit=True) plt.show()
关键优化点说明
- 网格分区碰撞检测:将粒子按位置分到网格,仅检测相邻网格的粒子,将O(n²)复杂度降到接近O(n)
- 动态时间步长:预测最近的碰撞时间,确保粒子不会在一步内穿过对方,彻底解决高速穿透问题
- Matplotlib渲染优化:复用scatter对象,仅更新数据,开启blit加速,大幅提升动画流畅度
- 正确弹性碰撞:采用考虑粒子质量的速度更新公式,碰撞时修正位置避免重叠
内容的提问来源于stack exchange,提问作者Onur Karakaş
相关产品推荐
相关产品推荐

