ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

JAX 多主机工具集实战:multihost_utils 六大 API 详解与源码级原理

JAX 多主机工具集实战:multihost_utils 六大 API 详解与源码级原理 JAX 多主机工具集实战multihost_utils 六大 API 详解与源码级原理【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本文以 JAX 仓库中的 jax.experimental.multihost_utils 模块 为主线系统讲解多进程multi-process环境下在主机host之间同步、广播、聚合数据以及全局/主机本地jax.Array互转的完整工具集。读完本文你将掌握broadcast_one_to_all、sync_global_devices、process_allgather、assert_equal、host_local_array_to_global_array与global_array_to_host_local_array六个核心 API 的语义、参数、典型用法与底层实现并能将其直接应用于分布式训练的数据加载、检查点保存与容错场景。背景多进程下的两种数据视图在 JAX 的多进程multi-controller程序中每个进程管理各自机器上的一组设备jax.local_device_count()所有进程的设备合起来构成全局设备集jax.device_count()。在这种模型下数据存在两种视角主机本地host local数据每个进程持有的一份数据通常来自本机磁盘的 batch各进程的数据可以互不相同例如各进程持有不同的 shard。全局global数据以jax.Array形式跨所有进程的设备分布通过NamedShardingPartitionSpec描述每个分片落在哪个设备的哪个索引区间。jax.experimental.multihost_utilsAPI 参考页正是为解决这两种视图之间的转换以及多进程之间的同步、校验、广播、聚合而设计的一组高层工具。其模块 docstring 的定位是“Utilities for synchronizing and communication across multiple hosts”。该模块包含的公开 API 如下以 RST 参考文档 中的 autosummary 为准API作用broadcast_one_to_all将源主机默认 host 0的数据广播到所有主机sync_global_devices在所有主机/设备之间创建同步屏障process_allgather跨进程聚合stack 或 concat各主机数据assert_equal校验所有主机持有相同的树形数据host_local_array_to_global_array主机本地数据 → 全局分片jax.Arrayglobal_array_to_host_local_array全局jax.Array→ 主机本地数组同步与校验sync_global_devices 与 assert_equalsync_global_devices全局屏障sync_global_devices(name: str)在所有主机/设备之间创建一个屏障barrier是所有多进程程序协调进度的基础设施。其实现非常简洁见 multihost_utils.pydef sync_global_devices(name: str): Creates a barrier across all hosts/devices. h np.uint32(zlib.crc32(name.encode())) assert_equal(h, fsync_global_devices name mismatch ({name}))它将传入的name字符串做 CRC32 哈希后交给assert_equal完成一次跨主机的集体通信。name参数用于标识这个同步点——所有主机必须传入完全相同的字符串否则校验失败并抛出AssertionError。测试 test_sync_global_devices_error 明确验证了这一点当 host 0 传入test message而其他主机传入test message2时所有进程都会抛出AssertionError。典型使用场景from jax.experimental import multihost_utils # 各进程完成本机数据准备后等待所有主机就绪 multihost_utils.sync_global_devices(data_ready) # 训练若干步之后再次同步确保各进程进度一致 multihost_utils.sync_global_devices(step_100_done)在 migrate_pmap 文档 中该函数被用于 pmap 迁移场景确保所有进程同步后再执行后续操作。注意 test_sync_global_devices_mesh_context_manager 表明它在jax.set_mesh上下文内也能正常工作。assert_equal跨主机数据一致性校验assert_equal(in_tree, fail_message: str )验证所有主机持有相同的 pytree 数据。其实现multihost_utils.py分为三步对每个叶子构造“期望值”若已是全局且非 fully-addressable 的jax.Array取本机可访问的分片否则转成 numpy 数组若是标量则先expand_dims再按进程数复制拼接调用process_allgather(in_tree, tiledTrue)收集所有主机的真实值逐叶子比较任一不一致即抛出AssertionError(f{fail_message}. Expected: {out}; got: {in_tree}.)。它要求各主机传入的 pytree 结构一致常用于调试时验证各进程的输入数据、随机种子或模型权重是否对齐。注意fail_message为空时仍会抛出AssertionError只是错误信息不附带自定义说明。广播与聚合broadcast_one_to_all 与 process_allgatherbroadcast_one_to_all从源主机广播到全体broadcast_one_to_all(in_tree, is_sourceNone)将数据从源主机广播到所有其他主机返回的 pytree 中每个叶子都是源主机的数据副本。关键语义multihost_utils.pyin_tree任意 pytree每个数组在各主机上必须形状一致is_source可选布尔值标记当前调用者是否为源。None时默认 host 0即jax.process_index() 0为源返回与in_tree结构相同的 pytree所有叶子均变为源主机的数据。其底层实现值得玩味它构造一个(processes, local_devices)的全局网格非源主机用np.zeros_like填充输入再通过host_local_array_to_global_array将数据以P(processes)分片方式放入网格最后在jax.jit(_psum, out_shardingsP())中对 process 轴做求和源主机数据与非源的零相加天然得到源数据从而借助一次集体归约完成广播——这是一个利用分片求和巧做广播的实现。测试 test_broadcast_one_to_all 验证各进程数据np.arange(4) process_index * 4广播后所有进程得到的都是np.arange(4)。该函数同样支持 pytree如(x, x)和uint8等非默认 dtype见 test_broadcast_one_to_all_uint8。process_allgather跨进程聚合process_allgather(in_tree, tiled: bool False)从所有进程收集数据返回 numpy 数组的 pytree。核心参数tiled决定聚合方式multihost_utils.pytiledFalse默认在位置轴 0 处堆叠stack输出形状为(num_processes, *x.shape)tiledTrue沿第 0 轴拼接concatenate输出形状为(num_processes * x.shape[0], *x.shape[1:])输入为标量时无论tiled取值如何输出都会被堆叠为(num_processes,)。对于输入类型还有区分非 fully-addressable 的全局jax.Array只支持tiledTrue结果是对该全局数组的完全复制replicate。若传tiledFalse会抛出ValueError(Gathering global non-fully-addressable arrays only supports tiledTrue)见 test_process_allgather_array_not_fully_addressablenumpy 数组或 fully-addressable 的jax.Array按tiled决定 stack 或 concat返回值为 numpy 数组np.asarray(out.addressable_data(0))可直接参与 Python 层逻辑。测试 test_process_allgather_stacked 与 test_process_allgather_concatenated 分别验证了两种模式在多维数组、标量上的形状与数值正确性。一个重要的工程细节模块将恒等函数_identity_fn定义在模块顶层multihost_utils.py并注释说明这是为了让process_allgather在多次调用时不重复编译。测试 test_process_allgather_cache_hit 用count_pjit_cpp_cache_miss验证了连续两次对不同数据的process_allgather只发生一次编译——即该 API 是可缓存、可高频调用的。在 multi_process 文档 中该函数被推荐为打印跨进程全局数组的标准方式当jax.Array跨越非本机可访问设备而无法直接取值时用process_allgather聚合后再打印。全局与主机本地数组互转host_local_array_to_global_array与global_array_to_host_local_array是数据加载与pjit迁移场景中最重要的两个 API。它们要求的全局网格必须是连续网格contiguous mesh——即每个主机的设备在网格中构成一个子立方subcube若网格不连续源码 docstring 明确建议改用jax.make_array_from_callback或jax.make_array_from_single_device_arrays。host_local_array_to_global_array主机本地 → 全局def host_local_array_to_global_array( local_inputs: Any, global_mesh: jax.sharding.Mesh, pspecs: Any)接受各主机可能不同的主机本地数据按照global_mesh/pspecs定义的分片将每台主机设备放入对应的数据切片组装成全局jax.Array。源码示例multihost_utils.py展示了三种典型分片global_mesh jax.sharding.Mesh(jax.devices(), x) pspecs jax.sharding.PartitionSpec(x) host_id jax.process_index() # 每个进程的数据 np.arange(4) * host_id按 x 轴跨设备分片 arr host_local_array_to_global_array(np.arange(4) * host_id, global_mesh, pspecs)当pspecs P()空 PartitionSpec时表示跨所有轴复制全局形状保持不变且该数组在每个主机、每个设备上都被复制。需要注意当 pspec 指示复制replication而各主机 local_inputs 不相同时属于未定义行为undefined behavior。实现的几个关键点multihost_utils.py若输入已是全局且非 fully-addressable 的jax.Array直接原样返回pspecs不允许为None需用P()表达复制意图否则抛ValueError测试 test_host_local_array_to_global_array_none_error若输入jax.Array的 sharding 与目标本地分片等价则直接复用底层 buffer 指针避免拷贝测试 test_host_local_array_to_global_array_same_sharding_array 用unsafe_buffer_pointer()验证零拷贝支持float0与 PRNG key 数组PRNGKeyArray等特殊类型通过pxla.batched_device_put将各切片放置到本地设备。global_array_to_host_local_array全局 → 主机本地def global_array_to_host_local_array( global_inputs: Any, global_mesh: jax.sharding.Mesh, pspecs: Any)与上一函数互为逆操作把全局jax.Array收缩为主机本地数组每个进程只保留本机设备可访问的部分。若输入已是 fully-addressable主机本地则原样返回multihost_utils.py。测试 test_global_array_to_host_local_array 验证全局形状(8, 2)、按P(x,y)分片的数组经P(x)收缩后得到形状(2, 2)的主机本地数组。与 pjit 配合的标准迁移模式jax_array_migration 文档 给出了二者与pjit配合的标准范式这也是从旧式 GDAGlobalDeviceArray迁移到jax.Array的机械式替换路径from jax.experimental import multihost_utils global_inps multihost_utils.host_local_array_to_global_array( local_inputs, mesh, in_pspecs) global_outputs pjit(f, in_shardingsin_pspecs, out_shardingsout_pspecs)(global_inps) local_outs multihost_utils.global_array_to_host_local_array( global_outputs, mesh, out_pspecs)语义上host_local_array_to_global_array是一种“类型转换”它将只有本地分片的值调整为其全局形状使pjit可以按全局语义消费。最典型的应用是多进程环境下的数据 batch——各进程读入各自的 batch 后先用它转成全局jax.Array再送入pjitjax_array_migration.mdbatch multihost_utils.host_local_array_to_global_array( batch, mesh, batch_partition_spec)对于完全复制fully replicated的输入各进程形状相同、P(None)分片形状本就是全局的无需经过转换jax_array_migration.md# P(None) 表示完全复制无需 host_local_array_to_global_array pjit(f, in_shardingsNone, out_shardingsNone)(key) # 混合输入key 复制、local_inp 按 data 轴分片 global_inp multihost_utils.host_local_array_to_global_array( local_inp, mesh, P(data)) global_out pjit(f, in_shardings(P(None), P(data)), out_shardings...)(key, global_inp)进阶能力抢占同步点与容错除上述 6 个公开 API 外multihost_utils.py 还提供了两个面向生产环境的进阶设施。reached_preemption_sync_point抢占preemption安全检查点reached_preemption_sync_point(step_id: int) - bool用于在集群调度器发出抢占通知默认是 SIGTERM时协调所有主机安全地保存检查点。机制如下multihost_utils.py任一主机收到抢占通知后通知会传播到所有主机并触发后台同步协议同步协议计算所有主机上报的 step_id 最大值以max 1作为安全保存检查点的 step各主机在每个训练步都调用本函数直到返回True即当前step_id等于安全 step时开始保存检查点。使用前提所有主机必须从相同 step 开始训练且每个训练步都调用。官方示例def should_save(step_id: int) - bool: # 抢占触发的按需检查点 if multihost_utils.reached_preemption_sync_point(step_id): return True # 常规周期检查点 return step_id - last_saved_checkpoint_step save_interval_steps若抢占同步管理器未初始化未启用jax_enable_preemption_service配置调用会抛出RuntimeError若分布式运行时未初始化则返回False。live_devices容错设备集合上下文管理器live_devices是一个低层原语用于让多控制器 JAX 程序具备容错能力。它是上下文管理器multihost_utils.py以with multihost_utils.live_devices(jax.devices()) as devices:的形式使用yield 当前存活的健康设备子集供主体代码在存活设备上运行。它具备两个关键语义屏障语义所有进程必须对“哪些设备存活”达成一致否则行为会发散。live_devices等待包含返回存活设备的所有进程都进入with块后才向所有进程返回相同的存活设备集合 A原子性语义当进程退出with块时要么所有持有 A 中设备的进程都成功执行完块内代码要么全部抛出异常——不可能出现部分进程进except分支、部分进程进else分支的撕裂状态。若检测到设备所在进程死亡或重启incarnation id 变化会抛出ProcessFailureError其failed_devices属性包含失败设备集合multihost_utils.py。try: with multihost_utils.live_devices(jax.devices()) as devices: # 在存活设备上运行 JAX 代码 pass except multihost_utils.ProcessFailureError as e: # 部分设备所在进程失败e.failed_devices 给出失败设备 pass注意该 API 处于活跃开发中源码标注 “UNDER ACTIVE DEVELOPMENT AND IS NOT STABLE”需要先初始化分布式运行时否则抛RuntimeError(Distributed JAX not initialized.)且传入设备集合必须包含至少一个本地设备。测试 test_live_devices 在无故障场景下验证了live_devices(jax.devices())返回全部设备。工程实践要点与注意事项基于源码实现与测试用例总结多进程编程中使用本模块的注意事项集体通信必须在所有进程上对称执行sync_global_devices、process_allgather等是集体操作。如 multi_process 文档 所警示不要只让process_index() 0发起device_put之类的通信否则会导致死锁——只有进程 0 发起集体通信而其他进程未参与时它会无限期等待。process_allgather适合高频调用由于恒等函数被提升到模块顶层并复用jit缓存连续调用不会重复编译test_process_allgather_cache_hit它不受外层jax.set_mesh上下文影响test_process_allgather_set_mesh。网格连续性约束host_local_array_to_global_array/global_array_to_host_local_array要求网格中每个主机的设备构成连续子立方否则应改用jax.make_array_from_callback或jax.make_array_from_single_device_arrays。pspecs不能传None表示复制必须显式使用P()否则抛出带提示的ValueError。主机本地输入与pjit的配合多进程环境下任何主机本地输入尤其是数据 batch都应先经host_local_array_to_global_array转为全局jax.Array再送入pjit完全复制输入P(None)则无需转换。相关文档与测试模块 API 参考docs/jax.experimental.multihost_utils.rst源码实现jax/experimental/multihost_utils.py多进程测试tests/multiprocess/multihost_utils_test.py多进程编程指南docs/multi_process.md含process_allgather调试用法与死锁规避jax.Array 迁移指南docs/jax_array_migration.md含与 pjit 配合的完整迁移模式pmap 迁移指南docs/migrate_pmap.md含sync_global_devices同步用法【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表