组织 JAX 多进程计算与分布式批次加载

组织 JAX 多进程计算与分布式批次加载

原文:The JAX authors。本文合并翻译整理 Introduction to multi-controller JAX 与 Distributed data loading,依据 2026 年 10 月 5 日的 latest 文档。原页版权:© Copyright 2024, The JAX Authors。

当计算需要的设备多于一台主机能够容纳的数量时,可以在 GPU 集群、Cloud TPU pod 或多台 CPU 主机上运行多个 JAX Python 进程。每个进程也称为 controller。它们先通过 jax.distributed.initialize() 加入同一个分布式系统,然后用统一的数组分片机制,像编程一台拥有很多设备的机器那样描述计算。

一个 jax.Array 可以跨越全部进程。只要各个参与进程以相同顺序对它执行相同 JAX 运算,XLA 就能切分计算,并在适当位置插入集合通信。设备之间有 TPU ICI 或 NVLink 等高速链路时会利用它们,否则使用可用的主机网络,例如 Ethernet 或 InfiniBand。建议先熟悉 分布式数组与自动并行化。

通常每个进程运行同一份 Python 脚本;与单进程最明显的差别,发生在外部数据进入 JAX 或结果离开 JAX 的时候。因此,本文先建立多进程计算模型,再把同一模型延伸到数据加载。

先运行一个完整的本地演示

将下面的程序保存为 toy.py,同时在各进程中启动。该教学例子固定使用总计 8 个设备。

# toy.py
import sys
import jax
import jax.numpy as jnp
from jax.sharding import NamedSharding, PartitionSpec as P
import numpy as np

proc_id = int(sys.argv[1])
num_procs = int(sys.argv[2])
jax.distributed.initialize("localhost:10000", num_procs, proc_id)

assert jax.device_count() == 8
mesh = jax.make_mesh((4, 2), ("i", "j"))
global_data = np.arange(32).reshape((4, 8))
sharding = NamedSharding(mesh, P("i", "j"))
global_array = jax.device_put(global_data, sharding)
assert global_array.shape == global_data.shape

for shard in global_array.addressable_shards:
    print(f"device {shard.device} has local data {shard.data}")

global_result = jnp.sum(jnp.sin(global_array))
print(f"process={proc_id} got result: {global_result}")

mesh 包含所有进程的设备,global_array 在逻辑上是一个数组,物理上分布在各设备上。jnp.sum 会对整个数组求和,而最终标量被复制到各进程,所以每个进程都能打印它。

原文提供的 Bash 启动方式如下:4 个进程,每个进程使用 2 个 CPU 设备。JAX_NUM_CPU_DEVICES 在这里用于演示设备数量,不代表额外的物理硬件。

export JAX_NUM_CPU_DEVICES=2
num_processes=4
range=$(seq 0 $(($num_processes - 1)))

for i in $range; do
  python toy.py $i $num_processes > /tmp/toy_$i.out &
done
wait

for i in $range; do
  echo "=================== process $i output ==================="
  cat /tmp/toy_$i.out
  echo
done

原文输出中,进程 0 的两个块是 0–3、4–7;进程 1 是 8–11、12–15;进程 2 是 16–19、20–23;进程 3 是 24–27、28–31,每块都是一行。设备标识依次为 cpu:0/1、cpu:2048/2049、cpu:4096/4097、cpu:6144/6145。四个进程都输出 -0.12398731708526611。

相同脚本也可在一个进程上使用 8 个 CPU 设备:

JAX_NUM_CPU_DEVICES=8 python toy.py 0 1

此时设备编号为 cpu:0 到 cpu:7,八个数据块的划分以及原文结果都相同。改变的是设备归属的进程数量,运算和全局数据没有改变。上述数值是原文输出,本次未运行脚本。

本地设备、可寻址设备与全局设备

一个进程能够直接向其本地设备传输数据、从其内存取回数据并启动计算,无需其他进程参与。设备通常通过 PCI 等方式连接到该进程所在主机。一个设备只能属于一个进程的本地设备集合,集合彼此不重叠;可用 jax.local_devices() 查询。

“可寻址”(addressable)在这里通常与“本地”同义。所有进程的设备合起来称为全局设备,由 jax.devices() 查询。各进程完成 initialize() 后,才建立全局设备视图。类似地,NumPy 本地数组仅在某一进程中可见,JAX 全局数组在概念上则跨越进程。

启动多个进程:JAX 不代替作业调度器

实际部署通常在多台主机上各启动一个或多个进程,可以用 SSH、Slurm 或 Kubernetes 完成。JAX 不会因为调用一次 Python 程序就自动启动其他进程。无论采用什么启动方式,每个进程都必须调用 jax.distributed.initialize()。在受支持的 Slurm、Kubernetes 和 Cloud TPU 环境中,所需参数可自动发现,因此可以不传参数。

