JAX 多控制器分布式容错编程实战:live_devices、心跳检测与集体通信取消
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
在 JAX 中,多控制器(multi-controller)分布式程序默认是**容错免疫(fault-intolerant)**的:任何一台机器崩溃,所有机器都会一起崩溃。本文基于 JAX 官方《Fault Tolerant Distributed JAX》指南(即本仓库 docs/501/fault-tolerance.rst),系统讲解如何在 GPU 上让多控制器 JAX 程序真正具备容错能力。读完本文,你将掌握三类核心技术能力:通过配置关闭“命运共享(fate sharing)”、利用核心容错 APIlive_devices判断哪些设备仍然存活并保证分布式代码块的原子执行,以及理解协调服务、NCCL 集体通信取消等底层实现原理,最终能写出可应对进程故障甚至支持故障进程恢复的分布式训练程序。
适用前提:JAX 的容错支持目前仍处于实验阶段,仅在 GPU 后端完整可用;在 TPU 上有粗糙边缘、可能有缺陷且接口随时可能变更,请自行评估风险。如果你需要在 TPU 上实现类似能力,官方建议参考 Pathways 方案。本仓库内所有示例脚本均位于 docs/_static/fault_tolerance/ 目录,可直接对照阅读。
背景:为什么多控制器 JAX 默认不可靠
在开始容错之前,需要先理解两个关键前提:
- 多控制器 JAX允许将一个 JAX 程序分布到多台机器上并行执行。相关的入门知识可参考仓库中的 多控制器 JAX 教程。
- 默认行为是“同生共死”:只要其中任何一个进程崩溃,所有其他进程也会主动崩溃。分布式系统术语称这种设计为命运共享(fate sharing)。
要构建容错的分布式 JAX 程序,你需要解决三个层层递进的问题:
- 不让存活进程崩溃(关闭命运共享);
- 不让存活进程永久卡死(取消卡在失败集体通信上的调用);
- 让存活进程知道“谁还活着”,并原子地推进程序(
live_devices及其屏障、原子性语义)。
下文将按照“基础 → 示例 → 实现原理”三部分依次展开,与官方文档结构保持一致。
第一部分:容错基础
1.1 默认机制演示:一个进程死,全体进程死
官方文档用一个极简脚本说明默认行为。其核心逻辑(完整文件见 docs/_static/fault_tolerance/while_loop.py)如下:
def main(_: Sequence[str]) -> None: jax.distributed.initialize( coordinator_address="localhost:9000", num_processes=_NUM_PROCESSES.value, process_id=_PROCESS_ID.value, local_device_ids=[_PROCESS_ID.value], heartbeat_timeout_seconds=10, ) while True: print(time.time()) time.sleep(1)脚本先调用jax.distributed.initialize初始化多控制器 JAX,然后进入死循环每秒打印一次当前时间。local_device_ids参数确保每个进程只被分配四块 GPU 中的一块;heartbeat_timeout_seconds稍后解释。
在一台拥有四块 GPU 的虚拟机上,用四个终端分别启动四个进程:
python example.py --i=0 --n=4 # 终端 1 python example.py --i=1 --n=4 # 终端 2 python example.py --i=2 --n=4 # 终端 3 python example.py --i=3 --n=4 # 终端 4此时四个进程每秒都会打印时间。现在杀掉第四个进程:
pkill -9 -f 'python example.py --i=3 --n=4'大约十秒后,其余进程会全部终止并打印类似下面的错误:
E0926 17:26:32.075402 157988 coordination_service_agent.cc:332] Polled an error from coordination service (this can be an error from this or another task). F0926 17:26:32.075587 157988 client.h:77] Terminating process because the JAX distributed service detected fatal errors. This most likely indicates that another task died; see the other task logs for more details. Disable Python buffering, i.e. `python -u`, to be sure to see all the previous output. absl::Status: UNAVAILABLE: The following tasks are unhealthy (stopped sending heartbeats): /job:jax_worker/replica:0/task:3 The tasks have crashed. Check the task logs for an earlier error, or scheduler events (e.g. preemption, eviction) to debug further.结论:当某个多控制器 JAX 进程发现同伴进程崩溃时,它决定自己也崩溃。jax.distributed.initialize的heartbeat_timeout_seconds参数决定了进程在断定同伴“已死亡”之前会等待多久——上面示例传入10,因此第一个到第三个进程大约在杀掉第四个进程十秒后崩溃。从错误信息中的“stopped sending heartbeats”可以看出,健康检测依赖进程间定期发送的心跳。
1.2 让进程存活:关闭命运共享
关闭命运共享只需在脚本中加入一行环境变量和一行配置,完整脚本见 docs/_static/fault_tolerance/dont_fail.py:
import os os.environ['XLA_FLAGS'] = '--xla_gpu_nccl_terminate_on_error=false' ... def main(_: Sequence[str]) -> None: jax.config.update("jax_enable_recoverability", True) jax.distributed.initialize(...) # 参数同上 while True: print(time.time()) time.sleep(1)这里两个开关的作用分别是:
| 开关 | 类型 | 作用 |
|---|---|---|
--xla_gpu_nccl_terminate_on_error=false | XLA 标志(写入XLA_FLAGS) | 禁止 GPU 集体通信在出错时触发进程自杀 |
jax_enable_recoverability | JAX 配置项(jax.config.update) | 启用可恢复性语义,允许进程在同伴死亡后继续运行 |
jax_enable_recoverability这一配置选项在源码中定义于 jax/_src/distributed.py,与分布式初始化逻辑同处一个模块。再次以四个进程运行脚本并杀掉第四个,你会观察到其余三个进程安然无恙地继续执行——命运共享已被成功关闭。
1.3 进程 0 的不可替代性
接下来尝试失败进程 0。你会发现:即使关闭了命运共享,所有四个进程仍然会全部终止,错误信息大致如下:
E0929 17:42:48.594192 1044529 coordination_service_agent.cc:332] Polled an error from coordination service (this can be an error from this or another task). F0929 17:42:48.594200 1044529 client.h:77] Terminating process because the JAX distributed service detected fatal errors. ... absl::Status: UNAVAILABLE: Failed to send RPC to coordination service. Either the leader task was preempted/died/restarted unexpectedly or this task is experiencing network issues. ...进程 0 是特殊的:它运行着一个名为**协调服务(coordination service)**的 RPC 服务,所有进程都通过它与彼此协调。如果协调服务本身失败,其他进程除了失败别无选择。这一点的详细原理见下文第三部分。
1.4 存活了但卡死了:陷入集体通信
上面演示的进程之间完全不通信。真实的多控制器 JAX 程序必然涉及进程间通信(否则就没有使用多控制器的意义)。现在给脚本加上每轮循环都执行一次分布式jnp.sum的逻辑,完整脚本见 docs/_static/fault_tolerance/collectives.py:
def main(_: Sequence[str]) -> None: ... n = jax.device_count() jax.set_mesh(jax.make_mesh((n,), ("i",))) x = jax.device_put(jnp.arange(n), jax.P("i")) while True: print(jnp.sum(x)) time.sleep(1)上述代码中,四个进程创建一个跨四进程分片的数组x,然后执行分布式jnp.sum(即 AllReduce)。再次运行并失败第四个进程,你会看到:前三个进程不会崩溃,但会卡死。这是默认行为——如果某个进程在参与分布式计算(如jnp.sum)时失败,其余参与该计算的进程会永远卡住。
1.5 取消失败的集体通信
要避免卡死,可以取消带有失败参与者的集体通信。这需要再补上若干 XLA 标志与环境变量,完整脚本见 docs/_static/fault_tolerance/cancel_collectives.py:
import os os.environ['XLA_FLAGS'] = ' '.join([ '--xla_gpu_nccl_terminate_on_error=false', '--xla_gpu_nccl_async_execution=true', '--xla_gpu_nccl_blocking_communicators=false', ]) os.environ['XLA_PYTHON_CLIENT_ABORT_COLLECTIVES_ON_FAILURE'] = '1' os.environ['XLA_PYTHON_CLIENT_USE_TFRT_GPU_CLIENT'] = '1' ... def main(_: Sequence[str]) -> None: jax.config.update("jax_enable_recoverability", True) jax.distributed.initialize(...) # 注意:正常运行不要这样做,应使用下方正式 API live_devices。 from jax.experimental.multihost_utils import _live_devices _live_devices(jax._src.distributed.global_state.client, jax.devices()) n = jax.device_count() jax.set_mesh(jax.make_mesh((n,), ("i",))) x = jax.device_put(jnp.arange(n), jax.P("i")) while True: print(jnp.sum(x)) time.sleep(1)各开关的作用说明如下:
| 开关 | 作用 |
|---|---|
--xla_gpu_nccl_terminate_on_error=false | 同前,禁止 NCCL 出错触发自杀 |
--xla_gpu_nccl_async_execution=true | 让 NCCL 集体操作异步执行,从而可被取消 |
--xla_gpu_nccl_blocking_communicators=false | 不阻塞地管理 NCCL communicator |
XLA_PYTHON_CLIENT_ABORT_COLLECTIVES_ON_FAILURE=1 | 允许客户端中止失败参与者的集体通信 |
XLA_PYTHON_CLIENT_USE_TFRT_GPU_CLIENT=1 | 使用支持取消语义的 TFRT GPU 客户端 |
jax_enable_recoverability=True | 启用可恢复性语义 |
脚本中插入的一次对jax.experimental.multihost_utils._live_devices的调用是文档作者为了让脚本在正式 API 讲解前先跑起来而使用的临时 hack——正常编程不应这样用,应使用下文马上介绍的live_devices正式 API。
再次运行并失败第四个进程:前三个进程一开始卡在jnp.sum中,约十秒后该调用被取消并抛出类似下面的异常:
jaxlib._jax.XlaRuntimeError: FAILED_PRECONDITION: Task with incarnation id 3446767950926952685 is not connected注意错误中的incarnation id(化身标识):每个进程每次启动都会生成一个随机化身标识,它用于区分“同一个进程”与“重启后的进程”,是实现容错语义的关键概念。
1.6 核心 API:live_devices
进程死亡后,存活的进程需要知道谁死了、谁还活着。这正是 JAX 核心容错 APIlive_devices的职责:它是一个上下文管理器,接收一组设备作为参数,并返回其中仍然存活的那部分设备。完整用法见 docs/_static/fault_tolerance/live_devices.py:
from jax.experimental.multihost_utils import live_devices ... def main(_: Sequence[str]) -> None: jax.config.update("jax_enable_recoverability", True) jax.distributed.initialize(...) while True: try: with live_devices(jax.devices()) as devices: print(f'{devices=}') n = len(devices) jax.set_mesh(jax.make_mesh((n,), ("i",), devices=devices)) x = jax.device_put(jnp.arange(n), jax.P("i")) print(jnp.sum(x)) except Exception as e: print('FAIL:', e) else: print('PASS') time.sleep(1)核心代码用live_devices(jax.devices())获得存活设备集合devices,只在这些设备上分片数组x并执行jnp.sum。如果jnp.sum执行期间有进程失败,该集体通信会被取消并在其余存活设备上抛出异常(严格来说集体通信并不保证一定失败,这一点详见 1.8 的“原子性”讨论)。
重要提示:
jax.devices()永远返回全部设备——即使其中某些设备所在的进程已经失败。要获知哪些设备真正存活,必须使用jax.experimental.multihost_utils.live_devices。
实际运行中会发生什么?
- 失败第四个进程后,存活的三个进程会捕获
jnp.sum抛出的异常,进入 while 循环的下一轮迭代;这一轮里devices不再包含已死进程的设备,三个存活进程继续正确执行。 - 重新启动第四个进程后,它的设备又会重新出现在
live_devices返回的存活设备集合中,四个进程随即恢复正常协同运行。
从源码看,live_devices的实现位于 jax/experimental/multihost_utils.py。其底层辅助函数_live_devices的逻辑是:收集所提供设备的进程 id 集合,调用客户端的get_live_nodes获取当前存活节点(连同各自化身 id),再过滤出真正存活的设备子集。源码明确注明该 API 仍在积极开发中、尚不稳定。
live_devices表面上很简单——“传一组设备,返回存活的那组”,但正如分布式系统中的许多事情一样,其中布满微妙的细节。下面两节解释它的屏障(barrier)语义与原子性(atomicity)性质。
1.7 屏障语义:所有进程必须看到同一份存活列表
多控制器 JAX 程序要求每个进程步调一致地执行:各进程应当以相同顺序执行相同指令,否则几乎必然导致死锁、崩溃或异常行为。
考虑一个具体场景:进程 1、2 调用live_devices,随后进程 4 失败,然后进程 3 才调用live_devices。此时进程 1、2 可能认为进程 4 还活着,而进程 3 认为它已死——各进程对“谁活着”的认知不一致,就会开始分叉(divergence)。
为避免这种情况,live_devices保证向每个进程返回相同的存活设备集合。其实现手段是一次屏障:live_devices(devices)调用会阻塞,直到每一个承载devices中设备且仍存活的进程都调用了live_devices。当所有存活进程都进入该屏障后,live_devices向每个进程返回同一份存活设备集合。
重要:
live_devices借助屏障保证它总是向每个存活进程返回相同的存活设备集合。
由于live_devices实现了屏障,使用不当就会死锁。官方建议:一个程序里只保留一个with live_devices代码块。多次调用live_devices难以推理且可能死锁。
1.8 原子性:要么全体成功,要么全体失败
所谓分布式计算的原子性,是指每个参与者对操作“成功还是失败”达成一致。在 1.6 的脚本中,进程在执行jnp.sum期间失败时,jnp.sum会在其余存活进程上中止并抛出异常——那么jnp.sum是原子的吗?
不是。当某个进程在集体操作执行期间失败时,剩余进程可能取消操作并抛异常,也可能成功完成操作。JAX 中的集体操作本身没有任何原子性保证。
如果集体操作不原子,多控制器进程就可能分叉:例如训练机器学习模型时某个进程失败,部分进程检测到失败并把模型回滚到检查点,另一部分进程却认为该步成功了继续训练。
为了解决这个问题,live_devices尽管集体操作不原子,仍提供自己的原子性保证:with live_devices块内的代码要么在所有进程上成功完成,要么在所有进程上抛出异常。具体来说,对下面的代码,要么所有进程执行分支 A,要么所有进程执行分支 B,绝不可能出现一部分进程执行 A、另一部分执行 B:
try: with live_devices(jax.devices()) as devices: ... # 主体代码 except Exception as e: ... # 分支 A else: ... # 分支 B注意:如果代码块因为集体通信失败(进程崩溃)之外的非确定性原因抛出异常(例如某个进程自身内存耗尽),该异常不会被传播给其他进程,此时原子性不被保证。
异步派发对原子性的影响:JAX 使用异步派发机制,jnp.sum这类操作不会阻塞到计算完成,而是返回充当 future 的jax.Array。这种异步性可能以意外方式与live_devices交互。例如:
x = ... y = ... try: with live_devices(jax.devices()) as devices: y = jnp.sum(x) except Exception as e: ... # 分支 A else: ... # 分支 B print(y)设想with live_devices块在所有进程上都成功执行(都走分支 B)。这只能保证每个进程都成功创建了一个 future 并赋给y;jnp.sum的实际计算可能被推迟到代码块之外。于是可能出现:部分进程成功完成jnp.sum并打印y的值,而另一些进程没能完成jnp.sum、在尝试打印y时抛出异常。
解决办法:在with live_devices块内使用jax.block_until_ready强制计算完成。如下代码能保证“要么所有进程成功执行jnp.sum,要么所有进程抛出异常”:
x = ... y = ... try: with live_devices(jax.devices()) as devices: y = jax.block_until_ready(jnp.sum(x)) except Exception as e: ... # 分支 A else: ... # 分支 B print(y)第二部分:实战示例
需要强调的是:live_devices本身并不“使程序容错”,它只是供你自行实现容错的底层工具,具体实现方式因应用形态而异。下面的示例用于演示而非规定,容错还有其他许多实现思路。
2.1 示例一:容错的数据并行训练
本示例在四个进程上以数据并行方式训练一个单参数线性模型y = α·x。示例刻意极度简化(你当然不会在四台机器上训练单参数模型),目的是把注意力集中在容错机制上。
为什么数据并行天然适合容错?因为每个进程都拥有一份完整的模型权重副本,进程失败后可以忽略它并继续训练。此示例可容忍任意数量(进程 0 除外)的进程失败,但假设失败的进程不会恢复——下一个示例将展示如何处理进程恢复。
完整脚本见 docs/_static/fault_tolerance/data_parallelism.py。脚本由以下几部分构成:
(1)开头的开关与参数定义(对应源码第 15–33 行):设置前文 1.5 节的全部 XLA 标志与环境变量,定义--i、--n两个命令行参数。
(2)两个“分片元数据”辅助函数:它们并不真正搬移数据,只是为既有数据创建带复制/分片 sharding 语义的进程级jax.Array视图:
def replicated(x: jax.Array, devices: list[jax.Device]): """返回在给定设备上复制的 x;不真正搬移数据。""" n = len(devices) mesh = jax.make_mesh((n, ), ("i", ), devices=devices) spec = jax.sharding.PartitionSpec(None) # 复制 = 无分片维度 sharding = jax.sharding.NamedSharding(mesh, spec) shards = [ jax.device_put(x.addressable_shards[0].data, d) for d in devices if d.process_index == jax.process_index() ] return jax.make_array_from_single_device_arrays(x.shape, sharding, shards) def sharded(x: jax.Array, devices: list[jax.Device]): """返回在给定设备上分片的 x;x 应与全局数组同形状。""" n = len(devices) mesh = jax.make_mesh((n, ), ("i", ), devices=devices) spec = jax.sharding.PartitionSpec("i") # 按首个轴分片 sharding = jax.sharding.NamedSharding(mesh, spec) m = sharding.addressable_devices_indices_map(x.shape) shards = [jax.device_put(x[m[d]], d) for d in jax.local_devices()] return jax.make_array_from_single_device_arrays(x.shape, sharding, shards)(3)主训练循环(对应源码第 99–125 行):
step = 0 while True: try: with live_devices(jax.devices()) as devices: print(f'=== Running step {step} with live devices = {devices} ===') # 复制模型权重。 weights = replicated(weights, devices) # 分片当前 batch。 batch_size = device_batch_size * len(devices) start = (step * batch_size) % len(X) stop = start + batch_size X_batch = sharded(X[start:stop], devices) Y_batch = sharded(Y[start:stop], devices) # 计算梯度并更新权重。 l, grad = loss_and_grad(weights, X_batch, Y_batch) new_weights = jax.block_until_ready(weights - learning_rate * grad) except Exception as e: print(f'Step {step} failed: {e}') else: print(f'Step {step} succeeded: loss = {l}') step += 1 weights = new_weights time.sleep(1)逐行解读这个循环的设计意图:
- 每一轮迭代先调用
live_devices获取当前存活设备; - 将权重复制到这些设备上、把训练数据分片到这些设备上(注意这只是创建带正确 sharding 元数据的 JAX 数组,不在设备间搬移数据);
- 调用
loss_and_grad(由jax.jit(jax.value_and_grad(loss))生成)计算梯度,再得到新权重。刻意把新权重赋给new_weights而非直接覆盖weights,是为了防止训练步失败时污染当前权重;同时调用jax.block_until_ready,确保退出live_devices块时每个进程都已真正算出新权重; - 若训练步执行期间没有进程失败,走
else分支:step递增、weights更新为new_weights。否则抛出异常走except分支:不更新step和weights,下一轮用新的存活设备集合重试这一步。
2.2 示例二:支持进程恢复的数据并行训练
现在扩展上面的示例,允许失败进程恢复。恢复后的进程需要拿到当前step与模型权重。由于进程 0 永不失败(回忆 1.3 节:进程 0 失败全体都会失败),由进程 0 向恢复中的进程发送当前 step 和权重。完整脚本见 docs/_static/fault_tolerance/data_parallelism_with_recovery.py。
(1)基于shard_map的点对点send/recv(源码第 69–90 行):发送方调用send,接收方调用recv。二者通过jax.lax.psum(AllReduce 求和)+ 复制 sharding 来传输数据——发送方持真值、接收方持全零占位,psum 恰好把值“送”到接收端:
def send(x: jax.Array, from_device: jax.Device, to_device: jax.Device): """将 x 从一个设备发送到另一个设备。""" devices = [from_device, to_device] psum = lambda x: jax.lax.psum(x, "i") mesh = jax.make_mesh((2, ), ("i", ), devices=devices) spec = jax.sharding.PartitionSpec(None) x = replicated(x, [from_device, to_device]) shard_map.shard_map(psum, mesh=mesh, in_specs=spec, out_specs=spec)(x) def recv(x: jax.Array, from_device: jax.Device, to_device: jax.Device): """接收来自匹配 send 的 x。""" to_device = jax.local_devices()[0] devices = [from_device, to_device] psum = lambda x: jax.lax.psum(x, "i") mesh = jax.make_mesh((2, ), ("i", ), devices=devices) spec = jax.sharding.PartitionSpec(None) x = jnp.zeros_like(x) x = replicated(x, [from_device, to_device]) return shard_map.shard_map(psum, mesh=mesh, in_specs=spec, out_specs=spec)(x)(2)allgather辅助函数(源码第 93–100 行):对单个 float 跨一组设备执行 AllGather,返回每个设备的数值列表:
def allgather(x: float, devices: list[jax.Device]) -> list[float]: """在给定设备上对 x 执行 AllGather。""" n = len(devices) mesh = jax.make_mesh((n, ), ("i", ), devices=devices) spec = jax.sharding.PartitionSpec('i') p = lambda x: jax.lax.all_gather(x, "i", tiled=True) f = jax.shard_map(p, mesh=mesh, in_specs=spec, out_specs=spec) return jax.block_until_ready(f(np.array([x] * len(devices)))).addressable_shards[0].data(3)修改后的训练循环(源码第 135–178 行):恢复是两步过程——先检测哪些进程在恢复,再由进程 0 把 step 和权重发给恢复进程:
step = 0 while True: try: with live_devices(jax.devices()) as devices: # 第 1 步:检测恢复中的设备。 # 对全部存活设备的 step 做 AllGather;恢复进程的 step 为 0, # 而进程 0 的 step 为正数,故 step 不等于进程 0 者即为恢复中。 print('all gathering steps...') steps = allgather(step, devices) print(f'{steps=}') recovering = [d for d, s in zip(devices, steps) if s != steps[0]] # 第 2 步:进程 0 向恢复中的设备发送 step 与权重。 for d in recovering: if jax.process_index() == 0: print('sending...') send(weights, jax.devices()[0], d) send(jnp.array([step]), jax.devices()[0], d) elif d.process_index == jax.process_index(): print('receiving...') weights = recv(weights, jax.devices()[0], d) step = recv(jnp.array([step]), jax.devices()[0], d)[0] # 之后与示例一相同:复制权重、分片 batch、计算并 block_until_ready。 ... except Exception as e: ... else: step += 1 weights = new_weights这里值得指出一个与incarnation id相关的深层要点:仅仅比较 step 并不能区分“失败后重启的进程”与“从未失败的进程”,如果恢复发生在两次调用之间,还可能引发匹配错乱。live_devices的正式实现通过跟踪进程化身 id 来严格处理这类情况(详见第三部分)。
第三部分:实现细节
如果只关心“如何编写容错程序”,前两部分已经足够。第三部分深入剖析多控制器 JAX 的架构与live_devices的语义及实现,帮助你在极端场景下也能理解 API 的行为。
3.1 协调服务:控制面、心跳与命运共享的引擎
启动多控制器 JAX 程序时,第一个进程(进程 0)会运行一个独立的 RPC 服务器,即协调服务(coordination service);同时所有进程(包括进程 0 自己)都创建到该服务的 RPC 客户端。具体来说,jax.distributed.initialize的coordinator_address参数就是协调服务的地址:它告诉进程 0 在哪个地址上启动服务器,也告诉所有进程去连接哪个地址。
协调服务实现了多控制器 JAX 的控制面(control plane)。例如:
- 它可以跨所有进程执行分布式屏障;
- 它实现了一个键值存储,进程可用来交换少量元数据。
需要特别注意的是,数据面(data plane)——即所有针对程序数据的集体操作——直接在进程之间完成,不经过协调服务。
协调服务最重要的功能之一是健康检查:每个进程周期性地向协调服务发送心跳;进程失败便停止发送心跳;若协调服务较长时间未收到某进程的心跳,就认定该进程已失败。默认情况下,协调服务一旦检测到进程失败,会向所有其他进程发送消息要求它们自我终止——这就是多控制器 JAX 程序“命运共享”的根源,也是它完全不具容错性的原因。
由此可归纳出开启容错必须做的两件事:(1) 移除命运共享,允许进程在同伴死亡后继续执行——通过
jax_enable_recoverability配置项开启;(2) 提供一种 API 让进程获知谁存活、谁已死——即live_devicesAPI。
实现live_devices的技术深度远超表象。官方文档采用逐步演进的教学路径:先提出一个更简单的live_processesAPI,再逐步修正缺陷,最终抵达live_devices。
3.2 从live_processes到live_devices:为什么朴素实现是错的
假设设计一个新 APIjax.live_processes(),期望它返回所有当前存活进程的集合。一个朴素的实现是:进程向协调服务发 RPC 请求,协调服务依据心跳信息直接回复它认为存活的一组进程。这样做正确吗?
不正确。多控制器 JAX 任务要求所有进程以相同顺序执行相同指令。一旦各进程因为对“谁存活”判断不一致而走上不同的代码路径,任务行为就会失控——大概率崩溃、挂起或产生垃圾值,而且极难排查。请看一个具体场景:三个进程的任务中,进程 0 和 1 几乎同时调用live_processes,恰在此刻进程 2 失败。协调服务可能告诉进程 0“所有进程都存活”,却告诉进程 1“只有进程 0 和 1 存活”。一旦进程对存活集合产生分歧,它们几乎必然分叉。
修补方案:给live_processes加上屏障语义。协调服务收到live_processes()请求后不立即回复,而是等每一个存活进程都调用了live_processes()之后,再把存活进程集合返回给所有进程。因为返回给所有进程的是同一份集合,各进程便不会分叉。
3.3 形式语义:基于线性化的一致性定义
分布式系统极其复杂:机器可在任意时刻失效,网络消息可能丢失、延迟、乱序。官方文档引入一套基于**线性化(linearizability)**的形式语义来界定live_processes的正确行为。系统被建模为若干进程,每个进程串行执行若干事件,共四种事件类型:
- 进程启动:假定启动后即连接协调服务,协调服务知晓其已启动;
- 进程失败:与启动不同,协调服务可能不会立即感知失败;
- 进程发送
live_processes请求给协调服务; - 进程接收来自协调服务的回复。
有效性定义:若live_processes返回一组存活进程 P,则必须存在某一瞬间,P 中每个进程都在live_processes屏障中、而所有其他进程都已死亡。实现live_processes的正确性标准就是:只允许有效执行发生。
由此可以得到若干看似反直觉却正确的推论:
- 返回 P 不代表 P 中进程此刻都活着、P 外进程此刻都死了,只表示曾存在某一时刻如此。
- 进程 1 调用了
live_processes却在收到回复前死亡:只要存在进程 0 在屏障内、进程 1 已死的时刻,执行依然有效(其请求可能已在网络中被丢弃)。 - 进程 0 收到回复
0,1时进程 1 刚死:仍然有效——协调服务可能已收到两个请求并回复,只是在回复传输途中进程 1 才失败。
修正失败时刻:分布式系统无法以 100% 精度探测失败。协调服务只是“一段时间收不到心跳就认定死亡”,它无法确定进程到底死于何时、甚至是否真死(也许只是网络分区)。因此形式语义允许把一次失败在时间上向前或向后移动(但不能越过同一进程的其他事件)——直观地说,可以把失败从“实际发生的时刻”移到“协调服务认为它发生的时刻”。例如进程 1 实际已死但协调服务还当它活着,把它的失败时刻向后推迟,就能构造出“两进程同时在屏障内”的合法瞬间;反之,进程 1 其实活着但被网络分区隔绝,协调服务判定其死亡,把失败时刻向前移动即可解释返回集合{0}的有效性。但失败不能越过该进程自己的其他事件,否则执行无效。
3.4 原子性如何实现:两次live_processes检查的思考
有了live_processes,尝试编写容错代码。下面这段代码“看起来”正确,实际含有一个微妙的 bug:
step = 0 while True: procs = jax.live_processes() # 获取存活进程 devices = [d for d in jax.devices() if d.process_index in procs] mesh = jax.make_mesh((len(devices),), ("i",), devices=devices) spec = jax.sharding.PartitionSpec("i") sharding = jax.sharding.NamedSharding(mesh, spec) x = jax.make_array_from_process_local_data(sharding, np.ones(1)) try: print(jnp.sum(x)) except: pass # jnp.sum 失败 else: step += 1 # jnp.sum 成功Bug 根源:若jnp.sum正跨进程集合 P 执行时 P 中某进程失败,jnp.sum在各进程上的表现可能不同——部分进程看到正确结果、部分抛出异常、还有部分得到错误结果。于是进程可能分叉:有的递增step,有的没有。在玩具代码里这种分叉无害,但在真实程序中会导致崩溃、死锁或垃圾输出。例如数据并行训练若分叉,部分进程把权重回滚到旧检查点、其余进程继续训练,就会产生无人认同的“弗兰肯模型”。
正确思路:想要“要么全体成功要么全体失败”的原子性,可以在代码块前后各调用一次live_processes:如果块前存活的进程集合与块后一致,说明代码块在所有存活进程上成功执行;只要有进程死亡,所有剩余进程就能一致认定代码块执行失败。但把它写对还有几个细节要处理:
- 代码块本身抛异常怎么办?需要捕获异常、仍完成第二次
live_processes、再重新抛出。 - 进程若在第一次调用后失败、第二次调用前又恢复了呢?前后集合相同但代码块实际失败过。解决方案:进程每次启动都会生成随机化身 id,除检查集合不变外,还要检查化身 id 未变。
- 恢复进程的第一次
live_processes与另一进程的第二次调用匹配上导致死锁怎么办?答案是只在单一程序点调用live_processes,让一次调用同时承担两个职责:既校验自上次调用以来进程集合未变,又生成本次原子代码块应使用的存活进程集合。
live_devices正是把这些细节全部封装抽象后的产物:它是一个上下文管理器,保证代码块原子执行。devices是所有存活进程上的设备列表;块 A 在这些进程上原子执行——要么每个进程都看到代码抛异常(分支 B),要么每个进程都看到代码成功(分支 C):
try: with live_devices() as devices: pass # A except Exception as e: pass # B else: pass # C3.5 取消集体通信的底层原理:NCCL 与通信器缓存
前面 1.4、1.5 节提到:集体通信的参与者失败时,其余进程会永久卡死,需要显式取消。需要明确的能力边界是:live_devicesAPI 在所有 JAX 后端(CPU、GPU、TPU)都受支持,但取消集体通信只有 GPU 后端支持,原因在于其实现依赖 NVIDIA 的集体通信库NCCL。
底层机制如下:
- GPU 后端用 NCCL 实现集体通信。一组进程要执行集体操作时,先组建一个NCCL communicator,之后可反复用该通信器执行集体操作。
- 创建 communicator 很昂贵(需要网络通信),因此 JAX 后端以参与进程集合及其化身 id 为键缓存 communicator。
- 在内部,JAX 客户端持续轮询协调服务以获取每个进程的当前状态。一旦客户端发现某进程死亡、或携带新化身 id 重启,就中止缓存键中包含该失败化身 id 的所有 communicator——这正是
jnp.sum能及时抛出FAILED_PRECONDITION: Task with incarnation id ... is not connected异常、而非永久挂起的原因。
结语与进一步阅读
容错分布式编程的本质困难在于:进程间“谁活着”无法被精确感知、集体操作本身不原子、异步派发会推迟计算的可见时机。JAX 给出的答案是live_devices这一“屏障 + 原子性”原语,配以心跳健康检查、jax_enable_recoverability、GPU 上的 NCCL 集体取消机制,构成一套自洽的容错编程模型。掌握它之后,你既能写出数据并行下忽略故障进程的训练循环,也能实现进程重启后从协调者恢复状态的高级方案。
想要继续深入,可以:
- 通读本文依据的官方指南 docs/501/fault-tolerance.rst,其中包含交互式可视化示例;
- 对照阅读全部可直接运行的示例脚本目录 docs/_static/fault_tolerance/;
- 阅读
live_devices的实现与完整文档字符串 jax/experimental/multihost_utils.py,以及jax_enable_recoverability等配置项的定义处 jax/_src/distributed.py; - 复习多控制器 JAX 的常规使用方式 docs/501/multiprocess.md,以及
jax.distributed.initialize的完整参数说明 jax.distributed 参考文档; - 若涉及大量分布式调试验证,可参考仓库中的多进程测试目录 tests/multiprocess/ 了解此类程序通常如何被组织与验证。
最后再次提醒:live_devices是有意暴露的底层原语,官方建议一个程序只保留一个with live_devices块,并在块内对关键结果调用jax.block_until_ready;当前的容错支持仍属实验特性且主要面向 GPU,投入使用前请结合自身负载做好充分压测与故障演练。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考