如何在JAX的顺序循环内实现带集合通信的设备级数据并行?
JAX代码从数据并行转单设备集合通信实现方案
要解决的两个核心问题
- 把原基于设备分片的数据并行实现,改成单设备上用集合通信模拟的版本
- 避开在
jax.lax.scan里用jax.pmap+jax.lax.psum导致的重复编译和冗余数据传输问题
核心修改逻辑
原代码是把数组分片到8个物理设备上跑并行,现在改成在单设备里用逻辑并行维度模拟多设备,再用集合通信操作模拟设备间的同步,同时保持scan的高效性,完全不用pmap就能实现原并行逻辑。
具体修改步骤与代码
1. 移除设备分片相关代码
直接删掉原代码里创建设备mesh、sharding的部分,以及jit装饰器里的out_shardings参数,因为现在只跑单设备。
2. 新增逻辑并行维度
给参数J和初始状态y加一个大小为8的维度(对应原8个设备),用vmap批量生成每个逻辑设备的参数,把物理设备并行转成单设备内的逻辑batch并行。
3. 重写scan内的计算逻辑
在calc_body里,先让每个逻辑设备独立计算更新量,然后用jax.lax.all_reduce模拟多设备间的求和同步,完全不用pmap,避免额外编译开销。
完整修改后的代码
import functools as ft import jax as jx import jax.numpy as jnp import jax.random as jrn import jax.lax as jlx # 对应原代码的8个设备,设为逻辑并行维度的大小 NUM_DEVICES = 8 @ft.partial(jx.jit, static_argnums=0) def test(m): key = jrn.PRNGKey(0) # 给每个逻辑设备分配独立的随机key key = jrn.split(key, NUM_DEVICES) # 批量生成每个逻辑设备的J矩阵,形状变为(8, m, m) J = jrn.vmap( lambda k: 0.1 * jrn.uniform(k, shape=(m, m), dtype='f8', minval=-0.1, maxval=0.1) + jnp.eye(m, dtype='f8') )(key) # 批量生成每个逻辑设备的初始y,形状变为(8, m, 1) y = jrn.vmap( lambda k: jrn.uniform(k, shape=(m, 1), dtype='f8', minval=0.0, maxval=1.0) )(key) def calc_body(y, _): # 每个逻辑设备独立计算J@y的更新量 update = 1.0e-06 * jnp.einsum('dmn,dni->dmi', J, y) # 用all_reduce模拟多设备间的求和同步,对应原数据并行的通信逻辑 update = jlx.all_reduce('sum', update, axis_name='devices') # 更新当前y new_y = y + update return new_y, None # 在scan里指定axis_name,让all_reduce能识别逻辑并行维度 y, empty = jlx.scan( ft.partial(calc_body), y, None, length=1000, axis_name='devices' ) return y # 编译并运行 test_comp = test.lower(2**14).compile() y = test_comp() print(f"输出形状:{y.shape}") # 输出(8, 16384, 1),对应8个逻辑设备的结果
关键细节说明
- 逻辑并行维度:用
vmap批量生成参数,替代原物理设备的分片,把多设备并行转成单设备内的batch维度计算,无需占用8个物理设备。 - 集合通信模拟:用
jlx.all_reduce的sum操作模拟原数据并行中设备间的更新同步,指定axis_name就能对应到我们新增的逻辑维度,完全无需pmap。 - scan的axis_name:在
jlx.scan里指定axis_name,让内部的集合通信能正确识别并行维度,不会触发额外的编译和数据传输。
为何要规避scan内用pmap?
原代码外层的test已经用jax.jit编译完成,如果在calc_body里加pmap,会导致:
- 每一次scan迭代都会触发
pmap的编译,产生大量冗余编译步骤,拖慢运行速度。 pmap会自动执行设备间的数据聚集和分散操作,而外层已经是单设备编译函数,这种额外通信会大幅降低性能。
内容的提问来源于stack exchange,提问作者DavidJ
相关产品推荐
相关产品推荐