初始化顺序:必须在调用 jax.devices()、jax.local_devices(),或通过 jax.numpy 等执行设备计算之前初始化。使用不访问设备的配置功能没有问题;如果后端已经初始化,随后调用分布式初始化会报错。

协调通信的安全边界:默认协调服务连接既不加密也不认证,演示配置应限定在受信网络。需要保护这些连接时,应按 JAX 的安全说明 核对 mTLS 配置。本文没有创建云资源、开放端口或执行部署命令。

GPU 示例

假设有 4 台虚拟机,每台 2 块 GPU,每台都运行下面的程序,并从命令行读取协调地址、当前进程编号和总进程数:

# gpu_example.py
import jax
import sys

coord_addr = sys.argv[1]
proc_id = int(sys.argv[2])
num_procs = int(sys.argv[3])

jax.distributed.initialize(
    coordinator_address=coord_addr,
    num_processes=num_procs,
    process_id=proc_id,
)
print("process id =", jax.process_index())
print("global devices =", jax.devices())
print("local devices =", jax.local_devices())

若第一台虚拟机的地址是示例私网地址 192.168.0.1,第一台运行 python3 gpu_example.py 192.168.0.1:8000 0 4,第二台把进程编号改成 1,依次到 3。原文中每个进程都看见 GPU 0–7;进程 0 的本地设备是 0、1,进程 1 是 2、3。

Cloud TPU 示例

原文以含 4 台主机的 v5litepod-16 为例,通过命令取得各主机地址,复制脚本并并行启动。执行警告:这些命令会调用 Cloud TPU API、向每台远端主机复制文件,并通过 SSH 升级安装 JAX TPU 包;它们会改动远端环境,也可能产生云资源费用。只有在明确授权的测试项目和目标 TPU 上使用,先核对项目、区域、实例名、SSH 目标和费用,勿直接对生产主机运行。

TPU_NAME=jax-demo
EXTERNAL_IPS=$(gcloud compute tpus tpu-vm describe $TPU_NAME --zone 'us-central1-a' \
                 | grep externalIp | cut -d: -f2)
cat << EOF > demo.py
import jax
jax.distributed.initialize()
if jax.process_index() == 0:
  print(jax.devices())
EOF
echo $EXTERNAL_IPS | xargs -n 1 -P 0 bash -c '
scp demo.py $0:
ssh $0 "pip -q install -U jax[tpu]"
ssh $0 "python demo.py" '

xargs 并行启动 SSH 命令,每台主机执行相同程序。只有进程 0 打印全局设备列表。原文列出了 16 个 TPU 设备,进程 0 对应编号 0、1、4、5;进程 1 对应 2、3、6、7;进程 2 对应 8、9、12、13;进程 3 对应 10、11、14、15,其坐标覆盖 4×4 网格。

这段是原文的环境相关操作示例,会复制文件、升级远端包并启动程序;应先核对账号、目标地址和部署版本,不把其输出当成本次实测结果。

Kubernetes 示例

Kubernetes 中每个 Pod 也运行相同 Python 程序。准备工作包括:使用包含 JAX 的镜像;选择 JobSet 或 indexed Job,使每个副本对应一个进程;为服务账号配置列出本作业 Pod 的权限,以便发现对等进程。原文另行提供 服务账号和 RBAC 文件。

下面保留原文的双副本 JobSet 结构。它包含 null 占位符及外部服务账号依赖,不是一份可以直接提交的完整清单。镜像、GPU 数量与私有镜像凭据需按环境填写;不需要私有镜像凭据时,应按 Kubernetes 配置规范处理该字段。

apiVersion: jobset.x-k8s.io/v1alpha2
kind: JobSet
metadata:
  name: jaxjob
spec:
  replicatedJobs:
  - name: workers
    template:
      spec:
        parallelism: 2
        completions: 2
        backoffLimit: 0
        template:
          spec:
            serviceAccountName: jax-job-sa
            restartPolicy: Never
            imagePullSecrets:
            - name: null
            containers:
            - name: main
              image: null
              imagePullPolicy: Always
              resources:
                limits:
                  cpu: 1
                  nvidia.com/gpu: null
              command:
                - python
              args:
                - -c
                - |
                  import jax
                  jax.distributed.initialize()
                  print(jax.devices())
                  print(jax.local_devices())
                  assert jax.process_count() > 1
                  assert len(jax.devices()) > len(jax.local_devices())

完成配置及前置资源以后,原文依次提交清单、查看 Pod 和日志。执行警告:kubectl apply 会按当前 kubeconfig 上下文向集群创建或更新资源,JobSet、ServiceAccount 与 RBAC 规则可能持续存在并扩展权限。运行前确认目标集群、命名空间、权限范围、所用镜像和可能的资源费用,并制定清理步骤。

kubectl apply -f example.yaml
kubectl get pods -l jobset.sigs.k8s.io/jobset-name=jaxjob
kubectl logs -l jobset.sigs.k8s.io/jobset-name=jaxjob

