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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 09:22:08