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

粒子模拟运行缓慢且高速粒子出现穿透碰撞问题求助

粒子模拟程序优化方案

问题概述

开发了一个粒子模拟程序,用于观测随机初始位置和速度的粒子间相互作用,但遇到两个问题:

  • 模拟整体运行缓慢
  • 粒子高速运动时互相穿透,无法触发正常碰撞

问题分析与解决

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ş

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 14:59:52