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

Python优化:3D粒子位置数组的2D网格框计数提速需求

粒子网格统计提速优化方案

原代码通过双重循环逐个遍历网格框、筛选粒子并统计,面对8000帧、12000个粒子的规模时,时间复杂度极高,必然导致速度缓慢。以下是针对性的提速优化方案:

一、单帧向量化优化

核心是用numpy的向量化操作替代嵌套循环,直接计算粒子所属网格索引,通过直方图统计数量,避免逐框筛选的冗余计算。

实现代码

import numpy as np

# 假设已知Lx,Ly = Lx*3
Lx = 12
Ly = Lx * 3
ratio = 6

# 计算网格数量:x方向ratio个,y方向3*ratio个
nx = ratio
ny = 3 * ratio

# 生成网格边界(覆盖粒子位置范围)
x_bins = np.linspace(-Lx/2, Lx/2, nx + 1)
y_bins = np.linspace(-Ly/2, Ly/2, ny + 1)

# 单帧粒子数据:shape=(12000, 3),前两列x、y坐标,第三列是属性值
coord = np.random.rand(12000, 3) * np.array([Lx, Ly, 1]) - np.array([Lx/2, Ly/2, 0])

# 1. 统计每个网格内的粒子数
counts, _, _ = np.histogram2d(coord[:, 1], coord[:, 0], bins=[y_bins, x_bins])

# 2. 计算每个粒子的网格索引(0-based)
y_idx = np.digitize(coord[:, 1], y_bins) - 1
x_idx = np.digitize(coord[:, 0], x_bins) - 1

# 过滤超出边界的粒子(理论上不会存在,可省略)
valid_idx = (y_idx >= 0) & (y_idx < ny) & (x_idx >= 0) & (x_idx < nx)
y_idx = y_idx[valid_idx]
x_idx = x_idx[valid_idx]
attr_vals = coord[valid_idx, 2]

# 3. 统计每个网格内的属性值总和
sum_attr = np.zeros((ny, nx))
np.add.at(sum_attr, (y_idx, x_idx), attr_vals)

# 4. 生成最终的boxes数组
boxes = np.zeros((ny, nx, 2))
# 满足条件的网格标记为类型1
mask = (counts > 0) & ((sum_attr - counts) / counts < 0.1)
boxes[mask, 0] = 1
# 填充粒子数量
boxes[:, :, 1] = counts

二、多帧批量优化

针对8000帧的大规模数据,直接批量处理所有帧,利用numpy的广播和批量索引操作,避免逐帧循环的开销。

实现代码

# 假设所有帧数据存储为shape=(8000, 12000, 3)的数组
all_coords = np.random.rand(8000, 12000, 3) * np.array([Lx, Ly, 1]) - np.array([Lx/2, Ly/2, 0])

# 将数据扁平化,方便批量处理:shape=(8000*12000, 3)
flat_coords = all_coords.reshape(-1, 3)
y_vals = flat_coords[:, 1]
x_vals = flat_coords[:, 0]
attr_vals = flat_coords[:, 2]

# 计算每个粒子对应的帧索引、网格索引
frame_idx = np.repeat(np.arange(8000), 12000)
y_idx = np.digitize(y_vals, y_bins) - 1
x_idx = np.digitize(x_vals, x_bins) - 1

# 过滤有效索引
valid = (y_idx >= 0) & (y_idx < ny) & (x_idx >= 0) & (x_idx < nx)
y_idx = y_idx[valid]
x_idx = x_idx[valid]
frame_idx = frame_idx[valid]
attr_vals = attr_vals[valid]

# 批量统计每个帧-网格的粒子数
counts_batch = np.zeros((8000, ny, nx), dtype=int)
np.add.at(counts_batch, (frame_idx, y_idx, x_idx), 1)

# 批量统计每个帧-网格的属性值总和
sum_attr_batch = np.zeros((8000, ny, nx))
np.add.at(sum_attr_batch, (frame_idx, y_idx, x_idx), attr_vals)

# 批量生成boxes数组
boxes_batch = np.zeros((8000, ny, nx, 2))
mask_batch = (counts_batch > 0) & ((sum_attr_batch - counts_batch) / counts_batch < 0.1)
boxes_batch[mask_batch, 0] = 1
boxes_batch[:, :, :, 1] = counts_batch

三、额外优化手段

  1. 内存优化:如果数据量过大无法一次性加载,可分批次处理(比如每次处理100帧),避免内存溢出。
  2. Numba加速循环:如果必须保留循环逻辑,用Numba的JIT编译加速,示例如下:
from numba import njit

@njit
def compute_boxes_numba(coord, Lx, Ly, ratio):
    lx = Lx / ratio
    ly = Ly / ratio / 3
    ny = int(Ly / ly)
    nx = int(Lx / lx)
    boxes = np.zeros((ny, nx, 2))
    for i in range(ny):
        for j in range(nx):
            x_l = -Lx/2 + lx * j
            y_l = -Ly/2 + ly * i
            x_h = x_l + lx
            y_h = y_l + ly
            mask = (coord[:,0] >= x_l) & (coord[:,0] <= x_h) & (coord[:,1] >= y_l) & (coord[:,1] <= y_h)
            temp_cord = coord[mask]
            n = temp_cord.shape[0]
            if n > 0:
                avg = (temp_cord[:,2].sum() - n) / n
                if avg < 0.1:
                    boxes[i,j,0] = 1
                boxes[i,j,1] = n
    return boxes

该版本比纯Python循环快几十倍。
3. 二进制文件预处理:若数据存储为多文件,用np.memmap直接映射文件,无需一次性加载全部数据,提升读取速度并降低内存占用。

内容的提问来源于stack exchange,提问作者Coldz

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 18:10:35