示例中的两个 Pod 最终为 Completed。它们都看到 [CudaDevice(id=0), CudaDevice(id=1)],各自的本地列表分别只有 0 或 1。确认每个进程拥有相同全局视图、不同本地设备后,再把内联程序替换成实际 JAX 作业。后文所有多进程片段都假定初始化已完成,并由需要参与的进程共同运行。

让 Mesh 跨越主机,并尊重网络拓扑

Mesh 把一个设备数组与各轴名称配对。直接构造可以展示原理,实际使用通常更适合 jax.make_mesh(),因为它能选择性能更好的设备次序。

from jax.sharding import Mesh

mesh = Mesh(jax.devices(), ("a",))
# 实际通常采用:
mesh = jax.make_mesh((jax.device_count(),), ("a",))

默认情况下,两种写法都能涵盖各进程的全局设备。但系统规模扩大后,设备间带宽不再均匀,应让通信量最大的轴落在最快的链路上。

JAX/TPU 术语 NVIDIA GPU 中的对应概念
ICI,芯片间高速互联 NVLink
DCN,切片之间的网络 InfiniBand、Ethernet、EFA、TCPXO 等
slice,切片 NVLink domain,例如 GB200-NVL72 的机架级域,或 HGX B200 NVL8 的单主机域

初始化时 JAX 会检测设备所属切片,并设置 slice_index。单切片通常可用 make_mesh;跨多个切片应使用 create_hybrid_device_mesh。原文指出,make_mesh 对多切片拓扑会报错。

from jax.experimental.mesh_utils import create_hybrid_device_mesh

# devices_per_slice、num_slices 需按实际拓扑定义
mesh = Mesh(
    create_hybrid_device_mesh((1, devices_per_slice), (num_slices, 1)),
    axis_names=("dcn", "ici"),
)

这里把 dcn 轴映射到切片间网络,ici 轴映射到切片内互联。DGX H100 一类系统中,每个节点往往就是一个 NVLink domain,即使只有两个节点也属于多切片。仅把 jax.devices() 重排成目标形状,不能可靠地保证合理通信性能。

跨主机数组与矩阵计算

有了 Mesh,就可用 NamedSharding 指定全局数组布局。例如,下面把 32×32 的全 1 数组沿第一轴切分:

arr = jax.device_put(jnp.ones((32, 32)), NamedSharding(mesh, P("a")))
if jax.process_index() == 0:
    jax.debug.visualize_array_sharding(arr)

在前述单切片、16 设备 TPU 例子中,原文的可视化将数组分成 16 个横条,设备顺序是 0、1、4、5、2、3、6、7、8、9、12、13、10、11、14、15。

下面的独立例子让编译器自动决定中间结果的分片方式,执行分布式矩阵乘法和 ReLU:

