ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

JAX 显式 sharding 模式下梯度(cotangent)的切分如何确定:unreduced 与 reduced 状态

JAX 显式 sharding 模式下梯度(cotangent)的切分如何确定:unreduced 与 reduced 状态 JAX 显式 sharding 模式下梯度cotangent的切分如何确定unreduced 与 reduced 状态【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax在 JAX 的显式 sharding 模式explicit sharding mode下sharding 是类型的一部分jax.typeof(x)可能打印出float32[8X,4]表示第一个轴沿 mesh 轴X切分。当你用jax.grad对这样的程序求导时梯度cotangent的切分不由编译器自动决定而是由前向程序的值类型shape、dtype、sharding局部地确定。本文的任务是读懂 cotangent 切分的判定规则学会使用unreduced与reduced两种状态显式放置反向传播中的 AllReduce并在一个微批梯度累加的例子里验证整个 step 只产生一次 AllReduce。文中示例基于 2 个 CPU 设备的单轴 mesh 环境。环境准备单轴 mesh 与显式模式按 docs/301/sharding-ad.md 的设定先准备两个 CPU 设备和一个单轴 meshimport jax import jax.numpy as jnp jax.config.update(jax_num_cpu_devices, 2) jax.set_mesh(jax.make_mesh((2,), (X,))) # explicit mode by defaultjax.set_mesh(jax.make_mesh(...))使程序默认运行在显式模式下用户代码决定计算中所有值的 sharding。显式模式的完整背景见 docs/201/sharding.md。判定规则cotangent 的切分是 primal 切分的函数先看一个数据并行 loss批数据x沿X切分权重w是复制replicated每个设备一份完整拷贝x jax.device_put(jnp.arange(8 * 4.).reshape(8, 4), jax.P(X, None)) w jax.device_put(jnp.arange(4 * 2.).reshape(4, 2) / 10., jax.P(None, None)) def loss(w, x): return jnp.sum((x w) ** 2) dw, dx jax.grad(loss, argnums(0, 1))(w, x) print(jax.typeof(w), -, jax.typeof(dw)) print(jax.typeof(x), -, jax.typeof(dx))切分输入的梯度以同样方式切分复制输入的梯度是复制的。这一条不仅对jax.grad的顶层输入输出成立对反向传播中的每个中间值都成立可以直接在 jaxpr 里核对print(jax.jit(jax.grad(loss)).trace(w, x).jaxpr)在上面的 jaxpr 中每个 cotangent 值的 sharding 都与其对应的 primal 一致最后两条等式计算dw两个沿X切分的值做dot_general得到复制unsharded结果随后做 transpose。JAX 强制「cotangent 切分是 primal 切分的函数」有两个原因理解它们能解释后面的两个新状态用户可控显式模式的目标是用户代码以可预测的局部方式决定所有 sharding。反向传播由 autodiff 生成而不是你手写cotangent 切分由 primal 切分决定意味着你在前向的切分决策就是反向的切分决策。编译器式的自动 sharding 模式没有这个保证那里的反向切分可以由编译器独立选择。消除歧义前向中一个变量被多次使用fan-out时autodiff 会在反向生成 cotangent 的加法。若 cotangent 切分可以脱离 primal 切分两个加数可能具有不同的 sharding加法就需要代码没有指定的通信。而 cotangent 切分是 primal 切分的函数时两个加数自动具有相同类型。通信从哪里来dw必须是复制的因为w是复制的但构建它的数据分散在各设备上——dw是沿X切分的两个数组的收缩contraction的转置。从切分轴上的收缩得到复制结果需要一个跨设备求和AllReduce分区器会把它插进最后那条dot_general内部。这就是数据并行训练中熟悉的梯度同步可以在本地纯静态地预测w是复制的每个设备只接触部分 batch所以反向传播中某处必须对各设备的梯度贡献求和。区别在于此时 AllReduce 在 jaxpr 里是隐式的藏在一条收缩维被切分的 op 内部。unreduced与reduced就是让这条通信显式化、可移动、可合并的手段。unreduced一次等待发生的归约先看一个收缩维被切分的 matmula jax.device_put(jnp.arange(4.).reshape(2, 2), jax.P(None, X)) b jax.device_put(jnp.arange(4., 8.).reshape(2, 2), jax.P(X, None)) print(jax.typeof(a)) print(jax.typeof(b))此时直接做a b会报错JAX 要求你显式声明意图因为这里存在真实的选择。文档给出的数据结构示意a: P(None, X) b: P(X, None) (device k has column k) (device k has row k) [ 0 | 1 ] [ 4 5 ] --------- [ 2 | 3 ] [ 6 7 ]每个设备可以无通信地把自己a的列乘上b的行每次本地 matmul 产生一个全形状的部分和真实答案是两份部分和的逐元素和。有两种处理方式方式一要求普通输出 sharding用 AllReduce 收尾c jnp.einsum(ij,jk-ik, a, b, out_shardingjax.P(None, None)) print(jax.typeof(c)) print(c)方式二停在归约之前得到 unreduced 数组c jnp.einsum(ij,jk-ik, a, b, out_shardingjax.P(None, None, unreduced{X})) print(jax.typeof(c))文档给出的类型读法是float32[2,2]{U:X}一个 2×2 数组沿 mesh 轴Xunreduced。每个设备持有全形状的部分和数组的真实值是这些部分的和。可以用addressable_shards检查各设备上的数据for shard in c.addressable_shards: print(fdevice {shard.device.id}:\n{shard.data})unreduced 数组就是「一次等待发生的归约」。要执行它把它 reshard 到普通 sharding延迟的 AllReduce 就会运行print(jax.reshard(c, jax.P(None, None)))延迟归约的价值在于先做更多工作、再用更少更大的集合通信支付代价。但有两条限制只有线性操作对 unreduced 数组有意义。沿相同轴 unreduced 的两个数组相加没问题部分和的和仍是部分和c2 jnp.einsum(ij,jk-ik, 2. * a, b, out_shardingjax.P(None, None, unreduced{X})) print(jax.typeof(c c2))这使得切分 matmul 的求和例如 LoRA 式的x W x A B可以在最后做一次 AllReduce而不是每个 matmul 各做一次。非线性操作无法对 unreduced 输入定义规则和的余弦不等于余弦的和对每个设备上的部分和取cos会算错因此这类操作直接报错jnp.cos(c) # raises error一个注意事项对c c2这类直线代码XLA 常常能自行合并相邻的 AllReduce编译后的程序两种写法可能一样好。类型层面的保证在编译器找不到合并点时才关键尤其是跨越循环迭代的情形——下面的微批例子就是这种情况。reduced选择 unreduced 的梯度回到 autodiff。已知 cotangent 切分是 primal 切分的函数且已有两条映射sharded 得到 sharded cotangentreplicated 得到 replicated cotangent。replicated 这一条正是反向通信的来源复制replication在转置下的对偶是归约跨设备求和。在CT(Replicated) Replicated下autodiff急切地执行该求和每个为复制 primal 产生 cotangent 的反向 op 都在现场隐式做 AllReduce。要反过来——让 autodiff 把 cotangent 留在「各设备部分和」的状态、由你决定何时归约——需要在 forward pass 里有一种新的类型因为 cotangent 类型必须是 primal 类型的函数。这就是reducedw_ jax.reshard(w, jax.P(None, None, reduced{X})) print(jax.typeof(w_))沿X为reduced的数组记作{R:X}物理上与 replicated 数组完全相同沿X的每个设备一份完整拷贝名字是 unreduced 的过去时态归约已经完成后的状态。前向中它表现得和 replicated 数组一模一样上面的 reshard 不移动任何数据。唯一区别是 autodiff 对它的处理。完整的 cotangent 映射表primal 类型cotangent 类型shardedXshardedXreplicatedreplicatedreduced{R:X}unreduced{U:X}unreduced{U:X}reduced{R:X}用 Replicated 得到 replicated 梯度用 Reduced 得到 unreduced 梯度。前向中的这个无通信 cast在转置下变成反向中的 reshard-from-unreduced也就是那次 AllReduce。前向代码里的一次免费 cast 固定了反向代码中集合通信的位置使反向通信成为写前向时可见、可摆放的对象。对比加与不加 cast 的 backward jaxprdef loss2(w, x): w jax.reshard(w, jax.P(None, None, reduced{X})) return jnp.sum((x w) ** 2) print(jax.jit(jax.grad(loss2)).trace(w, x).jaxpr)原 jaxpr 里是埋着隐式 AllReduce 的dot_general这份 jaxpr 中 dot 直接产出一个显式的f32[2,4]{U:X}值最后由一条reshard转置后的 cast把它变成 replicated 的dw。数学相同、总通信量相同但归约成了程序中可见、可移动的对象。可见之后就可以移动。假设权重被使用两次两个 head、两个微批、一个 LoRA 分支def loss_fanout(w, x1, x2): w jax.reshard(w, jax.P(None, None, reduced{X})) return jnp.sum(x1 w) jnp.sum(x2 w) x1 jax.device_put(jnp.ones((8, 4)), jax.P(X, None)) x2 jax.device_put(jnp.ones((8, 4)), jax.P(X, None)) print(jax.jit(jax.grad(loss_fanout)).trace(w, x1, x2).jaxpr)两个反向 dot 都产生{U:X}贡献fan-out 加法在unreduced状态下进行加法是线性的随后一条reshard为整个梯度做一次 AllReduce。不加 cast 的话每个贡献会在各自的 dot 内部被分别归约。这里融合由类型保证而不是留给编译器模式匹配。实战微批梯度累加中每个 step 只留一次 AllReduce典型场景是梯度累加。一个真实的训练 step 有两个编译器无法看穿的循环模型内部对 layer 的 scan以及累加梯度、最后更新一次的微批 scan。梯度 AllReduce 应该每个 step 只做一次但如果权重是 replicated 的每个微批的反向传播都会同步自己的梯度贡献发生在两个循环内部而 XLA 不能替你把集合通信提升出循环。用 reduced 权重梯度以 unreduced 形式出来累加器在整个 scan 中保持 unreduced最后归约一次def predict(stacked_ws, xs): # stacked_ws: [layer, features, features] def apply_layer(xs, w): return jnp.tanh(xs w), None final_xs, _ jax.lax.scan(apply_layer, xs, stacked_ws) return final_xs def loss3(stacked_ws, batch): return jnp.sum(predict(stacked_ws, batch) ** 2) jax.jit def step(stacked_ws, xs): # xs: [microbatch, batchX, features] def microbatch_step(grad_acc, xs_mb): grads jax.grad(loss3)(stacked_ws, xs_mb) # ws are reduced, so grads are unreduced -- and we can check it! assert jax.typeof(grads).sharding.spec.unreduced {X} return grad_acc grads, None grad_acc jax.reshard(jnp.zeros_like(stacked_ws), jax.P(unreduced{X})) grad_acc, _ jax.lax.scan(microbatch_step, grad_acc, xs) grads jax.reshard(grad_acc, jax.P()) # the one AllReduce ws jax.reshard(stacked_ws, jax.P()) # free: full copies already return ws - 0.01 * grads stacked_ws jax.device_put(jnp.stack([jnp.eye(4) / 2] * 3), jax.P(reduced{X})) xs jax.device_put(jnp.ones((5, 2, 4)), jax.P(None, X, None)) new_ws step(stacked_ws, xs) print(jax.typeof(new_ws))这段程序的所有类型都能在本地推导权重是{R:X}所以每个微批的梯度以{U:X}出来即使它们由对 layer 的 scan 计算得出unreduced 数组支持加法所以微批 scan 的 carry 可以累加它们一次 reshard 到 replicated 就是整个 step 的唯一 AllReduce。注意 scan body 里的assert因为 sharding 是 JAX 类型的一部分「梯度是 unreduced」是某个值的属性可以用jax.typeof在 traced 代码中、在 trace 阶段检查。更新后的权重以 replicated 形式返回下一步把它 cast 回{R:X}是免费的。验证方式把编译后的 HLO 文本里所有 all-reduce 指令的 op name 打印出来确认其位置和数量import re def print_all_reduces(jitted, *args): hlo jitted.lower(*args).compile().as_text() for line in hlo.splitlines(): if all-reduce( in line or all-reduce-start( in line: print(re.search(rop_name([^]*), line).group(1)) print_all_reduces(step, stacked_ws, xs)对 reduced 权重的step文档说明编译后的程序恰好有一个 AllReduce位于两个循环之外。再把同一个 step 写成普通 replicated 权重的版本对比jax.jit def step_replicated(stacked_ws, xs): def microbatch_step(grad_acc, xs_mb): grads jax.grad(loss3)(stacked_ws, xs_mb) # replicated: AllReduce inside! return grad_acc grads, None grad_acc jnp.zeros_like(stacked_ws) grad_acc, _ jax.lax.scan(microbatch_step, grad_acc, xs) return stacked_ws - 0.01 * grad_acc ws_replicated jax.reshard(stacked_ws, jax.P()) print_all_reduces(step_replicated, ws_replicated, xs)文档指出此版本打印出的 op name 形如while/body/.../while/body说明该 AllReduce 位于转置后的 layer scan 内部、微批 scan 内部每层每个微批执行一次。把梯度归约从两个循环里提升出来从每层每微批一次到每 step 一次在生产 LLM 训练中带来了可观收益——文档举例某个案例将每 step 花在梯度归约上的时间削减了数倍。编译器无法自行完成这个变换因为它无法跨循环模式匹配集合通信。限制与边界unreduced 数组只支持线性操作。非线性操作如jnp.cos作用于 unreduced 输入会直接报错因为对部分和应用非线性函数不等于对和应用。XLA 的自动合并不能依赖。对直线代码 XLA 常能自行合并相邻 AllReduce编译结果可能两种写法一样好类型保证的价值在编译器找不到的场景尤其是循环内。选择是逐数组的。设计文档解释了为什么不让CT(Replicated) Unreduced很多代码返回复制值比如 loss 值把它们的 cotangent 变成 unreduced 会给现有程序引入意外的通信要求。保持Replicated ↔ Replicated、另加Reduced ↔ Unreduced这一对让每个数组可以独立选择梯度以 replicated 形式到达归约由系统急切代做或以 unreduced 形式到达归约位置由你决定选择通过前向中一次无通信 cast 完成。手动模式有对应机制。在jax.shard_map手动模式见 docs/201/shard-map.md中每个 mesh 轴同样有四种状态varying / invarying / unreduced / reducedcotangent 映射是显式模式的镜像unreduced 与 reduced 互换前向中jax.lax.pcast(..., toreduced)这个免费 cast 在转置下变成反向中为之支付代价的jax.lax.psum。同样地reduced 权重模式在手动模式中逐字成立给权重 reduced 类型输入梯度即以 unreduced 出来反向传播中不出现任何psum。完整的 collective 与转置对照表在 docs/301/sharding-ad.md 的手动模式一节。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表