
JAX 501 系统主题深度指南多进程编排、故障容错与跨进程产物管理【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 官方文档按学习深度划分为 101/201/301/401/501 几个层级其中「501: systems topics」回答的是同一个核心问题当一个 JAX 任务不再只由一个 Python 进程驱动少数设备而是变成更大的系统的一部分时该如何运行它。本文以 docs/501/index.rst 为主线完整梳理这一层级的七大主题多控制器计算、分布式数据加载、分布式容错、网络安全、导出与序列化、持久化编译缓存、传输守卫并对文档「Smaller notes」中直接收录的并发线程约束、多进程协调工具与产物兼容性策略做源码级展开。读完你将能够判断你的多进程 JAX 程序何时共享命运share fate、如何借助jax.live_devices让训练在机器故障后存活下来以及如何让编译产物安全地跨越进程与版本长期复用。501 讲的是什么从单进程驱动设备到把 JAX 嵌入更大的系统回顾前面的学习路径101/201/301 等所有教程都隐含一个前提一个 Python 进程驱动若干个设备。而 501 的内容边界恰恰相反——索引页开宗明义这些页面讨论的是把 JAX 作为更大系统的一部分来运行多个进程与多台主机、长生命周期且可重启的任务、比创建它的进程活得更久的产物以及让这一切保持又快又安全的基础设施。也就是说本层级的关注点有四类规模scale多进程、多主机、多控制器生命周期longevity任务可以长跑、可以被重启、可以在故障后恢复产物artifacts序列化的编译产物比创建它的进程活得更久跨版本、跨平台复用安全与正确性保障safety控制面加密、无意的数据搬运被拦截、并发线程的调度被约束。这套体系的架构直觉可以用下面这张多主机 TPU pod 的示意图来理解每一台 host 上运行一个 JAX Python 进程host 通过 PCI 挂载多块 TPU 芯片芯片之间由高速互连ICI相连GPU、CPU 等其他平台上原理相同——这正是索引页所有主题共用的运行模型。主题地图501 的七个条目与八个入口页面索引页用一份带注释的编号清单概括了本层级全部内容。结合仓库内实际文件可将七个条目整理如下第八个文件是第 5 条内并列导出的 shape-polymorphism#主题索引页定位一句话仓库入口文档1多控制器 JAXmultiprocess每台主机一个进程、跨进程的 mesh 与 array、基于jax.device_put的运行时级流水并行、从各进程本地数据构造全局 arraydocs/501/multiprocess.md2分布式数据加载data-loading把每个 batch 的各个分片送到正确的主机覆盖数据并行与模型并行两种负载docs/501/data-loading.md3分布式容错fault-tolerance用jax.live_devices在机器故障后存活barrier 语义、原子性、带完整训练示例的恢复流程docs/501/fault-tolerance.rst4网络安全security为协调服务与 TPU runtime 内部服务启用 mTLS并说明哪些连接仍不受保护docs/501/security.md5导出与序列化export把 staged-out 的计算导出、序列化供后续或跨平台执行docs/501/export.md6符号形状导出shape-polymorphism用符号形状导出让一个产物服务多种输入尺寸docs/501/shape-polymorphism.md7持久化编译缓存compilation-cache进程重启后跳过重复编译并在多节点间共享编译产物docs/501/compilation-cache.md8传输守卫transfer-guard记录或禁止非预期的 host-device 数据传输docs/501/transfer-guard.rst以下按主题逐个展开说明其解决的核心问题与关键 API。主题 1多控制器 JAX——每个进程就是一个控制器索引页指出多控制器 JAX 的要点是在每个 host 上运行一个或多个Python 进程这些进程有时被称作控制器controller。整个集群通过各进程调用jax.distributed.initialize联通后jax.Array可以跨进程分布只要每个进程对同一份数组施加相同顺序的相同运算编程体验就接近一台插满了所有设备的巨型机器。主题 1 覆盖三个实操层次均可在 docs/501/multiprocess.md 中找到完整示例跨进程 mesh 与全局数组直接用jax.make_mesh((jax.device_count(),), (a,))构造包含所有进程设备的 mesh再以NamedSharding配合jax.device_put把数据分片到全局设备运行时级流水并行jax.device_put可以把数据从一组设备搬到另一组设备且由于 JAX 的异步分发机制不同设备队列上的jit函数与device_put只要输入就绪即可并行执行——这正是做 microbatch pipeline parallelism 的底层能力从进程本地数据构造全局数组当在每个进程内分别持有全局数据的某一片时可调用jax.device_put、jax.make_array_from_process_local_data或最通用的jax.make_array_from_single_device_arrays注意后两者不做数据搬运只是把已就位的各设备分片组装成全局jax.Array视图。与之配套的底层实现位于 jax/_src/distributed.py。其中State.initialize的参数coordinator_address、num_processes、process_id、local_device_ids、heartbeat_timeout_seconds等见 jax/_src/distributed.py就是索引页反复提到的故障检测与集群发现机制的入口。主题 2分布式数据加载——让每个设备拿到本该属于它的分片当训练数据分散在多进程/多主机环境中主题 2 要解决的是数据分片到设备的归属问题。索引页特别强调一个易错点如果数据分片放错了设备计算本身不会报错计算无从知道正确的数据应该是什么但最终结果往往是错的。这类问题更广义地适用于任何从非 JAX 数据源构造跨进程jax.Array的场景例如从 checkpoint 加载模型权重、加载大尺寸空间分片图像。docs/501/data-loading.md 给出的统一思路是先决定Sharding分片方案——它描述了每个全局设备需要哪一块全局数据用Sharding.addressable_devices()查出当前进程需要为哪些设备准备数据再在四种高层策略中选一种落地每个进程都加载全局数据 / 逐设备数据管线 / 按进程合并的数据管线 / 任意方便方式加载后放进计算内再做 reshard。文档中还指出分布式数据加载通常比单个进程加载全量再经 RPC 分发每个进程各自加载全量更高效但也更复杂——后两者实现简单代价是可能阻塞训练循环并占用额外网络带宽。主题 3分布式容错——用live_devices让训练在机器故障后继续多控制器 JAX 在默认情况下不具备容错能力任何一台机器故障全体机器都会随之退出进程间共享命运。索引页引出的容错方案集中在jax.live_devices与配套的 barrier/原子性语义上。本主题的完整文章 docs/501/fault-tolerance.rst 篇幅最长仓库还在 docs/_static/fault_tolerance 下随文附带了 7 个可直接运行的示例脚本构成一条渐进式的故障演练路径while_loop.py默认行为基线。四进程每进程独占一块 GPUpkill -9杀掉 4 号进程后其余进程约在heartbeat_timeout_seconds10后全部退出——这就是fate sharingdont_fail.py加入XLA_FLAGS--xla_gpu_nccl_terminate_on_errorfalse并设置配置项jax_enable_recoverabilityTrue进程不再互相拖垮但协调服务运行在 0 号进程上0 号进程一旦死亡全体仍会失败collectives.py进程间开始执行集体通信jnp.sum后故障进程会让剩余进程永久卡死在集体操作里cancel_collectives.py追加--xla_gpu_nccl_async_executiontrue、--xla_gpu_nccl_blocking_communicatorsfalse等 XLA flags并设置XLA_PYTHON_CLIENT_ABORT_COLLECTIVES_ON_FAILURE1使带失败参与者的集体操作可被取消、抛出异常而不是挂死live_devices.py引入核心 API——with live_devices(jax.devices()) as devices:返回当前存活的设备子集故障后的下一轮循环只在这组设备上重建 mesh 并继续执行被杀进程重启后其设备会重新出现在存活集合中。在配置项与配置源码层面可以验证jax_enable_recoverability在 jax/_src/distributed.py 中定义默认值为FalseAllows a multi-controller JAX job to continue running, even after some tasks have failed这正是容错能力需要显式打开的原因。live_devices的语义有两处必须注意详见 docs/501/fault-tolerance.rstBarrier 语义live_devices(devices)会阻塞直到devices中设备所在的所有存活进程都调用了live_devices随后向每个进程返回完全相同的存活设备集合避免各进程对谁活着产生分歧而走向发散原子性集体操作本身不保证原子——故障可能让某些进程成功完成而另一些取消并抛异常。live_devices通过 barrier 补上这一层保证with live_devices(...)块体要么在全体进程上都正常完成要么在全体进程上都抛异常。但需要注意异步分发的干扰块内只是创建了 future并不代表计算已完成因此应把jax.block_until_ready(jnp.sum(x))放在块内才能把计算真正完成也纳入原子性保证。文章第三部分还剖析了实现机制0 号进程运行一个独立的 RPC 服务协调服务 / coordination service构成多控制器 JAX 的控制面分布式 barrier、key-value 元数据交换、健康检查进程周期性发送心跳协调服务据此判定死亡并广播fate sharing通知而真正搬运程序数据的数据面各集体操作则在进程间直连不经过协调服务。主题 4网络安全——用 mTLS 保护协调服务连接多进程协调依赖网络而 docs/501/security.md 明确指出一个默认事实jax.distributed.initialize启动的 gRPC 协调服务连接默认既不加密也不认证。若攻击者能触达coordinator_address、伪装成协调者或截获流量就可能观察或篡改集群拉起过程、key-value 存储等功能。要直接加固可向jax.distributed.initialize传入三个 mTLS 参数或等价环境变量三者缺一不可mtls_cert_file或JAX_MTLS_CERT_FILE本进程从 CA 签发的证书mtls_key_file或JAX_MTLS_KEY_FILE用于向对端证明自己持有证书对应私钥mtls_ca_file或JAX_MTLS_CA_FILE用于校验对端证书的 CA 证书。可选地设置verify_secure_credentialsTrue或环境变量JAX_DISTRIBUTED_VERIFY_SECURE_CREDENTIALS可在未提供任何 mTLS 参数时直接崩溃从而杜绝意外使用不安全连接提供mtls_peer_uri_prefix或JAX_MTLS_PEER_URI_PREFIX可改变对端身份校验方式默认客户端按 CA 签名 主机名校验服务端身份。这些参数最终汇集到 jax/_src/distributed.py 的_get_mtls_kwargs回退读取对应配置项mtls_cert_file等且低于 jaxlib 0.11.2 时会直接拒绝启用 mTLS。主题 5导出与序列化——让计算产物活过创建它的进程201 层的 AOT API 产出的对象通常只在当前进程内用于调试或编译执行当需要把 staged-out 的计算序列化到另一个进程、另一台机器甚至留档复现时就要用主题 5 的jax.export。索引页给出的动机很明确序列化产物比创建它的进程活得更久因此必须回答跨版本还能不能用这类兼容性问题。仓库文档 docs/501/export.md 给出了最小闭环import jax from jax import export def f(x): return 2 * x * x exported: export.Exported export.export(jax.jit(f))( jax.ShapeDtypeStruct((), jax.numpy.float32)) serialized: bytearray exported.serialize() # 序列化为字节串 rehydrated: export.Exported export.deserialize(serialized) # 重新水合子主题shape-polymorphismdocs/501/shape-polymorphism.md则允许导出带符号形状symbolic shapes的产物使一个编译产物能够服务多种输入尺寸而不必为每种形状各导出一次。主题 6持久化编译缓存——跨进程重启与跨节点复用编译产物重复的 JIT 编译是大规模训练与推理中的常见开销。主题 6 的 docs/501/compilation-cache.md 介绍可选的磁盘编译缓存开启后 JAX 会把编译好的程序写盘从而在重复运行相同/相似任务时跳过编译。最小启用方式须在首次编译前设置import jax import jax.numpy as jnp jax.config.update(jax_compilation_cache_dir, /tmp/jax_cache) jax.config.update(jax_persistent_cache_min_entry_size_bytes, -1) jax.config.update(jax_persistent_cache_min_compile_time_secs, 0) jax.jit def f(x): return x 1 f(jnp.zeros((2, 2)))缓存目录也可在 shell 中用export JAX_COMPILATION_CACHE_DIR/tmp/jax_cache指定。文档特别给出安全警告缓存被当作可信内容处理——把缓存放到他人可写的共享目录等价于允许任何能写缓存目录的人在你这台机器上执行任意代码。缓存键已内含 jaxlib 版本因此跨版本切换表现为缓存未命中cache miss而不会误用旧产物详见下文兼容性小节。主题 7传输守卫——记录或禁止非预期的 host-device 搬运JAX 会在类型转换、输入分片等过程中隐式地在 host 与设备之间、设备与设备之间搬运数据。主题 7 的 docs/501/transfer-guard.rst 把传输分成两类并允许用户分级管控显式传输jax.device_put*()与jax.device_get()调用隐式传输除此之外的搬运例如打印DeviceArray。守卫等级guard level包括allow默认静默放行一切、log记录并放行隐式传输、disallow禁止隐式传输、log_explicit、disallow_explicit。配置途径与其它 JAX 选项一致命令行--jax_transfer_guardGUARD_LEVEL、jax.config.update(jax_transfer_guard, GUARD_LEVEL)或with jax.transfer_guard(GUARD_LEVEL): ...线程局部上下文。还可以按传输方向细分使用带方向后缀的选项host_to_device/device_to_device/device_to_host例如--jax_transfer_guard_host_to_devicedisallow。Small notes索引页直接收录的三个系统要点索引页在主题地图之外把几个小到足以就地讲清的系统议题直接收录在正文里是本文档最有干货的部分值得逐条深入。要点一Python 并发与线程的红线——jax.thread_guard可以在多个 Python 线程中调用jax.jit、jax.grad这类 API但绝不能在线程里操作被 trace 函数内部正在追踪tracing的值——那样做大概率只会得到一个莫名其妙的错误。多控制器 JAX 下线程还有额外的隐患每个进程在给定设备上必须按相同顺序入队相同的操作而线程可能在不同进程中按不同顺序调度工作导致非确定性崩溃。为此 JAX 提供了jax.thread_guard。它是一个可作上下文管理器使用的配置对象启用后若某个 JAX 操作从设置守卫时的线程以外的线程发出就会抛错。用法与配置源码一一对应import jax with jax.thread_guard(True): ... # 只有本线程发出的多进程 JAX 操作被允许其底层定义位于 jax/_src/config.py这是一个名为jax_thread_guard的布尔状态默认False并通过update_thread_local_hook把取值同步到 C 层guard_lib.update_thread_guard_global_state见 jaxlib/guard_lib.cc 与其类型桩 jaxlib/_jax/guard_lib.pyi。jax.thread_guard(True)之所以能作为with语句使用是因为配置State.__call__返回StateContextManagerjax/_src/config.py。守卫的精确语义可以对照多进程测试 tests/multiprocess/thread_guard_test.py 得到验证守卫开启时把跨进程数组交给ThreadPoolExecutor里的线程执行jit函数慢路径与快路径 JIT 均覆盖随后block_until_ready会抛出包含 thread guard was set 的RuntimeError/ValueError若数组只分布在进程本地设备上则不触发错误——守卫只拦截跨进程的调度失序风险上下文管理器支持冗余嵌套同一线程内重复with jax.thread_guard(True)不报错也支持在嵌套内临时jax.thread_guard(False)关闭并在退出内层后自动恢复不同线程间嵌套守卫不被支持会抛出 Nested thread guards in different threads are not supported。顺带一提线程局部配置还有一个众所周知的特性——新建线程默认采用全局选项而不是创建它的作用域里的线程局部选项这条规则同样适用于本主题与传输守卫等所有线程局部配置。要点二多进程协调工具箱——jax.experimental.multihost_utils索引页点名的这套工具全部实现在 jax/experimental/multihost_utils.py源码 docstring 对每个函数的定位如下函数定位源码位置sync_global_devices(name)命名跨进程 barrier把名字做 CRC32 后经assert_equal校验所有 host/设备到达同一同步点jax/experimental/multihost_utils.pybroadcast_one_to_all(in_tree, is_sourceNone)把源进程默认 0 号进程的值广播给所有进程内部用 host-local → global 再psum的思路实现jax/experimental/multihost_utils.pyprocess_allgather(in_tree, tiledFalse)从每个进程收集值tiledFalse时沿新位置轴堆叠tiledTrue时拼接非完全可寻址数组需用tiledTruejax/experimental/multihost_utils.pyassert_equal(in_tree, fail_message)校验各进程拥有相同的值树不一致即抛AssertionErrorjax/experimental/multihost_utils.pyhost_local_array_to_global_array(local_inputs, global_mesh, pspecs)把各进程本地的数组组装成全局分片jax.Array每个设备按 mesh/pspec 拿到对应切片jax/experimental/multihost_utils.pyglobal_array_to_host_local_array(...)host_local_array_to_global_array的逆操作把全局数组视图还原为每主机本地的数组jax/experimental/multihost_utils.py典型用法示例要打印一个跨进程 sharded 的jax.Array直接print会抛RuntimeErrorFetching value forjax.Arraythat spans non-addressable ... devices is not possible此时可以先把它复刻到所有进程再在 0 号进程打印或在循环中每轮用sync_global_devices(epoch)保证各进程在同一个命名点上对齐避免数据并行训练步进失步。要点三长期存活产物的兼容性承诺——跨版本宁可失效不可错用由于导出模块、编译缓存条目这类产物活得比创建它的进程更久JAX 明确承诺了各自的生命周期边界索引页把它们归结为三条策略导出模块拥有显式的兼容窗口compatibility windows具体保证见 docs/501/export.md持久化编译缓存条目不做跨版本承诺但缓存键包含 jaxlib 版本——因此 JAX/版本升级会表现为缓存未命中cache miss触发重新编译而不是错误地复用不兼容的旧产物宁可多花编译时间也不承担错用风险一般 API 稳定性规则由 docs/api_compatibility.md 统一描述。进一步探索从索引页到源码的阅读路径如果你想沿着 501 的主题继续深入仓库提供了完整的一手材料多控制器入门与三种建全局数组的方法docs/501/multiprocess.md配套集群检测、mTLS 等底层实现在 jax/_src/distributed.py分布式数据加载策略docs/501/data-loading.md图解素材见 docs/_static/distributed_data_loading分布式容错完整演练docs/501/fault-tolerance.rst 及其配套脚本 docs/_static/fault_tolerance含数据并行、带恢复的数据并行两个完整训练示例多进程工具实现jax/experimental/multihost_utils.py它同时是sync_global_devices、broadcast_one_to_all、process_allgather、assert_equal以及两种 global/local 数组互转函数的唯一实现点线程守卫语义验证tests/multiprocess/thread_guard_test.py 与 C 侧 jaxlib/guard_lib.cc。需要重申的是本主题描述的分布式容错能力仍属实验性支持docs/501/fault-tolerance.rst 顶部明确提示当前仅在 GPU 上完整可用、存在粗糙边缘并可能变更若你在 TPU 上需要替代方案官方建议评估 Pathways。把这一点与上文各主题的配置前提放在一起就能安全地把 501 的知识用于真实的长时间、多主机、可重启的 JAX 生产任务。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考