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

分片JAX数组的离散差分计算能否通过CPU分片实现加速?

能否通过CPU核心分片的方式加速JAX数组上的离散差分计算?

本文尝试遵循JAX自动并行文档的方法,针对不同分片数组对JAX NumPy API调用进行AOT编译,通过多种设备网格测试差分方向上的横向与纵向分区,测得的运行时间显示此类分片方式并无性能收益。

os.environ["XLA_FLAGS"] = (
    f'--xla_force_host_platform_device_count=8'
)

import jax as jx
import jax.numpy as jnp
import jax.experimental.mesh_utils as jxm
import jax.sharding as jsh

def calc_fd_kernel(x):
    # 沿第一轴计算一阶差分
    return jnp.diff(
        x, 1, axis=0, prepend=jnp.zeros((1, *x.shape[1:]))
    )

def make_fd(shape, shardings):
    # 编译差分核工厂函数
    return jx.jit(
        calc_fd_kernel,
        in_shardings=shardings,
        out_shardings=shardings,
    ).lower(
        jx.ShapeDtypeStruct(shape, jnp.dtype('f8'))
    ).compile()

# 创建待分片的二维数组
n = 2**12
shape = (n,n,)

x = jx.random.normal(jx.random.PRNGKey(0), shape, dtype='f8')

shardings_test = {
    (1, 1,) : jsh.PositionalSharding(jxm.create_device_mesh((1,), devices=jx.devices("cpu")[:1])).reshape(1, 1),
    (8, 1,) : jsh.PositionalSharding(jxm.create_device_mesh((8,), devices=jx.devices("cpu")[:8])).reshape(8, 1),
    (1, 8,) : jsh.PositionalSharding(jxm.create_device_mesh((8,), devices=jx.devices("cpu")[:8])).reshape(1, 8),
}

x_test = {
    mesh : jx.device_put(x, shardings)
    for mesh, shardings in shardings_test.items()
}

calc_fd_test = {
    mesh : make_fd(shape, shardings)
    for mesh, shardings in shardings_test.items()
}

for x_mesh, calc_fd_mesh in zip(x_test.values(), calc_fd_test.values()):
    %timeit calc_fd_mesh(x_mesh).block_until_ready()

测试运行结果

  • (1, 1) 分片模式:48.9 ms ± 414 µs per loop (7次运行的均值±标准差,每次10循环)
  • (8, 1) 分片模式:977 ms ± 34.5 ms per loop (7次运行的均值±标准差,每次1循环)
  • (1, 8) 分片模式:48.3 ms ± 1.03 ms per loop (7次运行的均值±标准差,每次10循环)

内容的提问来源于stack exchange,提问作者DavidJ

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 06:33:26