mesh = jax.make_mesh(
    (jax.device_count() // 2, 2), ("a", "b"),
    axis_types=(jax.sharding.AxisType.Auto,) * 2,
)

def device_put(x, spec):
    return jax.device_put(x, NamedSharding(mesh, spec))

x = device_put(jnp.ones((4096, 2048)), P("a", "b"))
y = device_put(jnp.ones((2048, 4096)), P("b", None))
z = jax.nn.relu(x @ y)

if jax.process_index() == 0:
    jax.debug.visualize_array_sharding(z)
    print(z.sharding)

原文结果沿 a 轴切分,在 b 轴复制,等价于 P("a", None)。16 个设备的图中共有 8 个横条,每条复制到两个设备:0/1、4/5、8/9、12/13、2/3、6/7、10/11、14/15。矩阵乘法需要沿 b 做求和通信。这里明确使用 Auto 轴;若使用默认 Explicit 轴,由于两个操作数都在收缩维度上分片,该运算会要求提供 out_sharding,不能机械省略这一差别。

集体运算顺序必须一致

跨进程数组上的运算可能包含集合通信屏障。所有参与设备的控制进程必须以相同顺序执行相同计算,否则其他设备可能永远等待。例如前三个进程执行 x @ y,最后一个却执行 y @ x,就可能挂起。这项一致性假设大多不会被自动检查。

最容易维护的做法是让所有进程执行相同代码,尤其检查依赖 jax.process_index() 的分支:打印可以只在一个进程做,涉及通信的计算不能随意移入这样的分支。

跨非本地设备分片的数组不能直接通过 np.asarray(z) 取回完整值,因为当前进程没有全部数据。原文对此给出 RuntimeError,提示可以使用 process_allgather 或查看 addressable_shards。直接打印该数组通常只显示形状和 dtype。

# 所有参与进程先共同执行复制:
w = device_put(z, P(None, None))
# 再只在一个进程显示已可用的值:
if jax.process_index() == 0:
    print(np.asarray(w))

不能把 device_put 放进仅进程 0 执行的分支,否则它可能等待其余进程而死锁。也可以使用 jax.experimental.multihost_utils.process_allgather。若只需要查看本地数据,z.addressable_shards 不需要通信,任意进程子集都可独立访问;但该属性不在 jax.jit 内使用。

只让部分设备参与

Mesh 也可以只覆盖设备子集,以便让不同设备并行处理不同任务。原文选择全局设备列表的后半部分:

num_devices = jax.device_count() // 2
mesh = jax.make_mesh(
    (num_devices,), ("a",),
    devices=jax.devices()[num_devices:],
    axis_types=(jax.sharding.AxisType.Explicit,),
)
sharding = NamedSharding(mesh, P("a"))
data = np.arange(64).reshape((8, 8))
x = jax.device_put(data, sharding)

if jax.process_index() == 0:
    jax.debug.visualize_array_sharding(x)
    print(x.sharding)

print(f"Devices attached to process {jax.process_index()}: {jax.local_devices()}")
for shard in x.addressable_shards:
    print(f"device {shard.device} has local data {shard.data}")

result = jnp.sum(jnp.sin(x))
print(f"process={jax.process_index()} got result: {result!r}")

在四主机 v5litepod-16 示例里,只有进程 2、3 的设备参与。布局设备顺序为 8、9、10、11、15、14、13、12。进程 0、1 的 addressable_shards 为空;进程 3 的本地块包括 16–23、24–31、32–39、40–47。进程 2、3 能打印数值结果 Array(0.09658563, dtype=float32),进程 0、1 只显示标量形状与 dtype。

顺序一致的要求只针对参与该分片的进程。这个例子可以只在进程 2、3 上做计算,也可以让其他进程执行相同代码并得到没有本地数据的数组引用。

跨进程传输与流水线

jax.device_put 可以在跨进程设备之间传输数据,使用可用的高速链路。调用必须发生在所有属于源分片或目标分片的主机上。典型用途是把流水线第一阶段的结果传给另一组设备上的第二阶段:

num_devices = jax.device_count() // 2
mesh_first_half = jax.make_mesh(
    (num_devices,), ("a",), devices=jax.devices()[:num_devices],
    axis_types=(jax.sharding.AxisType.Explicit,),
)
mesh_second_half = jax.make_mesh(
    (num_devices,), ("a",), devices=jax.devices()[num_devices:],
    axis_types=(jax.sharding.AxisType.Explicit,),
)
sharding_first_half = NamedSharding(mesh_first_half, P("a"))
sharding_second_half = NamedSharding(mesh_second_half, P("a"))

data = np.arange(64).reshape((8, 8))
x = jax.device_put(data, sharding_first_half)

@jax.jit
def f(x):
    return x  # 这里换成第一阶段的 JAX 运算

@jax.jit
def g(x):
    return x  # 这里换成第二阶段的 JAX 运算

y = f(x)
z = jax.device_put(y, sharding_second_half)
result = g(z)

这里所有进程都执行 y = f(x),即使某个进程没有第一阶段的本地设备,也能拿到后续传输需要的 y 引用。JAX 异步派发允许输入已准备好的不同设备计算并行进行,进而形成微批次流水线:

# 至少需要4个进程。每个阶段用不同进程的第一个本地设备。
pipeline_stages = [f, g, f, g]
devices = [
    jax.local_devices(process_index=i)[0]
    for i in range(len(pipeline_stages))
]
microbatches = [
    np.arange(512**2).reshape((512, 512))
    for _ in range(12)
]
results = []
# 所有进程运行循环,包括没有阶段设备的进程。
for mb in microbatches:
    for d, s in zip(devices, pipeline_stages):
        mb = jax.device_put(mb, d)
        mb = s(mb)
    results.append(mb)

各微批次没有数据依赖,一台设备结束前一批次后即可开始下一批次,并与后续传输重叠。原文特意按进程选择设备;直接使用 jax.devices()[:4] 往往会把四个阶段都放在进程 0。

当前源文还列出限制:跨进程 device_put 需要支持跨主机传输的 TPU/GPU 后端,或通过 jax_cross_host_transfer_socket_address 启用基于 DCN 的传输;源、目标分片目前必须拥有相同设备数和相同分片形状。

从外部数据构造全局数组的三个接口

第一种方法是每个进程加载完整数组,再用 device_put 取出其本地设备需要的部分,前面的 toy 例子正是这样。另两种更常用:每个进程只加载本地所需数据,然后调用 make_array_from_process_local_data;或者为每个本地设备准备独立数组,再用 make_array_from_single_device_arrays 组装。

按进程加载特别适合数据并行批次,因为无需逐个指定哪一小批落在哪一台设备。下面是原文的随机数据演示;它假定进程和设备分布规则、批次能整除相应数量。

batch_size = 1024
per_process_batch_size = batch_size // jax.process_count()
per_device_batch_size = batch_size // jax.device_count()

mesh = jax.make_mesh((jax.device_count(),), ("batch",))
sharding = NamedSharding(mesh, P("batch"))
process_batch = np.random.rand(per_process_batch_size, 2048, 42)

global_batch = jax.make_array_from_process_local_data(sharding, process_batch)
assert global_batch.shape[0] == batch_size
assert process_batch.shape[0] == per_process_batch_size
assert global_batch.addressable_shards[0].data.shape[0] == per_device_batch_size

按设备构造的接口更通用,但数据移动由调用者显式负责。下面给每个设备放一个值,逻辑上组成全局编号数组:

from jax.experimental import multihost_utils

shape = (jax.process_count(), jax.local_device_count())
mesh = jax.make_mesh(shape, ("i", "j"))
sharding = NamedSharding(mesh, P("i", "j"))
local_arrays = [
    jax.device_put(
        jnp.array([[jax.process_index() * jax.local_device_count() + i]]),
        device,
    )
    for i, device in enumerate(jax.local_devices())
]
global_array = jax.make_array_from_single_device_arrays(
    shape=shape, sharding=sharding, arrays=local_arrays,
)
assert np.all(
    multihost_utils.process_allgather(global_array)
    == np.arange(jax.device_count()).reshape(global_array.shape)
)

该构造函数把已经准备好的本地单设备数组视为一个全局数组,组装本身不搬移数据。也可以只覆盖进程 0、1;其他进程传空列表,并明确 dtype:

num_participating_processes = 2
shape = (num_participating_processes, jax.local_device_count())
devices = (jax.local_devices(process_index=0)
           + jax.local_devices(process_index=1))
mesh = jax.make_mesh(
    shape, ("i", "j"),
    axis_types=(jax.sharding.AxisType.Explicit,) * 2,
    devices=devices,
)
sharding = NamedSharding(mesh, P("i", "j"))

if jax.process_index() in (0, 1):
    local_arrays = [
        jax.device_put(
            jnp.array([[jax.process_index() * jax.local_device_count() + i]]),
            device,
        )
        for i, device in enumerate(jax.local_devices())
    ]
else:
    local_arrays = []

array = jax.make_array_from_single_device_arrays(
    shape=shape, sharding=sharding, arrays=local_arrays, dtype=jnp.int32,
)
if jax.process_index() in (0, 1):
    for shard in array.addressable_shards:
        assert shard.data.size == 1
else:
    assert not array.addressable_shards

分布式加载的核心问题:哪个设备需要哪块数据

分布式加载通常比“由一个进程读取全部数据再通过 RPC 分发”或“每个进程都读取完整数据”更省读取和传输开销,但实现也更复杂。训练循环可能被数据加载拖慢,重复加载还会占用额外网络带宽。

更重要的是正确性:如果数据被送到错误设备上,运算仍可能正常运行。JAX 无法知道输入本来应该是什么,结果可能悄悄出错。这些原则也适用于从检查点加载模型权重、加载大型空间分片图像等场景,不只适用于批次样本。

每个 jax.Array 都关联一个 Sharding,描述全局每个设备需要的切片。创建外部输入时,应先根据并行策略或数据产生方式选好 Sharding,再用其 addressable_devices 确定当前进程负责的设备。

例如,一个 64×128 数组分到 4 个进程、每进程 2 个设备,共 8 个设备。若沿第二维一维切分,每个设备取得 64×16,连续分配时进程 0 的两台设备持有前 64×32 的数据。也可以改变设备对应顺序,或做二维切分;无论布局如何变化,加载器都必须提供与声明一致的切片。

四种加载组织方式

方法 每个进程做什么 主要取舍
加载完整全局数据 读取完整值,只把本地设备需要的部分传过去。 最直接,但丢弃了许多重复读取的数据;全局数据很小时可能可以接受。
每设备一个流水线 为每个本地设备创建加载器,只读它需要的切片。 读取量有效,也便于单独考虑每个设备,但多个并发加载器可能影响性能。
每进程一个合并流水线 用单个加载器取得全部本地设备需要的数据,然后按设备切分。 原文将其视为加载效率最高的组织方式,但映射和数据选择逻辑最复杂;实际性能仍需测量。
便于加载的布局,再在计算内重新分片 例如每进程先读全局数据的四分之一,再由加速器互联转换成目标布局。 保持每进程单流水线,全局数据只加载一次,也更灵活;代价是额外加速器通信和两套 Sharding 描述。

第四种方法适用于加载器无法直接提供目标切片的情况。例如目标是二维布局,但每个进程更容易加载一列,可以先构造与列式输入一致的全局数组,再在计算开始处使用 jax.lax.with_sharding_constraint() 转成目标布局。重分片会使用 TPU ICI、NVLink 等加速器链路,也可能拖慢通信受限的工作负载。

完整复制与部分复制

复制表示多个设备持有同一份数据。四种加载方式都能用于复制,只是某些进程可能需要加载相同切片。

完整复制时,每个设备拥有整个数组。4 个进程、每进程 2 个设备,就有 8 份完整数据,每份独占一台设备。部分复制时有多份完整数据,但每份又跨多个设备切分;其布局有许多可能性。

例如,一份数据在同一进程的两台设备之间切分,四个进程各持一份,总计 4 份,那么每个进程都要加载完整数据。另一种布局仍让每份数据跨两个设备,但设备来自不同进程:原文例子中进程 0、1 都需要第一行,进程 2、3 都需要第二行。不能仅凭“部分复制”就推断每个进程的输入相同。

纯数据并行:批次与副本的对应可以交换

纯数据并行在每台设备上复制一个模型,各模型副本获得不同的 per-replica batch。全局输入数组是当前步骤所有副本批次沿 batch 轴的拼接,每个分片就是一个副本批次。

关键性质是:在原文所讨论的批次顺序无关的数据并行场景中,不必关心某个副本拿到哪一小批,因为各副本执行相同计算。可以重新排列全局批次中的小批次,再把每进程的批次拆给本地设备。一般数组任意交换分片会改变含义,这个简化只适用于这里的数据并行语义。

这样每个进程可以维护一条独立数据流。下面用 tf.data 演示;初始化应先完成,代码只取第一批,实际训练需要每步继续取批次。

import tensorflow as tf
# jax、np 已在前文导入,分布式初始化已经完成。
ds = tf.data.Dataset.from_tensor_slices(
    [np.ones((16, 3)) * i for i in range(100)]
)
ds = ds.shard(num_shards=jax.process_count(), index=jax.process_index())
per_process_batch = next(ds.as_numpy_iterator())

mesh = jax.make_mesh((jax.device_count(),), ("batch",))
sharding = jax.sharding.NamedSharding(mesh, P("batch"))
global_batch_array = jax.make_array_from_process_local_data(
    sharding, per_process_batch
)

编辑说明:这里使用规范命名空间 jax.sharding.NamedSharding,并将原文迭代器的 .next() 改写为 Python 的 next(...)。示例数据和划分意图不变。

数据并行加模型并行:同一副本必须得到同一批数据

模型并行将一个模型副本切分到多台设备。纯模型并行只有一个模型副本,输入通常在相关设备上完整复制;数据并行与模型并行结合时,有多个模型副本,每个副本跨多设备,同一副本的设备需要相同的批次,不同副本使用不同批次。

两个进程各有两台设备;模型副本A跨两个进程的第一台设备,必须都接收批次A;模型副本B跨两个进程的第二台设备,必须都接收批次B。
原创概念图:批次可以在模型副本之间交换,但同一副本各设备持有的数据必须一致。图为简化的跨进程模型并行布局。

模型副本限制在一个进程内

先考虑 2 个进程、每进程 4 台设备,每个模型副本占用同进程中的 2 台设备。每个进程有 2 个副本,全局共 4 个副本。全局批次按 4 个副本切分,每块在所属副本的两台设备上复制。

可以沿用按进程读取的数据流,再让分片配置完成副本内复制。某个小批次究竟分给哪一个模型副本可以交换,但一个模型副本内部不能混用两个不同小批次。JAX 不会检测“声明复制、实际数据不同”的错误,后续运算可能静默出错。

per_process_batches = [np.ones((16, 3)) * i for i in range(100)]
ds = tf.data.Dataset.from_tensor_slices(per_process_batches)
ds = ds.shard(num_shards=jax.process_count(), index=jax.process_index())
per_process_batch = next(ds.as_numpy_iterator())

num_model_replicas_per_process = 2
num_model_replicas_total = num_model_replicas_per_process * jax.process_count()

mesh_devices = np.array([
    jax.local_devices(process_idx)
    for process_idx in range(jax.process_count())
])
mesh_devices = mesh_devices.reshape(num_model_replicas_total, -1)
for replica_devices in mesh_devices:
    num_processes = len(set(d.process_index for d in replica_devices))
    assert num_processes == 1

mesh = jax.sharding.Mesh(
    mesh_devices, ["model_replicas", "data_parallelism"]
)
sharding = jax.sharding.NamedSharding(
    mesh, P("model_replicas")
)
global_batch_array = jax.make_array_from_process_local_data(
    sharding, per_process_batch
)

原文的第二个 Mesh 轴命名为 data_parallelism,本文保留该名称;判断数据是否复制应看 PartitionSpec,它只沿 model_replicas 切分,没有沿第二轴切分,所以在第二轴复制。每行设备必须属于同一进程,代码中的断言检查了这一前提。

模型副本跨越多个进程

如果一个副本放不进单进程,或为了利用更好的互联而采用跨进程设备布局,加载就需要额外协调。回到 4 个进程、每进程 2 台设备的例子,全局仍有 4 个模型副本,每个副本跨 2 台设备,但这两台设备属于不同进程。

原文的布局要求进程 0 与 2 加载相同的两份副本批次,进程 1 与 3 加载另一组相同的两份批次。每个进程还必须保证这两份批次没有送反。即使形状完全正确,只要同一副本的不同设备拿到不同批次,声明的复制关系就被破坏,JAX 仍可能不报错。

一种可行组织方式是按模型副本索引划分输入流水线,而不是按进程索引划分。每个进程只读取其设备所属副本的分片,并保持相同顺序;共享模型副本的进程加载相同批次。之后可用 jax.make_array_from_callback() 按全局切片回调创建数组。

这也是从“程序能启动”走向“结果正确”的最后一步:初始化和全局形状检查只能确认部分条件,数据内容、设备映射、复制一致性以及集合调用顺序都需要按实际任务验证。

合并来源:多控制器 JAX 入门;分布式数据加载。本文保留两篇的全部技术章节、示例意图和限制,重复长设备输出改用中文及编号说明。示意图由本文编辑自绘。所有代码只做静态审阅;未启动进程、连接主机、安装包、部署 Kubernetes、执行训练或测量性能。原页面标注 © Copyright 2024, The JAX Authors,但未在已保存页面中单独列出文档文字的开放许可;本文不为文档文字补加未证实的许可标签。JAX 项目代码仓库采用 Apache License 2.0,本文所用代码片段按该许可提供,完整文本见下方。本文保留 JAX 作者署名与原始链接。

Apache License 2.0

本文所用 JAX 代码片段依据 Apache License 2.0 提供。完整许可文本如下。

                                 Apache License
                           Version 2.0, January 2004
                        http://www.apache.org/licenses/

   TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION

   1. Definitions.

      "License" shall mean the terms and conditions for use, reproduction,
      and distribution as defined by Sections 1 through 9 of this document.

      "Licensor" shall mean the copyright owner or entity authorized by
      the copyright owner that is granting the License.

      "Legal Entity" shall mean the union of the acting entity and all
      other entities that control, are controlled by, or are under common
      control with that entity. For the purposes of this definition,
      "control" means (i) the power, direct or indirect, to cause the
      direction or management of such entity, whether by contract or
      otherwise, or (ii) ownership of fifty percent (50%) or more of the
      outstanding shares, or (iii) beneficial ownership of such entity.

      "You" (or "Your") shall mean an individual or Legal Entity
      exercising permissions granted by this License.

      "Source" form shall mean the preferred form for making modifications,
      including but not limited to software source code, documentation
      source, and configuration files.

      "Object" form shall mean any form resulting from mechanical
      transformation or translation of a Source form, including but
      not limited to compiled object code, generated documentation,
      and conversions to other media types.

      "Work" shall mean the work of authorship, whether in Source or
      Object form, made available under the License, as indicated by a
      copyright notice that is included in or attached to the work
      (an example is provided in the Appendix below).

      "Derivative Works" shall mean any work, whether in Source or Object
      form, that is based on (or derived from) the Work and for which the
      editorial revisions, annotations, elaborations, or other modifications
      represent, as a whole, an original work of authorship. For the purposes
      of this License, Derivative Works shall not include works that remain
      separable from, or merely link (or bind by name) to the interfaces of,
      the Work and Derivative Works thereof.

      "Contribution" shall mean any work of authorship, including
      the original version of the Work and any modifications or additions
      to that Work or Derivative Works thereof, that is intentionally
      submitted to Licensor for inclusion in the Work by the copyright owner
      or by an individual or Legal Entity authorized to submit on behalf of
      the copyright owner. For the purposes of this definition, "submitted"
      means any form of electronic, verbal, or written communication sent
      to the Licensor or its representatives, including but not limited to
      communication on electronic mailing lists, source code control systems,
      and issue tracking systems that are managed by, or on behalf of, the
      Licensor for the purpose of discussing and improving the Work, but
      excluding communication that is conspicuously marked or otherwise
      designated in writing by the copyright owner as "Not a Contribution."

      "Contributor" shall mean Licensor and any individual or Legal Entity
      on behalf of whom a Contribution has been received by Licensor and
      subsequently incorporated within the Work.

   2. Grant of Copyright License. Subject to the terms and conditions of
      this License, each Contributor hereby grants to You a perpetual,
      worldwide, non-exclusive, no-charge, royalty-free, irrevocable
      copyright license to reproduce, prepare Derivative Works of,
      publicly display, publicly perform, sublicense, and distribute the
      Work and such Derivative Works in Source or Object form.

   3. Grant of Patent License. Subject to the terms and conditions of
      this License, each Contributor hereby grants to You a perpetual,
      worldwide, non-exclusive, no-charge, royalty-free, irrevocable
      (except as stated in this section) patent license to make, have made,
      use, offer to sell, sell, import, and otherwise transfer the Work,
      where such license applies only to those patent claims licensable
      by such Contributor that are necessarily infringed by their
      Contribution(s) alone or by combination of their Contribution(s)
      with the Work to which such Contribution(s) was submitted. If You
      institute patent litigation against any entity (including a
      cross-claim or counterclaim in a lawsuit) alleging that the Work
      or a Contribution incorporated within the Work constitutes direct
      or contributory patent infringement, then any patent licenses
      granted to You under this License for that Work shall terminate
      as of the date such litigation is filed.

   4. Redistribution. You may reproduce and distribute copies of the
      Work or Derivative Works thereof in any medium, with or without
      modifications, and in Source or Object form, provided that You
      meet the following conditions:

      (a) You must give any other recipients of the Work or
          Derivative Works a copy of this License; and

      (b) You must cause any modified files to carry prominent notices
          stating that You changed the files; and

      (c) You must retain, in the Source form of any Derivative Works
          that You distribute, all copyright, patent, trademark, and
          attribution notices from the Source form of the Work,
          excluding those notices that do not pertain to any part of
          the Derivative Works; and

      (d) If the Work includes a "NOTICE" text file as part of its
          distribution, then any Derivative Works that You distribute must
          include a readable copy of the attribution notices contained
          within such NOTICE file, excluding those notices that do not
          pertain to any part of the Derivative Works, in at least one
          of the following places: within a NOTICE text file distributed
          as part of the Derivative Works; within the Source form or
          documentation, if provided along with the Derivative Works; or,
          within a display generated by the Derivative Works, if and
          wherever such third-party notices normally appear. The contents
          of the NOTICE file are for informational purposes only and
          do not modify the License. You may add Your own attribution
          notices within Derivative Works that You distribute, alongside
          or as an addendum to the NOTICE text from the Work, provided
          that such additional attribution notices cannot be construed
          as modifying the License.

      You may add Your own copyright statement to Your modifications and
      may provide additional or different license terms and conditions
      for use, reproduction, or distribution of Your modifications, or
      for any such Derivative Works as a whole, provided Your use,
      reproduction, and distribution of the Work otherwise complies with
      the conditions stated in this License.

   5. Submission of Contributions. Unless You explicitly state otherwise,
      any Contribution intentionally submitted for inclusion in the Work
      by You to the Licensor shall be under the terms and conditions of
      this License, without any additional terms or conditions.
      Notwithstanding the above, nothing herein shall supersede or modify
      the terms of any separate license agreement you may have executed
      with Licensor regarding such Contributions.

   6. Trademarks. This License does not grant permission to use the trade
      names, trademarks, service marks, or product names of the Licensor,
      except as required for reasonable and customary use in describing the
      origin of the Work and reproducing the content of the NOTICE file.

   7. Disclaimer of Warranty. Unless required by applicable law or
      agreed to in writing, Licensor provides the Work (and each
      Contributor provides its Contributions) on an "AS IS" BASIS,
      WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
      implied, including, without limitation, any warranties or conditions
      of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
      PARTICULAR PURPOSE. You are solely responsible for determining the
      appropriateness of using or redistributing the Work and assume any
      risks associated with Your exercise of permissions under this License.

   8. Limitation of Liability. In no event and under no legal theory,
      whether in tort (including negligence), contract, or otherwise,
      unless required by applicable law (such as deliberate and grossly
      negligent acts) or agreed to in writing, shall any Contributor be
      liable to You for damages, including any direct, indirect, special,
      incidental, or consequential damages of any character arising as a
      result of this License or out of the use or inability to use the
      Work (including but not limited to damages for loss of goodwill,
      work stoppage, computer failure or malfunction, or any and all
      other commercial damages or losses), even if such Contributor
      has been advised of the possibility of such damages.

   9. Accepting Warranty or Additional Liability. While redistributing
      the Work or Derivative Works thereof, You may choose to offer,
      and charge a fee for, acceptance of support, warranty, indemnity,
      or other liability obligations and/or rights consistent with this
      License. However, in accepting such obligations, You may act only
      on Your own behalf and on Your sole responsibility, not on behalf
      of any other Contributor, and only if You agree to indemnify,
      defend, and hold each Contributor harmless for any liability
      incurred by, or claims asserted against, such Contributor by reason
      of your accepting any such warranty or additional liability.

   END OF TERMS AND CONDITIONS

   APPENDIX: How to apply the Apache License to your work.

      To apply the Apache License to your work, attach the following
      boilerplate notice, with the fields enclosed by brackets "[]"
      replaced with your own identifying information. (Don't include
      the brackets!)  The text should be enclosed in the appropriate
      comment syntax for the file format. We also recommend that a
      file or class name and description of purpose be included on the
      same "printed page" as the copyright notice for easier
      identification within third-party archives.

   Copyright [yyyy] [name of copyright owner]

   Licensed under the Apache License, Version 2.0 (the "License");
   you may not use this file except in compliance with the License.
   You may obtain a copy of the License at

       http://www.apache.org/licenses/LICENSE-2.0

   Unless required by applicable law or agreed to in writing, software
   distributed under the License is distributed on an "AS IS" BASIS,
   WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
   See the License for the specific language governing permissions and
   limitations under the License.
© 版权声明
THE END
喜欢就支持一下吧
点赞0 分享
评论 抢沙发

请登录后发表评论

    暂无评论内容