JAX实现3D扩散核较CUDA慢5倍,求性能优化方案
JAX扩散核性能优化问题
问题背景
我将现有GPU CUDA代码移植为JAX版本的扩散核,结果正确但运行速度比CUDA慢约5倍。目前diff函数采用了和等效C++代码相同的公式,不确定实现是否最优,希望找到进一步提升性能的方法。
现有代码
import jax import jax.numpy as jnp import numpy as np from jax import jit from functools import partial from timeit import default_timer as timer # Numpy-like operation. @partial(jit, static_argnums=(6, 7, 8)) def diff(at, a, visc, dxidxi, dyidyi, dzidzi, itot, jtot, ktot): i_c = jnp.s_[1:ktot-1, 1:jtot-1, 1:itot-1] i_w = jnp.s_[1:ktot-1, 1:jtot-1, 0:itot-2] i_e = jnp.s_[1:ktot-1, 1:jtot-1, 2:itot ] i_s = jnp.s_[1:ktot-1, 0:jtot-2, 1:itot-1] i_n = jnp.s_[1:ktot-1, 2:jtot , 1:itot-1] i_b = jnp.s_[0:ktot-2, 1:jtot-1, 1:itot-1] i_t = jnp.s_[2:ktot , 1:jtot-1, 1:itot-1] at_new = at.at[i_c].add( visc * ( + ( (a[i_e] - a[i_c]) - (a[i_c] - a[i_w]) ) * dxidxi + ( (a[i_n] - a[i_c]) - (a[i_c] - a[i_s]) ) * dyidyi + ( (a[i_t] - a[i_c]) - (a[i_c] - a[i_b]) ) * dzidzi ) ) return at_new itot = 384; jtot = 384; ktot = 384; float_type = jnp.float32 nloop = 30; ncells = itot*jtot*ktot; dxidxi = float_type(0.1) dyidyi = float_type(0.1) dzidzi = float_type(0.1) visc = float_type(0.1) @jit def init_a(index): return (index/(index+1))**2 ## FIRST EXPERIMENT. at = jnp.zeros((ktot, jtot, itot), dtype=float_type) index = jnp.arange(ncells, dtype=float_type) a = init_a(index) del(index) a = a.reshape(ktot, jtot, itot) at = diff(at, a, visc, dxidxi, dyidyi, dzidzi, itot, jtot, ktot).block_until_ready() print("(first check) at={0}".format(at.flatten()[itot*jtot+itot+itot//2])) # Time the loop start = timer() for i in range(nloop): at = diff(at, a, visc, dxidxi, dyidyi, dzidzi, itot, jtot, ktot).block_until_ready() end = timer() print("Time/iter: {0} s ({1} iters)".format((end-start)/nloop, nloop)) print("at={0}".format(at.flatten()[itot*jtot+itot+itot//4]))
性能优化建议
1. 简化计算逻辑,减少冗余运算
当前代码中的二阶差分计算可简化,(a[i_e] - a[i_c]) - (a[i_c] - a[i_w])等价于a[i_e] - 2*a[i_c] + a[i_w],减少一次减法操作降低计算量:
# 替换原计算部分 laplacian = ( (a[i_e] - 2*a[i_c] + a[i_w]) * dxidxi + (a[i_n] - 2*a[i_c] + a[i_s]) * dyidyi + (a[i_t] - 2*a[i_c] + a[i_b]) * dzidzi ) at_new = at.at[i_c].add(visc * laplacian)
2. 用JAX LAX原语替代手动切片,优化内存访问
手动切片会生成多个子数组增加内存开销,改用jax.lax.shift实现数组偏移,避免创建中间子数组:
from jax import lax @jit def diff(at, a, visc, dxidxi, dyidyi, dzidzi): ktot, jtot, itot = a.shape i_c = jnp.s_[1:ktot-1, 1:jtot-1, 1:itot-1] # 用shift实现各方向偏移 a_e = lax.shift(a, (0, 0, 1), padding_value=0.0)[i_c] a_w = lax.shift(a, (0, 0, -1), padding_value=0.0)[i_c] a_n = lax.shift(a, (0, 1, 0), padding_value=0.0)[i_c] a_s = lax.shift(a, (0, -1, 0), padding_value=0.0)[i_c] a_t = lax.shift(a, (1, 0, 0), padding_value=0.0)[i_c] a_b = lax.shift(a, (-1, 0, 0), padding_value=0.0)[i_c] a_c = a[i_c] laplacian = ( (a_e - 2*a_c + a_w) * dxidxi + (a_n - 2*a_c + a_s) * dyidyi + (a_t - 2*a_c + a_b) * dzidzi ) return at.at[i_c].add(visc * laplacian)
同时移除静态参数,直接通过数组shape获取维度,提升JIT编译灵活性。
3. 用卷积实现拉普拉斯算子,利用GPU优化
扩散核本质是3D拉普拉斯运算,JAX的jax.lax.conv_general_dilated调用GPU高度优化的卷积内核,可大幅提升性能:
from jax import lax @jit def diff(at, a, visc, dxidxi, dyidyi, dzidzi): # 定义3D拉普拉斯卷积核,对应各方向的二阶差分系数 kernel = jnp.array([ [[0, 0, 0], [0, dzidzi, 0], [0, 0, 0]], [[0, dyidyi, 0], [dxidxi, -2*(dxidxi+dyidyi+dzidzi), dxidxi], [0, dyidyi, 0]], [[0, 0, 0], [0, dzidzi, 0], [0, 0, 0]] ], dtype=a.dtype) # 卷积计算拉普拉斯,添加batch和channel维度适配卷积接口 laplacian = lax.conv_general_dilated( a[None, ..., None], kernel[None, ..., None], window_strides=(1,1,1), padding='VALID', dimension_numbers=('NCDHW', 'OIHWD', 'NCDHW') )[0, ..., 0] # 移除额外维度 # 更新at数组,边界保持原值 return at.at[1:-1,1:-1,1:-1].add(visc * laplacian)
4. 替换Python循环为JAX LAX循环,减少CPU-GPU通信
原Python for循环每次迭代都需同步GPU结果,改用jax.lax.fori_loop让整个循环在GPU上执行,消除同步开销:
from jax import lax # 定义循环体 def loop_body(i, at): return diff(at, a, visc, dxidxi, dyidyi, dzidzi) # 执行循环 start = timer() at = lax.fori_loop(0, nloop, loop_body, at).block_until_ready() end = timer()
5. 优化数组内存布局
JAX默认使用C风格行优先布局,若原CUDA代码是列优先(Fortran风格),可调整数组维度顺序为(itot, jtot, ktot),让内存访问更连续,提升缓存命中率:
# 初始化时调整维度顺序 at = jnp.zeros((itot, jtot, ktot), dtype=float_type) a = init_a(index).reshape(itot, jtot, ktot)
同时对应修改diff函数中的切片和偏移逻辑。
内容的提问来源于stack exchange,提问作者Chiel
相关产品推荐
相关产品推荐

