ARTICLE DETAIL

资讯详情

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

JAX 分布式数组与自动并行化实战:从 `jax.Array`、`Sharding` 到 `jax.jit` 编译器的自动并行

JAX 分布式数组与自动并行化实战:从 `jax.Array`、`Sharding` 到 `jax.jit` 编译器的自动并行 机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载本指南基于 JAX 官方 Notebookdocs/notebooks/Distributed_arrays_and_automatic_parallelization.md整理而成系统讲解 JAX v0.4.1 引入的统一数组对象模型jax.Array如何使用PositionalSharding与NamedSharding把数据切分到多设备内存如何让jax.jit的计算跟随数据布局自动并行化计算以及如何在神经网络中组合数据并行与张量并行。读完本文你将掌握jax.device_put、Mesh/PartitionSpec、jax.lax.with_sharding_constraint的完整用法并理解随机数在并行环境下锐利边角的成因与规避方式。⚠️ 本文全部示例需要8 台设备GPU/TPU才能运行代码会在设备数不足时主动报错if len(jax.local_devices()) 8: raise Exception(Notebook requires 8 devices to run)。若在单卡环境运行建议配合 XLA 的虚拟多设备或 TPU 多核环境。引言与一个快速示例jax.Array是 JAX 中表示数组的统一数据类型即使物理存储横跨多台设备它依然是一个逻辑上的数组。把jax.Array与jax.jit结合使用可以获得基于编译器的自动并行化——无需手写任何设备通信代码编译器会根据数据布局自动把计算拆分到各设备上并行执行。先看一个完整的快速示例。首先用jax.experimental.mesh_utils.create_device_mesh按照硬件拓扑生成一个设备网格再用PositionalSharding描述数据如何在设备间分布import os import functools from typing import Optional import numpy as np import jax import jax.numpy as jnp from jax.experimental import mesh_utils from jax.sharding import PositionalSharding # 创建 Sharding 对象将值分布到 8 台设备上 sharding PositionalSharding(mesh_utils.create_device_mesh((8,))) # 生成一个随机数组 x jax.random.normal(jax.random.key(0), (8192, 8192)) # 使用 jax.device_put 将其分布到设备上 y jax.device_put(x, sharding.reshape(4, 2)) jax.debug.visualize_array_sharding(y)jax.debug.visualize_array_sharding实现在 jax/_src/debugging.py会以可视化网格的形式展示数组每个分片shard存放在哪台设备的内存里。对y应用元素级计算结果同样跨设备分布z jnp.sin(y) jax.debug.visualize_array_sharding(z)jnp.sin的求值被自动并行化输入值存储在哪台设备上输出值就分布在哪台设备上各设备并行计算自己的分片。用%timeit对比单设备与 8 设备分片的性能# x 只存在于单台设备 %timeit -n 5 -r 5 jnp.sin(x).block_until_ready() # y 被切分到 8 台设备 %timeit -n 5 -r 5 jnp.sin(y).block_until_ready()Sharding描述数组值在设备内存中的布局要并行化多设备计算第一步是把输入数据布局到多台设备上。在 JAX 中Sharding对象描述分布式内存布局它可以配合jax.device_put产出一个具有分布式布局的值。Sharding是一个接口interface任何实现该接口的类都可以与device_put等函数一起使用。仓库中常见的实现包括PositionalSharding、NamedSharding、GSPMDSharding等见 jax/_src/sharding_impls.py 与 jax/_src/sharding.py。PositionalSharding按位置描述布局默认创建的数组如jax.random.normal的返回值带有单设备Sharding所有数据都存放在一台设备上。切分它需要两步用mesh_utils.create_device_mesh((8,))创建Devices的numpy.ndarray。该函数会考虑硬件拓扑来决定Device的排列顺序因此设备编号往往不是按数值顺序排列的——底层网格反映的是设备的环形拓扑。基于该网格构造PositionalSharding并配合device_put使用from jax.experimental import mesh_utils devices mesh_utils.create_device_mesh((8,)) from jax.sharding import PositionalSharding sharding PositionalSharding(devices) x jax.device_put(x, sharding.reshape(8, 1)) jax.debug.visualize_array_sharding(x)PositionalSharding的行为像一个元素为设备集合的数组。PositionalSharding(ndarray_of_devices)固定了设备顺序和初始形状之后可以任意reshapesharding.reshape(8, 1) # 8 行 1 列的设备排列 sharding.reshape(4, 2) # 4 行 2 列的设备排列要在device_put中使用某个sharding需要把它 reshape 成与x.shape一致congruent的形状——即秩相同且每个维度都能整除x对应维度def is_congruent(x_shape: Sequence[int], sharding_shape: Sequence[int]) - bool: return (len(x_shape) len(sharding_shape) and all(d1 % d2 0 for d1, d2 in zip(x_shape, sharding_shape)))例如把shardingreshape 成(4, 2)再用于device_putsharding sharding.reshape(4, 2) print(sharding) y jax.device_put(x, sharding) jax.debug.visualize_array_sharding(y)这里的y与x表示同一个值只是它的分片切片被存放在不同设备的内存中。不同的PositionalSharding形状会产生不同的分布式布局例如reshape(1, 8)会沿第二维切成 8 份。用replicate表达复制有时我们不仅要切开存储还希望**复制replicate**某些分片——即把同一分片的值拷贝到多台设备的内存中。PositionalSharding通过归约方法replicate表达复制sharding sharding.reshape(4, 2) print(sharding.replicate(axis0, keepdimsTrue)) y jax.device_put(x, sharding.replicate(axis0, keepdimsTrue)) jax.debug.visualize_array_sharding(y)可视化会显示x沿第二维被切分为两份每份又在 4 台设备的内存中各存一份副本。从源码看jax/_src/sharding_impls.pyreplicate的实现是对内部_ids一个DeviceIdSet数组沿指定轴做集合并集self._ids.sum(axisaxis, keepdimskeepdims)DeviceIdSet重载了__add__为集合合并因此它语义上类似 NumPy 的归约方法.sum()/.prod()但与之不同的是keepdimsTrue是默认值被归约的轴不会被挤压掉print(sharding.replicate(0).shape) # (1, 2) print(sharding.replicate(1).shape) # (4, 1) y jax.device_put(x, sharding.replicate(1)) jax.debug.visualize_array_sharding(y)NamedSharding用名字表达布局PositionalSharding按轴位置描述布局而NamedSharding提供了更易读的命名方式先用Mesh定义设备网格并给网格轴命名再用PartitionSpec把数组轴映射到网格轴。from jax.sharding import Mesh from jax.sharding import PartitionSpec from jax.sharding import NamedSharding from jax.experimental import mesh_utils P PartitionSpec devices mesh_utils.create_device_mesh((4, 2)) mesh Mesh(devices, axis_names(a, b)) y jax.device_put(x, NamedSharding(mesh, P(a, b))) jax.debug.visualize_array_sharding(y)可以封装一个 helper 简化重复代码devices mesh_utils.create_device_mesh((4, 2)) default_mesh Mesh(devices, axis_names(a, b)) def mesh_sharding( pspec: PartitionSpec, mesh: Optional[Mesh] None, ) - NamedSharding: if mesh is None: mesh default_mesh return NamedSharding(mesh, pspec) y jax.device_put(x, mesh_sharding(P(a, b))) jax.debug.visualize_array_sharding(y)这里P(a, b)表示x的第一、二维分别切分在网格轴a、b上。切换为P(b, a)即可把两个轴映射到不同的设备排列y jax.device_put(x, mesh_sharding(P(b, a))) jax.debug.visualize_array_sharding(y)PartitionSpec中的None只是占位符用于对齐数组轴不表达任何切分而未被PartitionSpec提及的网格轴会得到复制# 这里的 None 表示 x 的第二维不切分 # 由于网格轴 b 未被提及分片会在 b 上复制。 y jax.device_put(x, mesh_sharding(P(a, None))) jax.debug.visualize_array_sharding(y)因为P(a, None)未提及网格轴b所以得到了沿b的复制。作为简写尾部None可以省略P(a, None)等价于P(a)但显式写出更清晰。只切分第二维可以在PartitionSpec第一维放Noney jax.device_put(x, mesh_sharding(P(None, b))) jax.debug.visualize_array_sharding(y) y jax.device_put(x, mesh_sharding(P(None, a))) jax.debug.visualize_array_sharding(y)固定网格时还可以把x的一个逻辑轴切分到多个网格轴上y jax.device_put(x, mesh_sharding(P((a, b), None))) jax.debug.visualize_array_sharding(y)NamedSharding的价值在于一次定义设备网格并命名网格轴之后每次device_put只需在PartitionSpec里引用名字即可语义一目了然。PartitionSpec是tuple的子类见 jax/_src/partition_spec.py其元素既可以是网格轴名字符串也可以是网格轴名的元组表示多维切分。计算跟随数据 sharding并被自动并行化有了分片输入数据编译器就能给出并行计算。特别是被jax.jit装饰的函数可以直接作用于分片数组而无需先把数据拷贝回单台设备。计算跟随 sharding基于输入数据的 sharding编译器决定中间值和输出值的 sharding并行化它们的求值必要时自动插入通信操作。以最简单的元素级计算为例from jax.experimental import mesh_utils from jax.sharding import PositionalSharding sharding PositionalSharding(mesh_utils.create_device_mesh((8,))) x jax.device_put(x, sharding.reshape(4, 2)) print(input sharding:) jax.debug.visualize_array_sharding(x) y jnp.sin(x) print(output sharding:) jax.debug.visualize_array_sharding(y)对元素级jnp.sin编译器选择的输出 sharding 与输入一致并把计算拆分成每个设备只算自己那一份——尽管我们写的jnp.sin仿佛是单机执行编译器却帮我们拆分并分发到了多台设备。不只是元素级操作。考虑输入已分片的矩阵乘法y jax.device_put(x, sharding.reshape(4, 2).replicate(1)) z jax.device_put(x, sharding.reshape(4, 2).replicate(0)) print(lhs sharding:) jax.debug.visualize_array_sharding(y) print(rhs sharding:) jax.debug.visualize_array_sharding(z) w jnp.dot(y, z) print(out sharding:) jax.debug.visualize_array_sharding(w)这里编译器选择的输出 sharding 能最大化并行度每台设备已经持有计算自身输出分片所需的输入分片无需任何通信。如何确认它真的在并行做一个简单的计时实验先把数据放到单台设备再对比分片版本x_single jax.device_put(x, jax.devices()[0]) jax.debug.visualize_array_sharding(x_single) np.allclose(jnp.dot(x_single, x_single), jnp.dot(y, z)) %timeit -n 5 -r 5 jnp.dot(x_single, x_single).block_until_ready() %timeit -n 5 -r 5 jnp.dot(y, z).block_until_ready()此外拷贝一个已分片的Array结果仍保持输入的分片布局w_copy jnp.copy(w) jax.debug.visualize_array_sharding(w_copy)总结该策略计算跟随数据放置。当我们用jax.device_put显式分片数据并对其应用函数时编译器会尝试并行化计算并决定输出 sharding。这是 JAX遵循显式设备放置策略的推广——相关内容可参见官方 FAQ 中Controlling data and computation placement on devices一节。显式 sharding 冲突时JAX 报错如果一次计算的两个参数被显式放置到不同的设备集合或使用了不兼容的设备顺序在这些有歧义的情况下会直接报错import textwrap from termcolor import colored def print_exception(e): name colored(f{type(e).__name__}, red) print(textwrap.fill(f{name}: {str(e)})) # 情形一两个参数被放在互不重叠的两组设备上 sharding1 PositionalSharding(jax.devices()[:4]) sharding2 PositionalSharding(jax.devices()[4:]) y jax.device_put(x, sharding1.reshape(2, 2)) z jax.device_put(x, sharding2.reshape(2, 2)) try: y z except ValueError as e: print_exception(e) # 情形二设备集合相同但顺序被打乱 devices jax.devices() permuted_devices [devices[i] for i in [0, 1, 2, 3, 6, 7, 4, 5]] sharding1 PositionalSharding(devices) sharding2 PositionalSharding(permuted_devices) y jax.device_put(x, sharding1.reshape(4, 2)) z jax.device_put(x, sharding2.reshape(4, 2)) try: y z except ValueError as e: print_exception(e)用jax.device_put显式放置或分片的数组被称为已提交committed到其设备因此不会被自动搬移。相对地未经jax.device_put显式放置的数组被未提交uncommitted地放在默认设备上与已提交数组不同未提交数组可以被自动搬移和重新分片——即使计算中其他参数被显式放在不同设备上未提交数组也能作为参数参与运算。例如jnp.zeros、jnp.arange、jnp.array的输出都是未提交的y jax.device_put(x, sharding1.reshape(4, 2)) y jnp.ones_like(y) y jnp.arange(y.size).reshape(y.shape) print(no error!)在jit代码中约束中间值的 sharding编译器会尽量自行决定函数中间值与输出的 sharding但我们也可以用jax.lax.with_sharding_constraint给出提示。它的用法与jax.device_put很相似区别在于它用于被 staged-out即被jit装饰的函数内部sharding PositionalSharding(mesh_utils.create_device_mesh((8,))) x jax.random.normal(jax.random.key(0), (8192, 8192)) x jax.device_put(x, sharding.reshape(4, 2)) jax.jit def f(x): x x 1 y jax.lax.with_sharding_constraint(x, sharding.reshape(2, 4)) return y jax.debug.visualize_array_sharding(x) y f(x) jax.debug.visualize_array_sharding(y)也可以把输出约束为完全复制jax.jit def f(x): x x 1 y jax.lax.with_sharding_constraint(x, sharding.replicate()) return y jax.debug.visualize_array_sharding(x) y f(x) jax.debug.visualize_array_sharding(y)通过with_sharding_constraint我们约束了输出的 sharding。除了尊重某个特定中间值的注解编译器还会利用这些注解去决定其他值的 sharding。实践中好的做法是依据结果最终如何被消费去注解计算的输出。实战示例神经网络中的数据并行与张量并行⚠️ 以下内容只是用jax.Array演示自动 sharding 传播的简单示例未必代表真实项目的最佳实践——真实场景可能需要更多地使用with_sharding_constraint。我们可以利用jax.device_put与jax.jit的计算跟随 sharding特性并行化神经网络计算。基于下面这个基本网络import jax import jax.numpy as jnp def predict(params, inputs): for W, b in params: outputs jnp.dot(inputs, W) b inputs jnp.maximum(outputs, 0) return outputs def loss(params, batch): inputs, targets batch predictions predict(params, inputs) return jnp.mean(jnp.sum((predictions - targets)**2, axis-1)) loss_jit jax.jit(loss) gradfun jax.jit(jax.grad(loss)) def init_layer(key, n_in, n_out): k1, k2 jax.random.split(key) W jax.random.normal(k1, (n_in, n_out)) / jnp.sqrt(n_in) b jax.random.normal(k2, (n_out,)) return W, b def init_model(key, layer_sizes, batch_size): key, *keys jax.random.split(key, len(layer_sizes)) params list(map(init_layer, keys, layer_sizes[:-1], layer_sizes[1:])) key, *keys jax.random.split(key, 3) inputs jax.random.normal(keys[0], (batch_size, layer_sizes[0])) targets jax.random.normal(keys[1], (batch_size, layer_sizes[-1])) return params, (inputs, targets) layer_sizes [784, 8192, 8192, 8192, 10] batch_size 8192 params, batch init_model(jax.random.key(0), layer_sizes, batch_size)8 路批量数据并行PositionalSharding(jax.devices()).reshape(8, 1)把批维度切到 8 台设备参数则完全复制到每台设备sharding PositionalSharding(jax.devices()).reshape(8, 1) batch jax.device_put(batch, sharding) params jax.device_put(params, sharding.replicate()) loss_jit(params, batch) step_size 1e-5 for _ in range(30): grads gradfun(params, batch) params [(W - step_size * dW, b - step_size * db) for (W, b), (dW, db) in zip(params, grads)] print(loss_jit(params, batch)) %timeit -n 5 -r 5 gradfun(params, batch)[0][0].block_until_ready() # 对照全部放到单台设备 batch_single jax.device_put(batch, jax.devices()[0]) params_single jax.device_put(params, jax.devices()[0]) %timeit -n 5 -r 5 gradfun(params_single, batch_single)[0][0].block_until_ready()jax.device_put对 pytree 是递归生效的batch被切成 8 份params中的每个张量都被复制到全部 8 台设备。4 路批量数据并行 2 路模型张量并行把shardingreshape 成(4, 2)就可以混合两种并行模式批维度切成 4 份同时把某些权重矩阵沿其维度切成 2 份张量并行。sharding sharding.reshape(4, 2) # 数据并行batch 沿第一维切 4 份第二维复制 batch jax.device_put(batch, sharding.replicate(1)) jax.debug.visualize_array_sharding(batch[0]) jax.debug.visualize_array_sharding(batch[1]) (W1, b1), (W2, b2), (W3, b3), (W4, b4) params # W1、b1、b3、b4 保持全复制 W1 jax.device_put(W1, sharding.replicate()) b1 jax.device_put(b1, sharding.replicate()) # W2 沿 axis0输出维切成 2 份张量并行 W2 jax.device_put(W2, sharding.replicate(0)) b2 jax.device_put(b2, sharding.replicate(0)) # W3 转置后沿 axis0 切 2 份等价于沿原 axis1 切分 W3 jax.device_put(W3, sharding.replicate(0).T) b3 jax.device_put(b3, sharding.replicate()) W4 jax.device_put(W4, sharding.replicate()) b4 jax.device_put(b4, sharding.replicate()) params (W1, b1), (W2, b2), (W3, b3), (W4, b4) jax.debug.visualize_array_sharding(W2) jax.debug.visualize_array_sharding(W3) print(loss_jit(params, batch)) step_size 1e-5 for _ in range(30): grads gradfun(params, batch) params [(W - step_size * dW, b - step_size * db) for (W, b), (dW, db) in zip(params, grads)] print(loss_jit(params, batch)) (W1, b1), (W2, b2), (W3, b3), (W4, b4) params jax.debug.visualize_array_sharding(W2) jax.debug.visualize_array_sharding(W3) %timeit -n 10 -r 10 gradfun(params, batch)[0][0].block_until_ready()其中W3使用sharding.replicate(0).T先把设备网格沿第 0 轴复制形状(1, 2)再转置得到(2, 1)这样W3形状(8192, 8192)的第二维被切成两份正好与切了第一维的W2匹配构成矩阵乘法链路中的张量并行。整个前向与反向传播过程中编译器自动在切分的批维度上做数据并行、在切分的权重维度上做张量并行并自动插入必要的跨设备归约。锐利边角随机数生成JAX 自带函数式、确定性的随机数生成器它是jax.random模块如jax.random.uniform等采样函数的基础。JAX 的随机数由基于计数器的 PRNG 产生因此原则上随机数生成应该是计数器值的纯映射——纯映射在原则上是可以平凡切分的操作既不需要跨设备通信也不需要设备间冗余计算。然而现存的稳定 RNG 实现由于历史原因不能自动切分。考虑下面这个例子一个函数生成均匀随机数并与输入按元素相加jax.jit def f(key, x): numbers jax.random.uniform(key, x.shape) return x numbers key jax.random.key(42) x_sharding jax.sharding.PositionalSharding(jax.devices()) x jax.device_put(jnp.arange(24), x_sharding)在分片输入上函数f的输出也是分片的jax.debug.visualize_array_sharding(f(key, x))但如果检查该分片输入上编译后的计算会发现其中包含通信collective-permutef_exe f.lower(key, x).compile() print(Communicating?, collective-permute in f_exe.as_text())一种规避方式是把实验性升级开关jax_threefry_partitionable打开。开启后编译计算中的collective-permute会消失jax.config.update(jax_threefry_partitionable, True) f_exe f.lower(key, x).compile() print(Communicating?, collective-permute in f_exe.as_text())输出依然是分片的jax.debug.visualize_array_sharding(f(key, x))jax_threefry_partitionable有一个重要告诫开启后生成的随机数值可能与未开启时不同即使使用同一个随机 keyjax.config.update(jax_threefry_partitionable, False) print(Stable:) print(f(key, x)) print() jax.config.update(jax_threefry_partitionable, True) print(Partitionable:) print(f(key, x))在jax_threefry_partitionable模式下JAX PRNG 依然是确定性的但它的实现是全新的且仍在开发中给定 key 生成的随机值在同一 JAX 版本或main分支的同一 commit内保持不变但可能随版本发布而变化。小结jax.Array是 JAX v0.4.1 起的统一数组模型即使物理存储横跨多台设备逻辑上仍是单个数组。Sharding是描述分布式内存布局的接口PositionalSharding按轴位置描述支持reshape、transpose、replicate集合并集归约默认keepdimsTrueNamedSharding配合Mesh与PartitionSpec用名字描述布局未被PartitionSpec提及的网格轴自动复制。jax.device_put用Sharding放置数据显式放置的数组被提交到设备不会自动搬移未提交数组则可以被自动搬移与重新分片显式 sharding 冲突时 JAX 报ValueError。计算跟随数据放置jax.jit编译的函数作用于分片输入时编译器自动决定中间值与输出的 sharding、并行化计算并插入必要的通信jax.lax.with_sharding_constraint可在jit内部约束中间值布局。通过给数据和权重设置不同的Sharding可组合出数据并行与张量并行的混合并行方案。随机数在并行环境下存在历史性的通信开销可用实验性开关jax_threefry_partitionable消除通信但需注意生成数值可能随版本变化。建议进一步阅读仓库中的关联资料JAX PRNG 设计文档、分布式数据加载指南、shard_map 教程以及 jax.sharding 模块文档 获取完整的 API 参考。赞分享机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载相关推荐JAX 分布式数组与自动并行化实战Mesh、Sharding 与三种并行模式全解析JAX 分布式数组与自动并行化实战Mesh、Sharding 与三种并行模式全解析 在 JAX 中分布式并行既可以是编译器替你全自动完成的也可以是由你人工智能机器学习深度学习编译器高性能计算JAX 分布式数组与自动并行化Mesh、Sharding 与三种并行模式的完整指南JAX 分布式数组与自动并行化Mesh、Sharding 与三种并行模式的完整指南 导读 本文是 JAX 官方 docs/parallel.md 文档的系统化人工智能机器学习深度学习编译器高性能计算JAX 分布式计算入门从数据 sharding 到 SPMD 并行jit 自动并行 / with_sharding_constraint / shard_map 全解析JAX 分布式计算入门从数据 sharding 到 SPMD 并行jit 自动并行 / with_sharding_constraint / shard_m机器学习深度学习上一篇GitHub_Trending/ra/rag-from-scratch历史版本回顾项目发展的关键里程碑下一篇copymanga漫画下载免费工具实战五关解锁批量下载与离线阅读创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表