ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

JAX 多控制器分布式容错编程实战:`live_devices`、心跳检测与集体通信取消

JAX 多控制器分布式容错编程实战:`live_devices`、心跳检测与集体通信取消 JAX 多控制器分布式容错编程实战live_devices、心跳检测与集体通信取消【免费下载链接】jaxComposable transformations of PythonNumPy 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_addresslocalhost:9000, num_processes_NUM_PROCESSES.value, process_id_PROCESS_ID.value, local_device_ids[_PROCESS_ID.value], heartbeat_timeout_seconds10, ) while True: print(time.time()) time.sleep(1)脚本先调用jax.distributed.initialize初始化多控制器 JAX然后进入死循环每秒打印一次当前时间。local_device_ids参数确保每个进程只被分配四块 GPU 中的一块heartbeat_timeout_seconds稍后解释。在一台拥有四块 GPU 的虚拟机上用四个终端分别启动四个进程python example.py --i0 --n4 # 终端 1 python example.py --i1 --n4 # 终端 2 python example.py --i2 --n4 # 终端 3 python example.py --i3 --n4 # 终端 4此时四个进程每秒都会打印时间。现在杀掉第四个进程pkill -9 -f python example.py --i3 --n4大约十秒后其余进程会全部终止并打印类似下面的错误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.pyimport os os.environ[XLA_FLAGS] --xla_gpu_nccl_terminate_on_errorfalse ... 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_errorfalseXLA 标志写入XLA_FLAGS禁止 GPU 集体通信在出错时触发进程自杀jax_enable_recoverabilityJAX 配置项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.pydef 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.pyimport os os.environ[XLA_FLAGS] .join([ --xla_gpu_nccl_terminate_on_errorfalse, --xla_gpu_nccl_async_executiontrue, --xla_gpu_nccl_blocking_communicatorsfalse, ]) 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_errorfalse同前禁止 NCCL 出错触发自杀--xla_gpu_nccl_async_executiontrue让 NCCL 集体操作异步执行从而可被取消--xla_gpu_nccl_blocking_communicatorsfalse不阻塞地管理 NCCL communicatorXLA_PYTHON_CLIENT_ABORT_COLLECTIVES_ON_FAILURE1允许客户端中止失败参与者的集体通信XLA_PYTHON_CLIENT_USE_TFRT_GPU_CLIENT1使用支持取消语义的 TFRT GPU 客户端jax_enable_recoverabilityTrue启用可恢复性语义脚本中插入的一次对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 核心 APIlive_devices进程死亡后存活的进程需要知道谁死了、谁还活着。这正是 JAX 核心容错 APIlive_devices的职责它是一个上下文管理器接收一组设备作为参数并返回其中仍然存活的那部分设备。完整用法见 docs/_static/fault_tolerance/live_devices.pyfrom 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,), devicesdevices)) 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、另一部分执行 Btry: 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 并赋给yjnp.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, ), devicesdevices) 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]): 返回在给定设备上分片的 xx 应与全局数组同形状。 n len(devices) mesh jax.make_mesh((n, ), (i, ), devicesdevices) 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(fStep {step} failed: {e}) else: print(fStep {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.psumAllReduce 求和 复制 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, ), devicesdevices) spec jax.sharding.PartitionSpec(None) x replicated(x, [from_device, to_device]) shard_map.shard_map(psum, meshmesh, in_specsspec, out_specsspec)(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, ), devicesdevices) spec jax.sharding.PartitionSpec(None) x jnp.zeros_like(x) x replicated(x, [from_device, to_device]) return shard_map.shard_map(psum, meshmesh, in_specsspec, out_specsspec)(x)2allgather辅助函数源码第 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, ), devicesdevices) spec jax.sharding.PartitionSpec(i) p lambda x: jax.lax.all_gather(x, i, tiledTrue) f jax.shard_map(p, meshmesh, in_specsspec, out_specsspec) return jax.block_until_ready(f(np.array([x] * len(devices)))).addressable_shards[0].data3修改后的训练循环源码第 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尝试编写容错代码。下面这段代码“看起来”正确实际含有一个微妙的 bugstep 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,), devicesdevices) 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要么每个进程都看到代码成功分支 Ctry: 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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表