完整指南:多进程场景下的数据切分、复制策略与实战代码)
JAX 分布式数据加载Distributed Data Loading完整指南多进程场景下的数据切分、复制策略与实战代码【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读当 JAX 程序运行在 多进程 / 多主机环境 中时一份完整训练数据的各分片往往散落在多个进程上如何让每个设备拿到它真正需要的那一份数据就成了正确性与性能的关键。本文以官方指南 docs/501/data-loading.md 为核心骨架系统讲解分布式数据加载的整体思路、为jax.Array构造Sharding的通用方法、四种高层加载方案、全量/部分复制Replication、以及纯数据并行与数据并行 模型并行两类工作负载下的完整tf.data实现并补充源码级佐证。读完本文你将能够在多进程 JAX 场景下独立设计数据管线并写出能避免数据分片放错设备这类隐性错误的生产代码。为什么需要分布式数据加载与两种替代方案的取舍指南开篇指出当 JAX 运行在 multi-host / multi-process 环境、且计算所需数据被拆分到多个进程时有三种数据获取方式可供选择分布式数据加载本文主题每个进程只加载本地设备需要的数据分片。通常更高效但实现更复杂。单进程加载全量全局数据由一个进程读取完整数据再通过 RPC 把需要的部分分发给其他进程。实现简单但代价高——训练循环可能因等数据而阻塞且每个进程都要消耗额外的网络带宽。所有进程都加载全量全局数据每个进程各自只使用所需部分。同样更简单但更浪费内存与 I/O 双重放大。在机器学习场景中训练循环在等待数据时会阻塞、带宽被白白消耗因此第 1 种方案分布式加载往往是更优选择。必须先牢记的隐患分片放错设备不会报错原文警告使用分布式数据加载时必须保证**每个设备例如每块 GPU / TPU都能访问到它运行计算所需的那一份输入分片**这正是它比上述两种替代方案更难正确实现的原因。如果错误的分片被放到了错误的设备上计算本身不会报错——因为计算无从得知输入数据本该是什么——但最终结果几乎必然是错的因为喂进去的数据和预期不一致。这一静默出错特性贯穿全文是后续所有方案设计时最重要的约束条件。通用方法为jax.Array定义Sharding分布式数据加载的核心对象是jax.Array与其关联的Sharding。思路不仅适用于批次数据也适用于任何不是由 JAX 计算直接产生的多进程jax.Array例如从 checkpoint 加载模型权重、加载一张空间上分片的大图。Sharding是什么一份数据到设备的布局说明书每个jax.Array都带有一个jax.sharding.Sharding它描述了全局数据的每个分片分别需要放在哪个全局设备上。当你从零创建一个jax.Array时必须同时为它创建一个Sharding——这是 JAX 理解数据如何跨设备排布的唯一途径。从源码看Sharding的基类定义在 jax/_src/sharding.py其核心 docstring 即Describes how ajax.Arrayis laid out across devices抽象接口包括device_set该Sharding覆盖的全部设备集合。在多控制器multi-controllerJAX 下是全局设备集合包含其他进程中不可寻址non-addressable的设备is_fully_replicated是否为全量复制即每个设备都持有完整数据的一份拷贝见 jax/_src/sharding.pyaddressable_devices当前进程可以寻址的设备子集实现见 jax/_src/sharding.py——它其实就是device_set中process_index等于当前进程的那些设备process_count()1时直接返回全集。Sharding具体选哪种由你的并行策略决定数据并行、模型并行等也可以根据原始数据在每个进程内如何产生来决定。创建好Sharding之后用sharding.addressable_devices()得到当前进程内需要为其加载数据的设备列表。可寻址设备addressable devices是比本地设备local devices更一般化的说法——目标始终是让每个进程的数据加载器把正确的数据喂给该进程的所有本地设备。直观示例1D 与 2D 切分假设需要一个形状为(64, 128)的jax.Array将其切分到 4 个进程、每进程 2 台设备共 8 台设备上。此时会产生 8 个互不相同的分片。切分方式有很多可以只沿第 2 维做一维切分每台设备拿到(64, 16)的分片图中每种颜色代表一个进程需要加载的分片——例如进程0的两台设备持有分片A、B对应全局数据的前(64, 32)上图出自原指南8 台设备各持一个分片进程 0 的两台设备负责分片 A 与 B。设备与分片的对应关系可以任意指定也可以做二维切分。但无论jax.Array如何被切分都必须保证每个进程的数据加载器装载该进程所需的分片。指南据此归纳出四种高层实现方法。四种高层实现方案Option 1每个进程都加载全局数据最省事、最浪费每进程执行两步① 加载它需要的全部全局数据值② 只把本进程本地设备需要的分片转移到本地设备上。这种方式在加载效率上并不划算每个进程都会丢弃本地设备用不上的数据总摄入量可能远超必需。但它的优点是简单、必定可用当全局数据量较小时例如加载小模型 checkpoint其额外开销完全可以接受。Option 2每个设备一条独立数据管线per-device data pipeline每个进程为它的每一台本地设备各设置一个数据加载器——即每台设备只加载自己需要的那一份分片。好处是按需加载效率高并且把每台设备独立看待、逐个思考通常比同时考虑进程内全部设备更简单对照 Option 3。潜在问题是同时运行多个并发数据加载器可能带来性能问题。Option 3合并的按进程数据管线consolidated per-process data pipeline每进程执行两步① 设置一个数据加载器一次性加载其所有本地设备需要的数据② 在转移数据前先把本进程的这批数据切分成各本地设备的分片。这是最高效的分布式加载方式——全局数据只被读取必要的份数、每个进程只有单个加载器。但它也最复杂既要搞清楚每台设备到底需要哪部分数据又要设计一个只加载这些数据理想情况下不多加载一字节的单一数据管线。Option 4先按方便的方式加载再在计算内部重新分片reshard inside computation这个概念最微妙但常常比前三种更好实现当恰好精确加载每台/每进程所需数据难以做到时仍然可以做到每进程加载1 / num_processes的数据——只是切分方式不对。回到上文二维切分的例子假设对每个进程而言加载数据的一整列更容易。做法是先用一个表达按列分片的Sharding创建jax.Array直接把它传入计算再调用jax.lax.with_sharding_constraint立即把列分片的输入重排成目标Sharding。由于重分片发生在计算内部会经由加速器互联链路完成例如 TPU ICI 或 NVLink。Option 4 与 Option 3 有相似收益每个进程仍然只有一个数据加载器全局数据在所有进程间恰好只被加载一次额外好处是数据加载方式更灵活。代价是它占用加速器互联带宽做重分片某些负载可能因此变慢并且输入数据必须额外表达为一种Sharding除目标Sharding之外即存在输入分片 目标分片两套布局信息。四种方案的本质区别可概括为数据从存储介质读到进程本地、再从进程本地摆到设备上这两个环节由谁来完成、各加载多少。小数据选 Option 1追求极致 I/O 且能精确分片选 Option 3难精确分片时优先考虑 Option 4。复制Replication全量复制与部分复制当多台设备持有相同的数据分片时就出现了复制。前面四种方案在复制场景下依然适用唯一区别是某些进程最终可能加载相同的数据分片。全量复制Full replication全量复制指所有设备都持有数据的完整拷贝——此时数据的分片就是整个数组值本身。回到 8 台设备每进程 2 台的例子若做全量复制最终会得到8 份完整数据每份都未分片、完整存在于单台设备上。部分复制Partial replication部分复制指数据存在多份拷贝且每份拷贝本身又被切分到多台设备上。对于同一个数组部分复制通常存在多种实现方式提示给定数组形状时全量复制的Sharding是唯一的即不存在多种全量复制布局。指南给出两个典型例子例一每份拷贝被切分到某进程的两台本地设备上共 4 份拷贝。这意味着每进程都必须加载完整的全局数据——因为它的本地设备合起来持有完整的一份。例二每份拷贝仍被切分到两台设备但每对设备跨越两个进程进程0粉与进程1黄都只需加载第一行数据进程2绿与进程3蓝都只需加载第二行数据。第二个例子说明部分复制会显著改变各进程的加载职责——即使布局复杂仍能通过正确设计Sharding让每个进程只读全局数据的一部分。实战场景一纯数据并行Data Parallelism在纯数据并行不含模型并行下每台设备上复制一份完整模型每个模型副本即每台设备接收不同的 per-replica batch每副本批数据。若把输入数据表达成单个jax.Array则该数组在一步中包含所有副本的数据称为global batch其中每个分片正好是某一个 per-replica batch。标准表达方式是沿设备做一维切分——即 global batch 就是所有 per-replica batch 沿 batch 轴拼接而成。沿用 4 进程 × 2 设备的例子进程0应拿到 global batch 的前四分之一8 个分片中的前 2 个进程1拿第二个四分之一依此类推。数据并行的关键性质无需关心哪个 batch 落在哪个副本上第一个四分之一到底是什么如何确保进程0恰好拿到它——好消息是数据并行本身让你无需回答这些问题每个设备对应一个执行相同操作的模型副本因此 global batch 内部哪个 per-replica batch 落到哪台设备根本无所谓。也就是说你可以在 global batch 内自由地重排 per-replica batch等价于随机化每台设备拿到的数据分片。对普通jax.Array而言重排数据分片通常不是好主意相当于对数组值做置换但对数据并行则完全合理——global batch 的顺序本就没有语义。这一性质极大地简化了数据加载每台设备只需要一条独立的 per-replica batch 流。而大多数数据加载器很容易实现每个进程一条独立管线 把本进程产出的 batch 切分成多个 per-replica batch即ds.shard(num_shardsjax.process_count(), indexjax.process_index())先按进程分片、再做 per-replica 切分。这是前文按进程合并管线思路的一个实例原文也说明可换用前文其他方案相对简单且高效。纯数据并行 tf.data的可运行示例import jax import tensorflow as tf import numpy as np ################################################################################ # Step 1: setup the Dataset for pure data parallelism (do once) ################################################################################ # Fake example data (replace with your Dataset) ds tf.data.Dataset.from_tensor_slices( [np.ones((16, 3)) * i for i in range(100)]) ds ds.shard(num_shardsjax.process_count(), indexjax.process_index()) ################################################################################ # Step 2: create a jax.Array of per-replica batches from the per-process batch # produced from the Dataset (repeat every step). This can be used with batches # produced by different data loaders as well! ################################################################################ # Grab just the first batch from the Dataset for this example per_process_batch ds.as_numpy_iterator().next() mesh jax.make_mesh((jax.device_count(),), (batch,)) sharding jax.NamedSharding(mesh, jax.sharding.PartitionSpec(batch)) global_batch_array jax.make_array_from_process_local_data( sharding, per_process_batch)这段代码演示了两个贯穿全文的关键 APIjax.process_index()/jax.process_count()返回本进程在多进程环境下的序号与进程总数。从源码看它们实现在 jax/_src/xla_bridge.pyprocess_index直接取get_backend(backend).process_index()process_count通过遍历所有设备的process_index取最大值加一得到。tf.data.Dataset.shard正是靠这两个值让每个进程各取不同的原始数据段从而保证全局数据只被读取一次jax.make_array_from_process_local_data(sharding, local_data)把当前进程本地已有的数据按给定Sharding组装成分布式jax.Array。它是更通用的jax.make_array_from_callback的一个常见特例见 jax/_src/array.py 的 docstringcreates distributed tensor using the data available in process... takes care of the index wrangling其内部会基于Sharding计算出本进程各地址设备应占据本地数组的哪些切片再经batched_device_put分发到各设备。源码 docstring 特别强调如果任意两个主机互为副本则local_data必须完全相同——这正是下面复制场景能正确工作的前提。make_array_from_process_local_data的行为可概括为你只需按进程提供拼接好的本地数据块和目标 Sharding索引换算哪个设备该拿哪一段交给 JAX。它支持更一般的混合 multi-host 复制与多轴分片但要求你正确计算 process-local 数据的大小与内容以满足切分约束global_shape可选缺省时按均匀切分从本地数据与 Sharding 推断非均匀切分则必须显式给出详见源码 docstring。实战场景二数据并行 模型并行Data Model Parallelism纯模型并行无数据并行时整个模型只有一份被切分到全部设备上数据通常在所有设备上全量复制。而本指南关注的是同时使用数据并行与模型并行的情况多个模型副本每个副本被切分到多台设备上数据在每个模型副本上部分复制同一副本内的各设备拿到相同的 per-replica batch不同副本间拿到不同的 per-replica batch。模型副本位于单进程内最简单的情形从数据加载角度最简单的是把每个模型副本限制在单个进程的本地设备内。把例子改为 2 进程 × 4 设备每个模型副本切分到某进程的 2 台本地设备上于是每进程有 2 个副本、全局共 4 个副本。此时输入仍表达为单个jax.Array并做一维切分每分片一个 per-replica batch但与纯数据并行不同引入部分复制把一维切分的 global batch 做成 2 份拷贝——因为每个模型副本由 2 台设备组成它们都需要同一份 per-replica batch。模型副本留在单进程内的好处是可以复用纯数据并行的整套设置只需额外把 per-replica batch复制到对应设备的本地设备上。**把 per-replica batch 复制到正确的设备上极其重要** 数据并行的关键 trick让你不必关心哪个 batch 落在哪个副本但**你必须在乎单个副本只能拿到一个 batch**。例如把同一 batch 复制给副本内的两台设备是正确做法但若不注意本地设备的装载顺序可能造出名义上复制、实际未复制的错误布局——尽管 Sharding以及并行策略声明数据是复制的。好消息是对进程内的模型并行如果误把本应复制的数据建成未复制状态JAX 会抛出错误拦截这正是make_array_from_process_local_data等构造路径会做的进程内一致性校验。但跨进程的模型并行则没有这种保护——详见下一小节。进程内模型并行 数据并行tf.data完整示例import jax import tensorflow as tf import numpy as np ################################################################################ # Step 1: Set up the Dataset with a different data shard per-process (do once) # (same as for pure data parallelism) ################################################################################ # Fake example data (replace with your Dataset) 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_shardsjax.process_count(), indexjax.process_index()) ################################################################################ # Step 2: Create a jax.Array of per-replica batches from the per-process batch # produced from the Dataset (repeat every step) ################################################################################ # Grab just the first batch from the Dataset for this example per_process_batch ds.as_numpy_iterator().next() num_model_replicas_per_process 2 # set according to your parallelism strategy num_model_replicas_total num_model_replicas_per_process * jax.process_count() # Create an example Mesh for per-process data parallelism. Make sure all devices # are grouped by process, and then resize so each row is a model replica. 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) # Double check that each replicas devices are on a single process. 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]) # Shard the data across model replicas. You dont shard across the # data_parallelism mesh axis, meaning each per-replica shard will be replicated # across that axis. sharding jax.sharding.NamedSharding( mesh, jax.sharding.PartitionSpec(model_replicas)) global_batch_array jax.make_array_from_process_local_data( sharding, per_process_batch)这段代码的关键点在于Mesh与PartitionSpec的组合先把所有进程的本地设备按进程排成二维np.ndarray再reshape(num_model_replicas_total, -1)使得每一行正好是一个模型副本代码用assert逐一核验副本内所有设备位于同一进程。Mesh的两条轴分别命名为model_replicas模型副本轴与data_parallelism数据并行轴。PartitionSpec(model_replicas)只沿副本轴分片、不沿data_parallelism轴分片——该轴上的设备因此自动共享同一份 per-replica batch即实现每副本内的复制。Mesh/NamedSharding/PartitionSpec的底层约定与make_mesh辅助函数可进一步在 jax/_src/sharding_impls.py 及 docs/201/sharding.md 中查看。模型副本跨进程分布需要跨进程协调当模型副本跨进程分布时可能因为单个副本放不进一个进程或设备分配本就如此数据加载就更有意思了。回到 4 进程 × 2 设备的配置把设备按如下方式分配给副本仍是 4 个模型副本、每个切分到 2 台设备唯一区别是每个副本的两台设备分属不同进程且每个进程只为两个副本各负责一份拷贝。这种跨进程拆分看似随意甚至多余在这个例子里确实如此但真实部署可能正是为了充分利用设备间的通信拓扑而这样安排设备。此时数据加载变复杂因为需要跨进程协调纯数据并行与副本在进程内的情形只要求每个进程加载互不相同的数据流而这里某些进程必须加载相同数据另一些进程必须加载不同数据上例中进程0粉与进程2绿必须加载相同的 2 个 per-replica batch进程1黄与进程3蓝也须加载相同但与进程 0/2 不同的 2 个 per-replica batch并且每个进程不能把自己负责的 2 个 per-replica batch 搞混虽然不关心哪个 batch 落到哪个副本但必须保证同一副本的所有设备拿到的是同一个batch——否则就会出现同副本内两台设备数据不一致的错误布局。原文警告截至 2023 年 8 月JAX **无法检测跨进程的 jax.Array 分片本应复制却未复制**的情况计算运行时会直接产出错误结果。因此这类场景必须格外小心自行确保每个进程装载的 batch 及复制关系正确。提示该行为以你实际使用的 JAX 版本为准。最终要把正确的 per-replica batch 放到每台设备上需要把全局输入数据表达为特定的jax.Array布局各进程只提供本进程负责的、且与协作进程保持一致的数据块再交由对应的Sharding描述复制与分片关系如下图所示从源码看构造分布式jax.Array的三条常用路径理解底层 API 有助于在不同方案间切换。指南与源码共同勾勒出如下工具集API位置用途与适用场景jax.make_array_from_process_local_data(sharding, local_data[, global_shape])jax/_src/array.py数据已在本进程、按给定Sharding组装分布式数组Option 2/3 的典型落点自动完成各地址设备到本地数组切片的索引换算jax.make_array_from_callback(sharding, callback, global_shape)jax/_src/array.py最通用的构造入口由回调函数按需取各分片make_array_from_process_local_data是其常见特例sharding.addressable_devices()jax/_src/sharding.py返回当前进程需为其提供数据的所有设备用于驱动每进程数据管线sharding.addressable_devices_indices_map(shape)jax/_src/sharding.py返回地址设备到全局数组局部切片的映射是每台设备应加载哪一段的权威依据jax.lax.with_sharding_constraint(x, sharding)jax.lax在计算内部把输入立即重分片到目标ShardingOption 4 的关键重分片经加速器互联完成jax.process_index()/jax.process_count()jax/_src/xla_bridge.py让Dataset.shard等实现按进程切分数据jax.make_mesh/jax.sharding.Mesh/NamedSharding/PartitionSpecjax/_src/sharding_impls.py用命名轴描述分片与复制策略配合make_array_from_process_local_data落地方案选择速查与易错点清单按场景选择方案全局数据很小如小 checkpoint直接Option 1每进程加载全量简单且开销可接受每台设备所需的 shard 可以精确定义、且进程内无设备间复制负担Option 2per-device 管线追求最少数据摄入、且能构造恰为进程所需的单一管线Option 3consolidated per-process 管线精确分片难以实现、但按进程加载 1/num_processes很容易Option 4加载后经with_sharding_constraint在计算内重分片注意它消耗互联带宽。必须避免的错误按指南明确列出分片放错设备不会报错只会静默产生错误结果——因此务必让加载逻辑直接由Sharding如addressable_devices_indices_map驱动而不是手写想当然的索引纯数据并行时不要尝试让进程 0 精确拿第一个四分之一——利用batch 落到哪个副本无所谓的性质用Dataset.shard(process_count, process_index)让每个进程各取一段即可出现复制尤其部分复制时同一副本内的所有设备必须拿到同一 per-replica batch进程内误建未复制数据 JAX 会报错拦截但跨进程场景 JAX 不检测截至指南编写时只能靠代码与测试保证若多个进程互为副本传给make_array_from_process_local_data的local_data必须完全一致源码 docstring 的明确约束。更进一步的背景知识可继续阅读仓库中的相关主题文档多进程执行基础、分片与并行策略Sharding、设备放置Placement以及并行训练教程。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考