ARTICLE DETAIL

资讯详情

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

从vmap到jax.sharding:JAX硬件级并行范式实战指南

从vmap到jax.sharding:JAX硬件级并行范式实战指南 一说起 JAX很多人的第一反应就是自动微分grad一行替换backward把梯度算得干干净净。但如果你只在 JAX 里用grad那基本等于进了宝山只拿了个塑料袋。我真正从 PyTorch 迁徙过来、折腾了一年之后觉得“回不去”的原因是 JAX 的并行计算 API——vmap、pmap、再到jax.sharding这套体系。它提供的不是一个两个便利函数而是一种硬件级并行范式你的 Python 函数会在编译期被改写成带显式设备布局和通信模式的计算图再交给 GPU / TPU 执行。同一份代码可以从 CPU 跑到单卡 GPU再跑到多机集群并行方式甚至可以让编译器帮你规划。这篇想聊的就是这套范式怎么理解、怎么落地以及我在实操中踩过的坑。适合正在做模型训练、推理加速、或者被多卡分布式搞到头大的工程师和 ML 研究者。1. 为什么 JAX 能成为“硬件级并行范式”的样板1.1 从 grad 到 jit先把 JAX 的本质搞清楚很多人以为 JAX 是个“带自动微分的 NumPy”这个理解没错但它严重低估了背后的东西。准确说JAX NumPy 风格的数值接口 基于 XLA 的编译后端 自动微分 可组合的函数变换 自动并行规划。这四个东西缺一不可但真正和其他框架拉开差距的是“编译”这个环节。当你调用jax.jit(f)时JAX 不会立刻执行f里的 Python 代码而是先把它“trace”成一份内部中间表示jaxpr再交给 XLA 做算子融合、内存规划、设备分配、通信插入最后生成一个可执行程序。这个流程意味着你的并行策略不是在运行时临时调度出来的而是在编译期就被确定下来的。这一点和 PyTorch 的 eager 执行有本质区别——PyTorch 的每一步算子都是运行时逐发到设备上执行多卡并行要靠torch.distributed在 Python 侧手动编排通信原语计算和通信之间经常需要你自己找同步点。我把这个区别用一个表格说清楚维度JAXPyTorch执行模式先编译后执行XLAeager 即时执行并行决策时机编译期由 SPMD 分区器决定运行期由 Python 运行时调度多卡通信编排编译器根据 sharding 自动插入开发者手动调用 allreduce 等原语代码形态逻辑上保持单设备视角设备布局是额外注解分布式代码中通信逻辑和计算逻辑混在一起新硬件适配有对应 XLA 后端即可框架层不重写通常需要写新的 kernel / 后端适配层所以 JAX 的并行 API 不是“在 PyTorch 并行方案上换个皮”它是一整套从语言到编译器的思路转变你写的是“全局、单设备、无通信”的逻辑然后在关键张量上声明“我希望它怎么分布”剩下的交给 XLA 去推断、补通信、调优。1.2 “硬件级”的含义不是 API 的并行是编译器的并行“硬件级并行”这个说法听上去很玄其实就是指JAX 把现代硬件上几个不同粒度的并行维度统一映射成了同一套可组合的函数变换。硬件上的并行大体分三层指令级并行ILP单个核心内部多条指令乱序执行或同时执行向量 / 线程级并行SIMD / SIMTGPU 上的一个 warp、CPU 上的一个 SIMD 寄存器组对一批数据同时做同一个操作多核 / 多设备级并行MIMD多块 GPU、多个 TPU 核、多台机器各算一部分再加通信。一般框架能把这三层都照顾到就已经不容易了JAX 的特别之处在于它在 API 层面用统一语义把这三层串起来。vmap对应向量 / 线程级并行pmap/jax.sharding对应多设备级并行XLA 的算子融合和循环优化则在指令级尽量把开销压掉。你可以把vmap和pmap嵌套使用也可以和jit组合因为它们本质都是“函数变换”变换可以叠加。这套东西能叫“范式”而不只是“工具集”核心在于GSPMDGeneralized SPMD机制。你写代码时只有一个“全局视角”比如“我有一批 1024 个样本我要把它们过一遍矩阵乘”然后你在PartitionSpec里说“batch 维按 data 轴切到 4 张卡上”XLA 的分区器会自动把整个前向和反向传给重写把必要的 all-gather、all-reduce、reshard 原语插入到合适的位置。最后在你的感知里代码仍然像单机程序一样干净。用一个生活化类比你坐在办公室说“这一万张发票我要今晚算出税额”这是全局逻辑DSP 团队、票据扫描团队、计算团队怎么分工、谁汇总给谁是细节有经验的管理者在收到指令的一瞬间会在脑内分派好任务流。GSPMD 就是这个“有经验的管理者”你只需要把资源清单和分工偏好告诉它。2. 并行计算 API 全家桶从 vmap 到 pmap 再到 sharding2.1 vmap把循环变成硬件批量化的开关vmap是很多人接触到 JAX 并行 API 的第一站。它解决的是最朴素的问题不要在手写 batch 循环了把“对单个样本做计算”的函数直接变成“对一批样本做计算”的函数并且让这批样本的计算在硬件层面尽可能并行。import jax.numpy as jnp from jax import vmap, jit def single(x, W): # x: (feat,), W: (out, feat) return jnp.tanh(W x) # 把第 0 维作为 batch 维W 保持共享 batched vmap(single, in_axes(0, None)) x_batch jnp.ones((128, 64)) W jnp.ones((32, 64)) y_batch batched(x_batch, W) # y_batch shape: (128, 32)关键在in_axes0表示这个参数的第 0 维是 batch 维None表示这个参数在所有 batch 里共享。out_axes控制输出结果的 batch 维放哪。默认都是 0但多输出或者多维 batch 时一定要检查。实际工程中vmap几乎总是和jit一起用。为什么因为vmap本身只是把你这个函数里的算子按 batch 维“展开改写”一遍真正把批量维度映射到 SIMD/SIMT 指令上的是 XLA 编译器。如果你只vmap不jit性能提升通常很有限有时还更慢jit(vmap(f))才会让编译器看到完整展开后的计算图然后做算子融合和向量化。我见过很多人犯一个错写成vmap(jit(f))觉得 “我要 keep 一个 jitted 的单样本函数”。这个顺序很微妙。你还是可以得到正确的计算结果但每个样本都会被单独编译成一个 kernel 调用批量维没有真正融合进计算图性能就浪费了。正确顺序基本是jax.jit(vmap(f))让 vmap 先做函数变换、jit 再做编译。使用vmap的另一个注意点是内存访问模式。如果 batch 维对应的底层连续性和硬件访问模式差异很大vmap展开后可能产生不连续的地址访问反而拖慢速度。遇到这种情况不要急着怀疑 vmap 没用先看 profile 里的访存命中率再考虑调整数据布局比如把 batch 维放最后。2.2 pmap曾经的多设备数据并行入口pmap是 JAX 早年为多设备设计的主力 API。它把一个函数复制到多个设备上每个设备处理不同分片的数据执行过程中如果你在函数内部声明了axis_nameJAX 会在对应位置自动插入集合通信原语比如lax.psum、all_gather。import jax from jax import pmap import jax.numpy as jnp def f(x): return jnp.sum(x, axis0) p_f pmap(f, axis_namebatch) out p_f(jnp.ones((8, 16))) # 8 个设备各处理 (1, 16)结果按设备堆叠shape (8, 16)这段代码在执行时会把(8, 16)按第 0 维切成 8 份每个设备算自己的(1, 16)的 sum结果收集回来后拼成(8, 16)的完整输出。看起来很方便但它有几个硬伤设备数和 batch 长度强绑定batch 必须能被设备数整除通信操作靠axis_name手动插代码里到处散落着lax.psum能读但不好维护一旦要混合模型并行和数据并行pmap的嵌套写法相对晦涩JAX 内部后来把pmap的实现逐步统一到 “自动分片” 的路线上新特性基本都在jax.sharding里。所以我的建议是老项目里的pmap能跑就留着新代码一律用下一节讲的jax.sharding。pmap不坏只是这套 API 的思想还是“人肉控制并行”而 JAX 后来的方向是“声明式并行 编译器代劳”。2.3 jax.sharding 和 NamedSharding声明式并行的新一代 APIjax.sharding把“并行”的概念彻底抽象成了三个组件设备网格Mesh、切分规范PartitionSpec、以及把两者绑在一起的NamedSharding。Mesh定义设备之间的网格拓扑from jax.sharding import Mesh import jax devices jax.devices() # 例如 4 块 GPU mesh Mesh(devices.reshape((2, 2)), (data, model))这里把 4 个设备排成 2x2 的网格两个逻辑轴分别叫data和model。你可以理解为横轴负责数据并行纵轴负责模型并行。PartitionSpec描述一个数组的各维度如何映射到 Mesh 轴from jax.sharding import PartitionSpec, NamedSharding # 一个 shape 为 (batch, hidden) 的数组batch 维按 data 切分hidden 维不切分 data_sharding NamedSharding(mesh, PartitionSpec(data, None)) # 一个 shape 为 (out, feat) 的权重矩阵两个维都不切分每设备完整复制 replicated_sharding NamedSharding(mesh, PartitionSpec())PartitionSpec里的位置对应数组维度序号第 0 个元素描述数组第 0 维如何切分第 1 个元素描述数组第 1 维如何切分None表示这一维保持完整即广播复制到所有相关设备字符串则对应 Mesh 中的轴名。如果整个数组都复制就传空元组。使用方式是在jit时通过in_shardings和out_shardings声明输入输出分布from jax.sharding import NamedSharding, PartitionSpec, Mesh jax.jit(in_shardings(data_sharding, replicated_sharding), out_shardings(data_sharding,)) def train_step(params, x_batch): def loss_fn(p): pred p x_batch return jnp.mean(pred ** 2) grads jax.grad(loss_fn)(params) new_params jax.tree.map(lambda p, g: p - 0.01 * g, params, grads) return new_params, pred这里我说一个新版 JAX 可以简化的问题代码里 params 本身是个 pytree里面可能有 embedding、linear 的 weight 和 bias 等不同类型的 tensor它们在同一个PartitionSpec下不一定都适用。所以实际工程里更稳健的做法是给每个参数分别指定 shardingparams_sharding jax.tree_util.tree_map( lambda x: NamedSharding(mesh, PartitionSpec(None, model)), params)把PartitionSpec(None, model)应用到所有形状为 2 维的参数上意思是第一个维度复制、第二个维度沿model轴切分。如果某个参数是 bias形状只有一维就得单独处理。这类“按形状分策略”的逻辑是 JAX 代码里最常见的样板。2.4 和 PyTorch DTensor 的对比很多人问我PyTorch 2.x 不也有 DTensor 吗不也能声明式并行吗对但两者设计起点不同。DTensor 是在已有 eager 体系和torch.distributed之上加一层“张量分布元数据”是典型的运行时解释方案JAX 的 sharding 是编译期方案XLA 的 GSPMD 分区器会看到完整数据流能对通信做全局融合与调度优化。另一个实际差别是体验用 DTensor 时你通常还要理解 pytorch 的分布式运行时、NCCL 进程组、自动微分引擎等一整套东西JAX 里你在逻辑设备的角度用PartitionSpec描述一遍剩下的错误多数能在编译期报出来。表格对比对比项JAX shardingPyTorch DTensor决策时机编译期运行期通信插入编译器自动运行时在算子执行中触发编程模型全局逻辑 分布注解全局逻辑 DTensor 分布元数据与 eager 的兼容不兼容走 XLA 编译兼容大部分 eager 业务学习曲线需要理解 jit/Mesh/Spec需要理解分布式运行时 / NCCL3. 实操从单卡到四卡搭一个并行训练脚本3.1 环境准备与设备感知安装 JAX 本身不复杂但不同硬件差别很大。CPU 环境直接pip install -U jaxGPU 环境建议安装带 CUDA 依赖的后端pip install -U jax[cuda12]如果你用的是 TPU在 Colab 或自己的 TPU 虚拟机里一般已经预装或者用pip install -U jax[tpu]安装完先确认设备可见import jax print(jax.devices()) print(jax.local_device_count())这一步很重要。很多人后续 sharding 写好了却发现设备列表和预期不一致原因往往是数据并行网格形状和设备总数对不上或者本地进程只看到了部分设备。多机场景还要留意jax.local_device_count()返回的是当前进程可见设备数和jax.device_count()有区别。3.2 单卡基线jit vmap 的 MLP 训练循环先写一个不依赖任何分布式概念的 MLP 训练骨架把基线跑稳再往上加并行。这里我用最简单的 MSE 回归优化器用手写 SGDimport jax import jax.numpy as jnp from jax import random, jit, grad, vmap def init_mlp(key, layer_dims): params [] for din, dout in zip(layer_dims[:-1], layer_dims[1:]): key, subkey random.split(key) w random.normal(subkey, (din, dout)) * 0.1 b jnp.zeros(dout) params.append((w, b)) return params def forward(params, x): for w, b in params: x jnp.tanh(x w b) return x jit def train_step(params, x_batch, y_batch): def loss_fn(p): pred forward(p, x_batch) return jnp.mean((pred - y_batch) ** 2) grads grad(loss_fn)(params) new_params jax.tree.map( lambda p, g: p - 0.01 * g, params, grads) return new_params, loss_fn(params) key random.PRNGKey(0) params init_mlp(key, [64, 128, 1]) x_batch random.normal(key, (256, 64)) y_batch random.normal(key, (256, 1)) for step in range(100): params, loss train_step(params, x_batch, y_batch) if step % 20 0: print(step, loss)注意这里forward内部是逐样本向量运算batch 维在 x_batch 的第 0 维上算子天然支持 batch 矩阵乘所以单个 kernel 本身就处理了数据并行。更常见的是你的函数写成了单样本逻辑def forward_single(params, x): for w, b in params: x jnp.tanh(x w b) return x forward_batched jit(vmap(forward_single, in_axes(None, 0)))然后就可以在任意调用点直接入 batch 数据。性能上如果模型小巧、算子少vmap展开后的图大概率会和一个手动写的 batched 版本差不多但如果你的单样本逻辑里有很多条件分支和动态 shapevmap的收益可能不明显压缩分支逻辑后再试才是正路。3.3 多卡数据并行Mesh NamedSharding 实战现在把上面的训练函数改造成四卡数据并行。核心思路非常朴素batch 维切到所有设备上模型参数每台设备都复制一份梯度经过跨设备求和后更新到同一组参数。我推荐在代码里明确写出 Mesh 和两类 shardingfrom jax.sharding import Mesh, PartitionSpec, NamedSharding import jax devices jax.devices() mesh Mesh(devices, (data,)) # 一维网格数据并行轴 # 数据按 batch 维切分 data_sharding NamedSharding(mesh, PartitionSpec(data, None)) # 参数全部复制 replicated NamedSharding(mesh, PartitionSpec()) jax.jit(in_shardings(replicated, data_sharding, data_sharding), out_shardings(replicated,)) def train_step_dp(params, x_batch, y_batch): def loss_fn(p): pred forward(p, x_batch) return jnp.mean((pred - y_batch) ** 2) grads grad(loss_fn)(params) # 梯度在设备间自动 all-reduce这里不完全对 # 因为 grads 的 sharding 继承自 paramsreplicated # XLA 会把梯度约减后同步到所有副本。 new_params jax.tree.map( lambda p, g: p - 0.01 * g, params, grads) return new_params, loss_fn(params)这里有个容易误解的点grads的 sharding 会跟随params的 sharding也就是每个设备本地算出一个完整但独立的部分梯度。为了让四份梯度一致XLA 会在反向结束时插入跨设备all-reduce。这是数据并行标准做法你不用手写但心里要清楚通信发生在哪一步。真正要让 sharding 起效数据必须先放到正确的设备布局上。JAX 里有一种做法是from jax import device_put x_device device_put(x_batch, data_sharding) y_device device_put(y_batch, data_sharding) params_device device_put(params, replicated) # 之后就能直接喂给 jitted 函数 params_device, loss train_step_dp(params_device, x_device, y_device)也可以用jax.make_array_from_callback或jax.random的randn配合sharding直接生成分布好的数据。如果你不device_putJAX 会尝试自动使用默认的“全部复制” sharding你的in_shardings和实际输入不匹配时就可能报错或者被静默转成复制性能就没了。执行前用可视化排查 sharding 是最省时间的习惯jax.debug.visualize_sharding(params_device, params sharding) jax.debug.visualize_sharding(x_device, x sharding)这段 code 不是必须写进生产但调试时几乎必备。3.4 模型并行与混合并行内存不够时怎么切数据并行解决的是吞吐量问题但碰到单卡显存放不下模型参数时就必须模型并行。JAX 的模型并行在声明式范式下同样简洁把权重矩阵的某个维度沿model轴切分。以一个简单线性层为例权重W的形状是(out, in)。如果按列切分也就是PartitionSpec(None, model)那么每个设备只持有权重的部分列。矩阵乘法x W在计算前XLA 可能需要先做all-gather把完整权重拼出来或者采取更聪明的split-K类算法。如果按行切分也就是PartitionSpec(model, None)那么每个设备只持有一部分输出行计算x W_local得到部分输出再对输出做all-reduce求和。这两种切法各有通信量差异。我给出一个按列切分权重的示例mesh Mesh(jax.devices().reshape((1, jax.local_device_count())), (data, model)) weight_spec PartitionSpec(None, model) weight_sharding NamedSharding(mesh, weight_spec) jax.jit(in_shardings(replicated, weight_sharding), out_shardings(replicated,)) def mlp_forward(params, x): # params 里的 weight 是沿 model 轴切分的 # 返回结果时由于输出维对应 weight 的行维所以输出也复制到每台设备 return forward(params, x)混合并行就是把data和model两个轴放到同一个 Mesh 里。比如 8 块 GPU 排成 2x4 网格batch 维沿data轴切成 2 份权重沿model轴切成 4 份同时获得数据并行吞吐和模型并行省显存的效果。这种写法在 JAX 里只是多几个PartitionSpec的事但通信开销也会叠加。千万别觉得 “编译器自动插入” 就是“无成本”实际 profiling 时通信往往是大头。3.5 大模型场景这套 API 在现代 LLM 里怎么用现在大家天天在业务层调用各种大模型 API但很少有人关注底层训练和推理框架是怎么在成千上万块加速卡上把模型组织起来的。JAX 这套并行 API 是底层框架层的答案之一很多开源 LLM 的 TPU / GPU 实现就是拿它写的。LLM 的并行化一般会叠加三种切法多头注意力的 head 维可以按model轴切分不同设备算不同 head再拼接FFN 的两个线性层权重可以分别按行 / 列切分配合all-reduce完成一次完整计算推理阶段的 KV cache 也可以按 batch 维 head 维同时切片降低单卡显存压力。在 JAX 里你不需要手写通信原语只需要给每个参数想清楚PartitionSpec怎么填。但你要知道每种切法对应的通信模式切法中间产物典型通信权重列切部分列权重参与局部矩阵乘AllGather 或 Split-K权重行切局部结果需要跨设备求和AllReducebatch 数据切每设备独立计算前向 / 反向梯度 AllReduce这段时间很多争论“自动分片会不会比手工分布式慢”我的实测是对于规则模型GSPMD 自动分片通常能逼近手工优化水平但不会自动超过。你要做的是用 profile 工具找出通信热点再手工补一两条 sharding 约束。3.6 性能调参经验batch 大小、通信重叠、编译时间并行程序的性能受三个因素控制计算效率、通信开销、同步等待。JAX 这套体系下最容易忽略的是 batch 大小和 mesh 切分的整除关系。如果你总 batch 是 1000设备数是 4那每台设备拿到 250 还好要是设备数是 3JAX 会自动做 padding内存和计算都会有一点浪费。宁可把 batch 调成设备数的整数倍。通信和计算重叠方面XLA 在常规矩阵乘 all-reduce 场景会自动尝试 overlap但复杂数据流里它不一定总能做到。手动优化时可以把一个大 batch 分成几个 micro-batch让第一个 micro-batch 的反向通信和第二个 micro-batch 的前向计算交叠——这在 JAX 里用lax.scan或显式 loop 都能做到代价是代码变复杂。编译时间也是 JAX 新手觉得很痛的点。首次jit一个大模型XLA 可能编译几分钟。解决办法是把模型拆成多个jit函数或者用jax.jit的缓存机制让常用 shape 只编译一次调试阶段可以先关jit用 eager 跑小规模验证确认逻辑没问题再开编译。真跑大模型时编译时间在整个训练周期里通常可接受。4. 常见问题与排查技巧实录4.1 vmap 后反而更慢是哪里不对劲我在 2.1 提过vmap必须和jit配合。除此之外还有一个常见原因vmap展开后的算子导致中间张量形状变化内存访问不连续。排查方法很简单用 profiler 看单个内核的执行时间如果发现大量小 kernel 而不是一个融合的大 kernel就说明编译器没有把处理逻辑合并起来。还有一个我踩过好几次的坑在一个已经有 batch 维的数据上再vmap一次也就是错误地给 batched 逻辑加了多余的外层循环性能直接掉一个量级。检查一下你的输入形状和in_axes到底匹配不匹配。4.2 sharding mismatch 报错Array 的分布和预期不符运行时经常见到这样的报错大意是某个jax.Array的 sharding 与你jit里声明的in_shardings不一致。原因通常有两种。第一种你在device_put时用的 sharding 和jit里声明的不一致第二种函数内部reshape/squeeze/ 转置改变维度数量或顺序导致后续张量的PartitionSpec对不上。解决方法是先jax.debug.visualize_sharding打印所有关键张量的分布看到底哪一层开始偏的然后在偏的位置显式插入jax.lax.with_sharding_constraint把这个张量的期望分布钉住。4.3 设备数量和 Mesh 轴对不上Mesh(devices.reshape((2, 2)), (data, model))要求设备总数等于 2x2。如果你只有 3 块 GPU这行就炸。更隐蔽的情况是本地总共有 4 块卡但其中一块被其他进程占用jax.devices()返回的可用设备不足或顺序和你预期不同。这时不要硬编码设备形状先打印一次设备列表确认可用情况再生成 Mesh。常见做法是让代码支持一个环境变量来覆盖设备数量方便在 Debug 机器上跑。4.4 NCCL 相关报错和通信卡死多卡并行里JAX 的跨设备通信在 GPU 上走 NCCL。一旦环境变量不对、网卡不互通、或者 NCCL 版本与驱动不匹配就会遇到各种意想不到的卡死和超时。排查时先跑一个极简单的 all-reduce 测试确认设备间通信正常调大NCCL_DEBUGINFO看日志有没有报错检查多机场景的共享内存和网卡设置。还有一点经常被忽略不要在jit函数内部随便调用jax.device_get或np.array()这类强制同步操作它们会打断异步执行流可能让通信和计算互相等待拖出超时。4.5 常见问题速查表症状可能原因建议操作编译时间太长大函数首次 jit拆分模块、固定 shape、使用缓存显存不足batch 未有效切分 / 中间张量全部复制可视化 sharding给中间结果加约束训练 loss 和单卡不一致梯度没有跨设备同步检查 params sharding 是否为 replicated多卡性能无提升通信开销大于计算收益增大单卡计算量或改为 batch 切分更细频繁 rechunk 或 reshard函数内形状变换打断了分布用 with_sharding_constraint 固定关键张量5. 个人选型心法我现在是怎么决定用哪套 API 的这几年折腾下来我在选型上沉淀了一套比较实用的判断逻辑。只写研究代码、甚至只想在 Colab 里复现论文时我一般只用jax.gradjax.jit并行部分先不碰把单机逻辑跑通是第一优先级。单卡上需要处理大量独立样本时我优先写单样本函数然后jit(vmap(f))因为它最容易让编译器看到完整意图。当实验开始上多卡我的默认选择不是pmap即使pmap写起来更快。我会直接搭一个一维MeshNamedSharding把数据切分和参数复制用声明式方式写出来。理由很简单从第一天就用 sharding 范式后面切换到混合并行时代码骨架不用推翻重写。如果模型大到单卡放不下我再把 Mesh 扩成二维给权重按model轴切分然后观察通信 profile决定是否要手工分开前向里的两块矩阵乘。我个人的一个习惯是每次写完并行代码先开visualize_sharding把所有关键张量的布局截图看一眼再跑训练。这一步能提前暴露八成问题。另外一个小技巧调试并行逻辑时临时把 batch 设成设备数的 1/4这样即使某步 reshard 写错了也能靠小数据量快速复现和定位不用每次都在大 batch 上肉眼盯日志。这套硬件级并行范式确实有学习曲线但一旦你把“并行是运行时的额外负担”这种思维切换成“并行是编译期对计算图的一种优化”就不会再想退回手工插通信的日子了。希望这篇分享能给正在从自动微分走向并行计算的你省下几个月的弯路。
返回列表