分片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
相关产品推荐
相关产品推荐

