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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 10:27:06