能否简化计算3D笛卡尔网格四阶精度偏导数的Python代码?
问题描述
需要计算形状为(Nmesh, Nmesh, Nmesh)(如Nmesh=512)的3D笛卡尔网格的一阶偏导数,采用四阶精度格式,公式为:
f'(n) = 2*(f(n+1)-f(n-1))/(3dh) - (f(n+2)-f(n-2))/(12dh)
编写的Python函数如下,但不同条件分支代码重复度高,希望进一步简化:
import numpy as np def partial( arr, Nmesh, dh, axis=0 ): ''' The partial derivative of the 3D Cartesian mesh with periodic boundary conditions, computed by the fourth-order accuracy scheme. ''' dh1 = 2/(3*dh) dh2 = -1/(12*dh) if (axis==0): arr_ = np.pad(arr, [(2,2), (0,0), (0,0)], mode='wrap') diff1 = arr_[3:Nmesh+3,:,:] - arr_[1:Nmesh+1,:,:] diff2 = arr_[4:Nmesh+4,:,:] - arr_[0:Nmesh,:,:] return dh1*diff1 + dh2*diff2 elif (axis==1): arr_ = np.pad(arr, [(0,0), (2,2), (0,0)], mode='wrap') diff1 = arr_[:,3:Nmesh+3,:] - arr_[:,1:Nmesh+1,:] diff2 = arr_[:,4:Nmesh+4,:] - arr_[:,0:Nmesh,:] return dh1*diff1 + dh2*diff2 elif (axis==2): arr_ = np.pad(arr, [(0,0), (0,0), (2,2)], mode='wrap') diff1 = arr_[:,:,3:Nmesh+3] - arr_[:,:,1:Nmesh+1] diff2 = arr_[:,:,4:Nmesh+4] - arr_[:,:,0:Nmesh] return dh1*diff1 + dh2*diff2
简化后的实现
可以通过动态生成padding参数和切片索引消除重复分支,简化后的代码如下:
import numpy as np def partial(arr, Nmesh, dh, axis=0): ''' 计算3D笛卡尔网格的一阶偏导数(周期边界条件),采用四阶精度格式。 ''' dh1 = 2 / (3 * dh) dh2 = -1 / (12 * dh) # 仅在指定轴的两端各补2个元素,其余轴不补 padding = [(2, 2) if i == axis else (0, 0) for i in range(3)] arr_padded = np.pad(arr, padding, mode='wrap') # 构造对应轴的切片索引,其余轴保持全选 def build_slice(start, end): return tuple( slice(start, end) if i == axis else slice(None) for i in range(3) ) diff1 = arr_padded[build_slice(3, Nmesh+3)] - arr_padded[build_slice(1, Nmesh+1)] diff2 = arr_padded[build_slice(4, Nmesh+4)] - arr_padded[build_slice(0, Nmesh)] return dh1 * diff1 + dh2 * diff2
简化思路
- 动态生成padding参数:通过列表推导式,根据传入的
axis参数,仅在目标轴的两端各填充2个元素,其余轴填充0个,避免重复写三次pad参数。 - 统一切片逻辑:用一个辅助函数
build_slice生成对应轴的切片索引,其余轴用slice(None)等价于:,这样不管axis是0、1还是2,都能复用同一段切片代码。 - 消除分支判断:去掉原有的
if/elif分支,所有轴的处理逻辑统一,后续修改格式或调整参数时只需维护一处代码,提升可维护性。
内容的提问来源于stack exchange,提问作者Stephen Wong
相关产品推荐
相关产品推荐

