尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

JAX 分布式数据加载全指南:jax.Array 分片、数据并行与模型并行的工程实践

JAX 分布式数据加载全指南:jax.Array 分片、数据并行与模型并行的工程实践 JAX 分布式数据加载全指南jax.Array 分片、数据并行与模型并行的工程实践【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本指南源自 JAX 官方文档 docs/distributed_data_loading.md围绕多进程multi-host/multi-process环境下如何正确、高效地把分布式存储的数据送入 JAX 计算展开先建立为 jax.Array 构造 Sharding的统一思维框架再依次讲解四种数据加载方案、复制replication、纯数据并行以及进程内/跨进程的数据模型并行并给出基于tf.data的可直接运行的完整代码。读完本文你将掌握Sharding.addressable_devices()、jax.make_array_from_process_local_data()、jax.lax.with_sharding_constraint()等核心 API 的适用场景并能针对自己的并行策略设计出正确的数据流水线。为什么需要分布式数据加载在运行于多主机或多进程环境下的 JAX 程序中计算所需的原始数据往往分散在多个进程上例如每个进程各持有一份本地文件分片。此时分布式数据加载指的是每个进程只加载/读取自己所需的那部分数据分片而不是把全局数据全部读进来。分布式数据加载通常更高效但也更复杂。与之相对的两种朴素替代方案分别是单进程加载全量全局数据切分后通过 RPC 把需要的部分发给其他进程每个进程都加载全量全局数据每个进程只使用其中需要的部分。这两种方案通常更简单但代价更高在机器学习训练中训练循环可能因等待数据而阻塞且每个进程都会消耗多余的网络带宽。使用分布式数据加载时必须保证**每个设备如每张 GPU 或每块 TPU都能访问到它执行计算所需的那份输入数据分片**。这正是分布式数据加载比上述替代方案更难正确实现的原因如果错误的数据分片落在了错误的设备上计算本身不会报错——因为计算无从知道输入数据本该是什么——但最终结果往往是错误的因为输入数据与预期不符。加载 jax.Array 的通用方法考虑从非 JAX 产生的原始数据创建一个jax.Array的场景。这些概念不仅适用于加载批量数据记录也适用于任何非 JAX 计算直接产生的多进程jax.Array例如1从 checkpoint 加载模型权重2加载一张很大的按空间分片spatially-sharded的图像。每一个jax.Array都关联着一个jax.sharding.Sharding它描述了每个全局设备需要全局数据中的哪一份分片。当你从头创建jax.Array时也必须同时创建它的Sharding这样 JAX 才能理解数据在设备间的布局。你可以创建任意想要的Sharding实践中通常根据所采用的并行策略来选择后文会详细讲解数据并行与模型并行也可以根据每个进程内原始数据的生产方式来选择。一旦定义了Sharding就可以调用jax.sharding.Sharding.addressable_devices()获取当前进程需要为之加载数据的设备列表。可寻址设备addressable devices是比本地设备local devices更一般的概念其目标正是确保每个进程的数据加载器能为该进程的所有本地设备提供正确的数据。从源码看该方法是Sharding抽象上的一个通用接口见 jax/_src/sharding.pyNamedSharding还针对多个对象共用同一 mesh的常见情况做了专门优化见 jax/_src/named_sharding.py。一个具体的分片示例假设需要一个形状为(64, 128)的jax.Array要把它切分到 4 个进程、每个进程 2 个设备共 8 个设备上。这会得到 8 个互不相同的数据分片data shard每个设备一个。分片方式有很多种例如沿数组第二维做一维分片让每个设备持有(64, 16)的分片上图中每个数据分片用不同颜色标识由哪个进程加载。例如假设进程0的 2 个设备分别持有分片A和B即全局数据最前面的(64, 32)部分。你也可以选择不同的分片→设备映射例如再例如二维分片无论jax.Array如何分片你都必须确保每个进程的数据加载器被提供/加载的是该进程所需的那部分全局数据分片。实现这一点有几种高层方法见下文 Option 1~4。四种数据加载方案Option 1每个进程加载全量全局数据采用该方案时每个进程加载所需的完整全局值只把需要的分片传输给本进程的本地设备。这不是高效的分布式数据加载方式因为每个进程都会丢弃其本地设备用不到的数据总数据摄入量可能高于实际所需。但它的优点是实现简单、可用对某些负载例如全局数据量很小而言性能开销可以接受。Option 2逐设备数据流水线per-device data pipeline该方案下每个进程为它的每个本地设备各设置一个数据加载器即每个设备拥有自己独立的、只加载所需分片的数据加载器。它的优点是在数据加载量上是高效的有时逐个设备独立考虑比一次性考虑进程的所有本地设备更简单可对比下文 Option 3。缺点是多个并发数据加载器有时会带来性能问题。Option 3合并的逐进程数据流水线consolidated per-process data pipeline采用该方案时每个进程设置单个数据加载器加载其所有本地设备所需的数据先把本地数据切分好再传输到每个本地设备。这是最高效的分布式数据加载方式但也是最复杂的既需要判断每个设备需要哪些数据又要构造一个只加载所有这些数据最好不多不少的单一数据加载器。Option 4以某种方便的方式加载在计算内部重新分片该方案解释起来更有挑战性但通常比 Option 1~3 更易实现。设想一个场景很难甚至不可能为逐设备或逐进程加载器精确构造恰好加载所需数据的数据加载器但仍然可能为每个进程构造一个加载1 / num_processes数据的加载器——只是分片方式不对。以前面的 2D 分片示例继续假设每个进程更方便加载数据的一整列column。此时你可以先创建一个带有按列分片Sharding的jax.Array直接把它传入计算然后立即使用jax.lax.with_sharding_constraint把按列分片的输入重分片reshard到目标分片。由于重分片发生在计算内部它将通过加速器互联链路如 TPU ICI 或 NVLink完成。Option 4 与 Option 3 有相似的好处每个进程仍然只有单个数据加载器全局数据在所有进程间恰好被加载一次额外的好处是数据加载方式更灵活。但代价是它使用加速器互联带宽执行重分片可能会拖慢某些负载并且它要求输入数据除了目标Sharding之外还要表达为另一个独立的Sharding。复制Replication复制描述的是多个设备持有同一份数据分片的情形。上述 Option 1~4 在存在复制时依然适用唯一的区别是某些进程可能最终加载了相同的数据分片。下面介绍完全复制与部分复制。完全复制Full replication完全复制指所有设备都持有数据的完整副本即数据分片就是整个数组的值。以下面示例为例总共有 8 个设备每进程 2 个最终会得到 8 份完整数据的副本每份副本都是不分片的即每份完整副本位于单个设备上部分复制Partial replication部分复制指数据存在多份副本且每份副本又被切分到多个设备上。对于给定数组通常存在多种部分复制方式注意对给定数组形状完全复制的Sharding永远是唯一的。下面给出两个示例。第一个示例中每份副本被切分到某个进程的两个本地设备上共 4 份副本。这意味着每个进程都需要加载全量全局数据因为它的本地设备合起来持有一份完整副本第二个示例中每份副本仍然被切分到两个设备上但每个设备对device pair横跨两个不同进程。进程0粉色和进程1黄色都只需要加载数据的第一行进程2绿色和进程3蓝色都只需要加载第二行上面已经梳理了创建jax.Array的高层方案接下来把它们应用到机器学习的分布式数据加载场景中。纯数据并行Data parallelism在纯数据并行不含模型并行时模型在每个设备上各复制一份replicate每个模型副本即每个设备接收不同的逐副本批次per-replica batch。把输入数据表示成单个jax.Array时该数组包含这一步所有副本的数据称为全局批次global batch其中每个分片恰好是一个逐副本批次。可以把它表示成跨所有设备的一维分片见下图——也就是说全局批次 所有逐副本批次沿批次轴拼接而成沿用这个框架你可能会得出结论进程0应获得全局批次的前四分之一8 个分片中的 2 个进程1获得第二个四分之一依此类推。但问题来了怎么知道前四分之一是什么又怎么确保进程0拿到的是前四分之一幸运的是数据并行有一个非常重要的性质让你根本不需要回答这些问题从而让整个设置大为简化。关于数据并行的关键性质这个关键性质是你不需要关心哪个逐副本批次落在哪个副本上因此也不在乎哪个进程加载哪个批次。原因是每个设备对应一个做同样事情的模型副本因此在全局批次内部哪个设备拿到哪个逐副本批次并不重要。这意味着你可以自由重排全局批次内的逐副本批次——换句话说你可以自由地随机化每个设备拿到哪个数据分片通常来说像上面这样重排jax.Array的数据分片并不是好主意——这实际上是在对数组的值做置换但对数据并行而言全局批次的顺序没有意义所以你可以如前所述自由重排全局批次中的逐副本批次。这个性质简化了数据加载每个设备只需要一个独立的逐副本批次数据流。大多数数据加载器都可以很容易地实现这一点——为每个进程创建一条独立的流水线再把得到的逐进程批次切成逐副本批次这正是前文合并的逐进程数据流水线方案的一个实例原文档在此处将其编号为 Option 2与上文的 Option 3 指同一方案编号在文档中存在细微出入你同样可以使用前文其他方案但该方案相对简单高效。下面是用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)这段代码的核心在最后一步jax.make_array_from_process_local_data(sharding, per_process_batch)根据给定的sharding把当前进程本地数据per_process_batch正确地放到该进程每个可寻址设备对应的分片上从而构造出全局jax.Array。从源码看该函数是make_array_from_callback的一个常见特例见 jax/_src/array.py它假定数据在进程内可用并替用户完成索引整理工作最常见的用法就是分片沿批次维展开、每个主机只加载自己的子批次同时它也支持多主机多轴复制与分片混合等更一般的情形此时需要自行正确计算进程本地数据的尺寸与内容以满足分片约束如果两个主机互为副本它们传入的本地数据必须完全一致。更底层的通用接口make_array_from_callback见 jax/_src/array.py则允许通过回调按全局索引取数据make_array_from_single_device_arrays见 jax/_src/array.py允许从一系列单设备数组组装全局数组。数据 模型并行Data model parallelism模型并行指把每个模型副本切分到多个设备上。若使用纯模型并行不含数据并行全局只有一个模型副本它被切分到所有设备上数据通常在所有设备上完全复制。本指南重点考虑同时使用数据并行与模型并行的情形把多个模型副本中的每一个都切分到多个设备上数据在每个模型副本上做部分复制——同一模型副本内的每个设备拿到相同的逐副本批次不同模型副本之间的设备拿到不同的逐副本批次。进程内模型并行Model parallelism within a process对数据加载而言最简单的做法是让每个模型副本都切分在单个进程的本地设备内。为了示例这里改为 2 个进程、每进程 4 个设备而不是之前的 4 进程 × 2 设备。考虑每个模型副本被切分到单个进程的 2 个本地设备上的场景结果是每进程 2 个模型副本、总共 4 个模型副本此时输入数据同样表示成单个jax.Array每个分片是一个逐副本批次的一维分片但有一个例外与纯数据并行不同这里引入了部分复制对一维分片的全局批次做 2 份副本原因在于每个模型副本由 2 个设备组成这 2 个设备各自都需要一份逐副本批次的副本。把每个模型副本保持在单个进程内会让事情更简单可以直接复用前面纯数据并行的设置只需额外把逐副本批次复制一份即可把逐副本批次复制到**正确的设备**上同样非常重要虽然前面关于数据并行的关键性质意味着你不关心哪个批次落到哪个副本上但**你必须保证单个副本只拿到单个批次**。例如下面这种加载方式是没问题的然而如果不注意把每个批次加载到哪个本地设备就可能意外制造出未复制的数据——尽管Sharding以及并行策略声称数据是复制的如果在一个进程内意外创建了本应复制却未复制的jax.ArrayJAX 会报错不过在跨进程模型并行时并不总能检测到见下一节。下面是用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 构造mesh_devices先按进程把设备分组再 reshape 成每个模型副本一行的形状并用断言确保每个副本的设备都在同一进程内相关辅助 APIjax.device_count()、jax.process_count()、jax.process_index()、jax.local_devices()的实现可参见 jax/_src/xla_bridge.pyPartitionSpec 只指定model_replicas轴数据沿模型副本轴切分而data_parallelism轴不参与切分因此每个逐副本分片都会沿该轴被复制make_array_from_process_local_data再次完成进程本地数据 → 全局分片布局的转换。跨进程模型并行Model parallelism across processes当模型副本跨进程分布时事情会变得更有意思。出现这种情况的原因可能是单个模型副本放不进单个进程或设备分配方式本就如此。回到之前的 4 进程 × 2 设备配置如果像下面这样把设备分配给副本这与前面进程内模型并行示例的并行策略完全相同——4 个模型副本、每个被切分到 2 个设备上。唯一不同的是设备分配每个副本的两个设备分属不同进程每个进程只为每个逐副本批次负责一份副本但要为 2 个副本负责。把模型副本这样跨进程切分看似随意且不必要在这个例子中确实如此但真实部署中可能会为了充分利用设备间的通信链路而采用这种设备分配方式。此时数据加载变得更复杂因为跨进程需要额外的协调。在纯数据并行和进程内模型并行的情况下只需保证每个进程加载一条唯一的数据流即可而现在某些进程必须加载相同的数据另一些进程必须加载不同的数据。在上面的例子中进程0粉色和进程2绿色必须加载相同的 2 个逐副本批次进程1黄色和进程3蓝色也必须加载相同的 2 个逐副本批次但与进程0、2的不同。此外还要确保每个进程不搞混自己的 2 个逐副本批次。虽然你不关心哪个批次落在哪个副本上数据并行的关键性质但必须保证一个副本内的所有设备拿到同一个批次。例如下面这种就是错误做法截至 2023 年 8 月JAX **无法检测**跨进程的 jax.Array 分片本应复制却未复制的情况运行计算时会直接产生错误结果。因此务必小心不要犯这种错误要保证每个设备拿到正确的逐副本批次需要把全局输入数据表示为下面这样的jax.Array小结如何为你的并行策略选择加载方案把整篇指南的核心决策路径收拢如下先确定并行策略纯数据并行 / 进程内模型并行 / 跨进程模型并行据此确定Sharding与 mesh 轴的定义再确定数据加载方案数据量小、图省事可选 Option 1每进程全量加载追求精确加载且设备数少可选 Option 2逐设备流水线追求最高效可选 Option 3合并的逐进程流水线难以精确加载时可选 Option 4方便加载 计算内with_sharding_constraint重分片利用数据并行的关键性质纯数据并行场景下批次顺序无意义每个进程独立加载一条数据流即可无需关心批次最终落到哪个副本警惕复制陷阱进程内复制错误 JAX 会报错跨进程复制错误则可能静默产生错误结果——务必按文档示例那样构造 mesh 并核对每个副本的设备归属。上述 API 与模式的正确性在仓库测试中有大量佐证例如 tests/multiprocess/array_test.py、tests/multiprocess/pjit_test.py 等测试文件均在实际多进程环境下演练了make_array_from_process_local_data、make_array_from_callback等分布式数组构造路径可作为深入阅读的参考。对于更完整的上下文还可对照阅读同一主题的 docs/501/data-loading.md本文档的上游版本与 docs/multi_process.md多进程环境基础。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表