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
三、额外优化手段
- 内存优化:如果数据量过大无法一次性加载,可分批次处理(比如每次处理100帧),避免内存溢出。
- 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
相关产品推荐
相关产品推荐

