如何在Python中优化多个二维粒子的均方位移(MSD)计算性能
二维多粒子均方位移计算优化
问题定义
需要计算多个二维粒子的均方位移(MSD),定义如下:
其中i为粒子索引,Dt为时间间隔,t为时间,vec(x)为粒子的二维位置,计算时需要对所有可能的时间t取平均。
基础Numpy实现
已基于Numpy完成基础版本实现,输入参数pos为三维np.array,维度为(粒子数, 时间步, 坐标维度),完整代码如下:
import numpy as np import matplotlib.pyplot as plt import time # 初始化模拟数据 np.random.seed(1) nTime = 10**4 nParticles = 3 pos = np.zeros((nParticles, nTime, 2)) # 维度顺序:粒子、时间步、坐标 for t in range(1, nTime): pos[:, t, :] = pos[:, t-1, :] + ( np.random.random((nParticles, 2)) - 0.5) # MSD直接计算实现 def MSD_direct(pos): Dt_r = np.arange(1, pos.shape[1]-1) MSD = np.empty((nParticles, len(Dt_r))) dMSD = np.empty((nParticles,len(Dt_r))) for k, Dt in enumerate(Dt_r): SD = np.sum((pos[:, Dt:,:] - pos[:, 0:-Dt,:])**2, axis = -1) MSD[:,k] = np.mean( SD , axis = 1) dMSD[:,k] = np.std( SD, axis = 1 ) / np.sqrt(SD.shape[1]) return Dt_r, MSD, dMSD start_time = time.time() Dt_r, MSD_d, dMSD_d = MSD_direct(pos) print("MSD_direct -- Time: %s s\n" % (time.time() - start_time)) # 绘图输出 plt.figure() for i in range(nParticles): plt.plot(pos[i,:,0]) plt.xlabel('t') plt.ylabel('x') plt.savefig('pos_x.png', dpi = 300) plt.figure() for i in range(nParticles): plt.plot(pos[i,:,1]) plt.xlabel('t') plt.ylabel('y') plt.savefig('pos_y.png', dpi = 300) plt.figure() for i in range(nParticles): plt.fill_between(Dt_r, MSD_d[i,:]+dMSD_d[i,:], MSD_d[i,:] - dMSD_d[i,:], alpha = 0.5) plt.plot(Dt_r, MSD_d[i,:]) plt.xlabel('Dt') plt.ylabel('MSD') plt.savefig('MSD.png', dpi = 300)
代码运行输出:MSD_direct -- Time: 7.793087720870972 s
生成的可视化结果如下:


当前Numpy版本的问题是代码中仍然保留了Dt维度的循环,无法通过完全向量化操作消除该循环进一步提升性能。
Numba优化实现
已基于Numba重写计算逻辑,相比原始Numpy版本性能提升约2倍,实现代码如下:
import numba as nb @nb.jit(fastmath=True,parallel=True) def MSD_numba(pos): Dt_r = np.arange(1, pos.shape[1]-1) MSD = np.empty((nParticles, len(Dt_r))) dMSD = np.empty((nParticles,len(Dt_r))) for i in nb.prange(nParticles): for Dt in Dt_r: SD = (pos[i, Dt:, 0] - pos[i, 0:-Dt, 0])**2 + (pos[i, Dt:, 1] - pos[i, 0:-Dt, 1])**2 MSD[i, Dt-1] = np.mean( SD ) dMSD[i, Dt-1] = np.std( SD ) / np.sqrt(len(SD)) return Dt_r, MSD, dMSD start_time = time.time() Dt_r, MSD_n, dMSD_n = MSD_numba(pos) print("MSD_numba -- Time: %s s" % (time.time() - start_time)) print("MSD_numba -- All close to MSD_direct: %r\n" %(np.allclose(MSD_n, MSD_d) ) )
运行输出:
MSD_numba -- Time: 4.520232915878296 s MSD_numba -- All close to MSD_direct: True
目前希望进一步优化代码性能,获得更高的运行效率。
内容的提问来源于stack exchange,提问作者Puco4
相关产品推荐
相关产品推荐

