JAX 新特性文档全景解读:分片自动微分、一等 VJP 对象与编译器控制
2026/9/11 6:28:38 网站建设 项目流程

JAX 新特性文档全景解读:分片自动微分、一等 VJP 对象与编译器控制

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

JAX 项目在 2026 年 8 月集中文档化了一批新特性,覆盖显式分片模式下的自动微分(Autodiff and sharding)、可被当作 pytree 传递与拆分的 VJP 对象、基于 hijax 的自定义微分与类型系统、多主机容错、以及逐函数/逐算子级别的编译器控制。本文以 docs/whats-new.md 的十二条条目为骨架,逐条展开其核心概念、类型规则与可复跑的代码示例,并补充仓库源码与测试中的实现依据,帮助你快速判断哪些能力可以直接引入自己的训练与推理代码。


一、分片与自动微分:cotangent 类型是 primal 类型的函数

在显式分片模式(explicit sharding mode)下,分片信息是 JAX 类型的一部分:jax.typeof(x)可能打印float32[8@X,4],表示前导轴沿网格轴X分片。新的 Autodiff and sharding 文档回答了"对这样的程序求导会发生什么"这一核心问题,其主线思想是:cotangent(梯度)类型由 primal(前向值)类型决定——反向传播中每个梯度值的形状、dtype 与分片方式,都完全由前向传播中对应值的类型决定。因此,只要确定了前向类型,反向类型(以及反向通信发生在哪里)就完全可预测。

1.1 为什么必须如此设计

文档给出了两条理由:

  1. 用户控制。显式分片模式的目标是让用户代码以可预测、局部的方式决定计算中的全部分片。反向传播虽然是自动微分生成的,但让 cotangent 分片成为 primal 分片的函数,意味着"你对前向分片的决定就是你对反向分片的决定"。编译器自动分片模式没有这一保证。
  2. 排除歧义。当前向变量被多次使用(fan-out)时,自动微分会在反向生成 cotangent 的加法。若 cotangent 分片可与 primal 分片无关,两个加数可能具有不同分片,加法就需要代码中从未指定的通信。让 cotangent 分片成为 primal 分片的函数,两个加数自动一致——它们是同一 primal 变量的 cotangent,类型必然相同。

这也是 JAX 自动微分的一般原则在分片维度上的推广:cotangent 类型始终是相应 primal 类型的函数(形状、dtype、如今再加上分片),由此每个算子的反向规则类型清晰、良类型的前向程序总能得到良类型的反向程序,且你只看前向代码就能预测反向类型。

1.2 数据并行示例:梯度同步的来源

一个典型的数据并行损失:批次x沿X分片,权重w复制(每台设备一份完整副本):

import jax import jax.numpy as jnp jax.config.update('jax_num_cpu_devices', 2) jax.set_mesh(jax.make_mesh((2,), ('X',))) # 默认进入显式分片模式 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)) # replicated -> replicated print(jax.typeof(x), '->', jax.typeof(dx)) # sharded -> sharded

分片输入得到同分片梯度,复制输入得到复制梯度——这条规则不仅对jax.grad的顶层输入输出成立,对反向传播中的每个中间值也成立。由于dw必须复制、但其数据分散在各设备上(是对沿@X分片轴的收缩结果的转置),产生复制结果需要一次跨设备求和,即 AllReduce,分区器(partitioner)会把它插入到那个dot_general内部。这正是数据并行训练中熟悉的梯度同步,而且可以纯局部地预测:w复制、每台设备只接触部分批次,反向某处必然求和各设备的梯度贡献。

1.3 新分片状态:unreduced 与 reduced

如果每个数组的每个轴都只是普通"分片",自动微分无需任何新东西。真正需要新机制的是复制(replication):复制在转置(transposition)下的对偶是归约(reduction),即跨设备求和,而必须有人决定这个和在哪里发生。于是文档引入两个新类型状态:

  • unreduced(未归约):每个设备持有完整形状的部分和(partial sum),数组的真值是这些部分的元素级求和。它"是一场等待发生的归约"。
  • reduced(已归约):物理上与复制数组完全一致(沿X每台设备一份完整副本),前向行为也与复制数组相同,唯一区别是自动微分如何对待它。

reduced之所以必须是一个新类型而非新操作,正是因为 cotangent 类型是 primal 类型的函数——要在前向请求"梯度保持未归约",就必须有对应的前向类型。完整的 cotangent 映射如下:

primal 类型cotangent 类型
分片@X分片@X
复制(replicated)复制(replicated)
归约{R:X}(reduced)未归约{U:X}(unreduced)
未归约{U:X}(unreduced)归约{R:X}(reduced)

使用复制权重得到复制梯度(自动微分即时完成归约);使用 reduced 权重得到 unreduced 梯度(归约由你决定何时执行)。

构造这两种类型的方式是out_shardingreshard中的关键字参数。例如一个收缩维被分片的矩阵乘,可以"在归约前停下":

a = 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)) c = jnp.einsum('ij,jk->ik', a, b, out_sharding=jax.P(None, None, unreduced={'X'})) print(jax.typeof(c)) # float32[2,2]{U:X}:沿 X 未归约

类型float32[2,2]{U:X}读作:一个 2×2 数组,沿网格轴X未归约。需要真值时用jax.reshard(c, jax.P(None, None))触发被推迟的 AllReduce。只有线性操作对 unreduced 数组有意义(部分和之和仍是和的部分和),非线性操作如jnp.cos会直接报错——"余弦的和"不等于"和的余弦"。因此可以先把多个分片矩阵乘(如 LoRA 风格的x @ W + x @ A @ B)的贡献累加成 unreduced,最后只做一次 AllReduce,而非每个 matmul 各做一次。

reduced的典型用法是在前向插入一次无通信的转换:

def loss2(w, x): w = jax.reshard(w, jax.P(None, None, reduced={'X'})) return jnp.sum((x @ w) ** 2)

对比 jaxpr 可以看到:原先dot_general内部隐含 AllReduce,而加上这次转换后,点积先产生显式的f32[2,4]{U:X}值,再由一个reshard(转置后的转换)变成复制的dw。数学与总通信量完全相同,但归约变成了程序中可见、可移动的对象。当权重被使用多次(两个 head、两个 micro-batch、LoRA 分支)时,各反向点积贡献的{U:X}在 fan-out 加法处未归约地相加,最后只做一次 AllReduce——这种融合由类型保证,而非交给编译器模式匹配。

1.4 实战示例:微批量梯度累积

文档给出的经典场景是梯度累积。真实训练步骤有两个编译器"看不穿"的循环:模型内部的层扫描(scan over layers)与累积梯度的微批量扫描,最后才做一次更新。若权重复制,每个微批量的反向都会在自己的梯度贡献上同步(AllReduce 位于两层循环内部),而 XLA 无法替你从循环中提出 collective。改用 reduced 权重后:

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, batch@X, features] def microbatch_step(grad_acc, xs_mb): grads = jax.grad(loss3)(stacked_ws, xs_mb) # ws 是 reduced,所以 grads 是 unreduced——而且可以当场断言! 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()) # 唯一的一次 AllReduce ws = jax.reshard(stacked_ws, jax.P()) # 免费:各设备已有完整副本 return ws - 0.01 * grads

一切局部地通过类型检查:权重是{R:X},每个微批量的梯度(即便由层扫描算出)都是{U:X};unreduced 数组支持加法,因此微批量扫描的 carry 可以累积它们;最后一次 reshard 到复制就是整个 step 的唯一 AllReduce。注意扫描体内的assert:因为分片是 JAX 类型的一部分,"梯度未归约"成为可在 traced 代码中、trace 时用jax.typeof检查的值属性。对比普通复制权重版本,其 HLO 中 AllReduce 的操作名形如while/body/.../while/body,位于转置的层扫描内部、微批量扫描内部,每层每微批量各运行一次。文档指出,把梯度归约从"每层每微批量一次"提升到"每步一次",在生产级 LLM 训练中带来了显著收益(某一案例中每步花在梯度归约上的时间下降数倍),而编译器无法自行完成此变换。

1.5 为什么这样设计

针对"为什么不让 Replicated 的 cotangent 直接是 Unreduced",文档的解释是:那样会剥夺选择权——大量代码返回复制值(如 loss),若其 cotangent 默认 unreduced,会给既有程序引入意外的通信需求。保留Replicated ↔ Replicated、另加Reduced ↔ Unreduced这对组合,你可以逐数组选择梯度是复制到达(归约已替你即时完成)还是 unreduced 到达(归约由你放置),而选择手段只是前向中一次无通信的转换。此外文档还给出了 Unreduced 的理论论证:若要求自动微分逐算子保持通信代价、复制可表达、复制标量乘分片向量(无通信前向算子)可表达、cotangent 与 primal 同形状,则"无通信地把分片 cotangent 映射回复制标量的 cotangent"必然迫出一个保持形状但只含各设备局部答案的状态——这正是 Unreduced。

1.6 手动模式:shard_map 中的对应物

上述机制在手动模式(manual mode,见 docs/201/shard-map.md)下同样成立,遵循同一规则"cotangent 类型是 primal 类型的函数"。shard_map内部沿手动轴有四种状态,与显式模式的四种状态一一对应:

显式模式(外部)手动模式(内部)沿X各设备持有
分片f32[8@X]varyingf32[4]{V:X}不同的值
复制f32[4]invaryingf32[4]相同的值
未归约f32[4]{U:X}unreducedf32[4]{U:X}真值的部分和
归约f32[4]{R:X}reducedf32[4]{R:X}相同的值,但 cotangent 未归约

in_specs/out_specs负责对应:P('X')输入绑定为内部 varying 值,P()绑定为 invarying,P(unreduced={'X'})/P(reduced={'X'})原样传入。cotangent 映射同样遵循上文表格:varying 与 invarying 各自是自己的 cotangent 类型,unreduced 与 reduced 互换。jax.lax.pcastjax.lax.psum在转置下配对:

前向类型转置类型
jax.lax.psumvarying → invaryingpcast(..., to='varying')invarying → varying
pcast(..., to='reduced')invarying → reduced对 unreduced 值做psumunreduced → invarying
pcast(..., to='unreduced')varying → unreducedpcast(..., to='varying')reduced → varying

第一行是经典故事(求和转置为标记为 varying);第二行是本文核心技巧的手动模式版本——前向一次免费的 reduced 转换,转置为反向为之买单的psum;第三行两个方向都免费。因此"前向代码中免费转换的放置位置,决定了反向psum的运行位置",reduced-weights 模式可以原样搬入手动模式:层函数用shard_map编写、权重以 reduced 类型传入,梯度以 unreduced 输出、反向全程无psum,跨微批量累积后每步仅一次 AllReduce。


二、一等 VJP 对象:前向与反向的独立编译与调度

新的 First-class VJP objects 文档阐述了一个关键事实:jax.vjp返回的可调用对象是一个一等公民值(pytree),其叶子就是前向保存的残差值。因此它可以像任何数据一样传入/传出编译函数、被序列化或卸载,其保存状态可被检查与编辑。这直接支持两种新能力:把前向与反向拆成独立编译的函数按自己的调度运行;以及用saveable_args把某些参数值(如权重)从保存状态中排除。

2.1 拆分前向与反向

jax.gradjax.vjp默认把前向与反向打包在一起,在jax.jit下编译为单个程序。用jax.vjp可直接构造分离版本,其可行性正源于 VJP 对象是 pytree:

from jax import grad, jit def fwd_and_bwd(f): def fwd(*args): return jax.vjp(f, *args) def bwd(f_vjp, y_bar): return f_vjp(y_bar) return jit(fwd), jit(bwd) def layer(W, x): return jnp.tanh(x @ W) fwd, bwd = fwd_and_bwd(layer) W1, W2 = jnp.ones((3, 3)), 2. * jnp.ones((3, 3)) x0 = jnp.ones((2, 3)) # 按自己的调度:先正向穿过两层,再反向 x1, res1 = fwd(W1, x0) x2, res2 = fwd(W2, x1) dW2, dx1 = bwd(res2, jnp.ones_like(x2)) dW1, dx0 = bwd(res1, dx1)

注意bwd中没有任何f的特定逻辑,它只是应用参数。每个fwd/bwd只编译一次,可任意次数、任意顺序调用,结果与端到端jax.grad完全一致。JAX 也把这一模式预打包为jax.fwd_and_bwd(带argnums选择要为哪些输入产生 cotangent,以及has_auxjitted等选项),其实现位于 jax/_src/api.py 的fwd_and_bwd定义中,并默认开启 jit。典型应用是在流水线或微批量调度中交错不同 micro-batch / 不同流水级的前向与反向,每函数只编译一次、多次复用。

2.2 VJP 对象保存什么

VJP 对象通过三个属性暴露保存状态:

  • args_res:反向所需"原样保留"的参数值,排列与参数镜像对应;反向不需要的参数以NotNeeded()哨兵出现。
  • opaque_residuals:前向过程中计算出的值。
  • structured_residuals:第三条通道,保留用户可理解结构的残差(详见下文)。

对于layerx @ W的反向需要原样的xW。若很多微批量经过同一层后才做反向,或每个 VJP 对象都要序列化/卸载,则每个对象都复制一份权重——权重通常是保存状态中最大的部分,且我们本已持有。这正是saveable_args的用武之地。

structured_residuals保存的是保持用户语义结构的残差:自定义微分规则可以把命名的 pytree 残差存于此,并在变换中保持结构——scan跨迭代堆叠条目、cond记录标记了哪个分支运行的和、shard_map沿前导网格轴堆叠各分片条目。JAX 还会对保存内容去重(一个值出现在多个名字下只存一次,这是优化而非保证)。对典型程序它通常为空,填充它属于 hijax primitive 规则的工作。

2.3 saveable_args:把权重排除在保存状态之外

saveable_argsjax.vjp的参数,一个 bool 的 tuple-tree(嵌套元组、bool 叶子),每个参数一项,默认为单个True(一切皆可保存)。凡是False覆盖之处,本应原样保存的参数值被替换为NotSaveable()哨兵:

y, f_vjp = jax.vjp(layer, W1, x0, saveable_args=(False, True)) print(f_vjp.args_res) # W1 的位置变成 NotSaveable() print(len(jax.tree.leaves(f_vjp))) # 3 而非 4:W1 不在保存状态中

NotSaveable是空 pytree 节点,因此展平 VJP 对象(序列化或卸载)时这些参数不产生任何叶子。应用 VJP 函数前必须先恢复缺失值,否则会抛出错误并点名仍需恢复的参数;恢复方式是指定f_vjp.args_res[0] = W1,或更函数式地用f_vjp = f_vjp.replace(args_res=[W1, ...])。组合起来,"轻量"流水线把权重直接传给反向函数,而非塞进保存状态:

def fwd_light(W, x): return jax.vjp(layer, W, x, saveable_args=(False, True)) def bwd_light(f_vjp, W, y_bar): f_vjp.args_res[0] = W return f_vjp(y_bar) fwd_light, bwd_light = jit(fwd_light), jit(bwd_light)

在 jax/_src/api.py 的vjp实现中可以看到,saveable_args先经_saveable_args_flags校验并展开为与参数树对齐的标志,随后在构造args_res时按keep = lambda x, s: ((x if s else NotSaveable()) if id(x) in used else NotNeeded())的逻辑逐叶决定保留、替换为哨兵或标记为不需要;而_vjp_not_saveable_error负责在未恢复时给出指明参数位置的错误信息。

两个细节值得注意:

  1. saveable_args只需是参数的"宽松树前缀":容器仅按子节点数量匹配(tuple 条目可对齐 dict 参数),单个 bool 会广播覆盖整个参数子树(默认True即如此)。恢复时可用原始 pytree 结构整体赋值,如g_vjp.args_res = [d]
  2. 只有"原样保存"的参数值受影响。从参数计算出的残差照常存入opaque_residualssaveable_args从不引发重算——关于保存-重算的权衡,参见 docs/301/remat.md。反向不需要的参数即使标记False也保持NotNeeded(),因此args_res精确显示哪些值必须恢复。

三、Refs:可原地读写、可与变换组合的可变数组

新文档化的 Refs 机制(锚点jax-101-refs,配套 docs/101/refs.md、jax-201-jit-refs见 docs/201/jit.md)引入jax.new_refjax.ref.new_ref):创建一个数组ref,可被原地读取和写入,并与 JAX 变换组合使用。它把"可变状态"以显式类型的方式纳入 JAX 的纯函数世界:在jit下支持原地更新,在自动微分下也有配套规则(jax-301-refs)。其实现位于 jax/_src/ref.py,核心类型Ref定义在 jax/_src/core.py 中,而 tests/state_test.py 提供了大量行为测试(包括与jitgradvmap等变换的组合)。适合用 refs 表达的场景包括状态化的内核(如 Pallas/手动模式)、需要原地累积缓冲区的循环体等。


四、hijax:新一代自定义微分、类型与残差机制

whats-new 中数条条目围绕 hijax(JAX 的原始 primitive 扩展机制)展开,彼此配套:

  • Custom derivatives with hijax primitives(docs/301/custom-derivatives.md):一个 primitive 可以同时携带两类微分规则与 batching 规则,是jax.custom_vjp/jax.custom_jvp之外更强大的替代方案——后者只能为整个函数指定一种微分方式,而 hijax primitive 可以精细到算子级别同时定义 JVP、transpose 与 batching。
  • New JAX types with hijax(docs/301/hijax-types.md):定义带有自己的 tangent 类型、batching 行为与分片方式的新 JAX 类型,并由你自己的 hijax primitives 消费——这正是 docs/301/sharding-ad.md 中unreduced/reduced这类"类型即分片状态"能力的扩展接口。
  • Structured residuals(锚点jax-301-structured-residuals):组织前向为反向保存的内容,将残差以用户可理解的结构(命名 pytree)存进 VJP 对象的structured_residuals通道,并在scan/cond/shard_map变换中保持结构。
  • Backward-pass logging(锚点jax-301-bwd-logging):把数据从反向传播中带出来,典型用途是梯度诊断(例如记录梯度范数、检查异常梯度),弥补了此前反向计算难以观测的空白。

这四条互为整体:hijax 类型是载体,hijax primitives 提供规则,structured residuals 是前向与反向之间的结构化数据通道,backward-pass logging 则把观测能力延伸进反向计算。对希望扩展 JAX 核心语义(而非仅组合现有算子)的开发者,这套机制是目前最完整的入口;仓库中的 tests/hijax_test.py 覆盖了类型、微分、batching 与分片规则的组合验证。


五、FFI with hijax:带规则的外部函数调用

FFI with hijax(docs/401/ffi.md)把外部函数接口(foreign function interface)文档围绕 hijax primitives 重写:外部调用现在可以携带自己的 batching、微分与分片规则,从而与vmapgrad以及分片输入自然组合。对性能敏感的算子(如自定义 CUDA kernel),这意味着不必再在"外部调用"与"可微分/可向量化"之间二选一。仓库中配套的示例见 examples/ffi,包含 Python 绑定与 C++/CUDA 实现的完整工程骨架。


六、容错:多主机作业中的设备故障恢复

Fault tolerance(docs/501/fault-tolerance.rst)关注多主机(multi-host)训练作业中设备故障的存活问题,核心 API 是jax.live_devices。在多机训练里,某台设备故障可能导致整个作业崩溃;容错机制允许作业探测当前仍然存活的设备集合,据此调整数据分片、shard_map网格或检查点恢复策略,从而把故障影响限制在可恢复的范围内。配套的实现细节与示例位于 docs/_static/fault_tolerance 下的 Python 演示文件中。


七、编译器控制:逐函数 flags 与逐算子元数据

Compiler control(docs/201/controlling-xla.md)给出两层 XLA 控制手段:编译器 flags(全局或逐函数地引导 XLA 如何编译)与XLA metadata(为编译后程序中的单个算子附加注解,供调试器与调度提示等编译器级工具读取)。

7.1 逐函数:jit 的compiler_options

jax.jit接受compiler_options字典,仅作用于该函数的编译,不影响程序其余部分。键是去掉--前缀的 XLA flag 名,值可以是普通 Python bool、数字或字符串:

f_opt = jax.jit(f, compiler_options={ "xla_embed_ir_in_executable": True, "xla_gpu_auto_spmd_partitioning_memory_budget_ratio": 0.5, })

一个限制:compiler_options必须放在顶层的jit上(即配置其编译的那个),而不能放在被另一个 jitted 函数内部调用的 jitted 函数上。除了 XLA 的 debug-option flags,compiler_options还接受 XLA 的编译投入度旋钮optimization_levelmemory_fitting_level,取值为jax.CompilerEffortLevel成员或其字符串名"O0""O3"

g = jax.jit(f, compiler_options={"optimization_level": jax.CompilerEffortLevel.O3})

未识别的键或非法值会在编译期立即报错(例如JaxRuntimeError: INVALID_ARGUMENT: No such compile option: 'not_a_real_flag'),拼写错误能立刻暴露。使用 AOT API 时(见 docs/201/aot.md),同一字典可在编译步骤传入:jax.jit(f).lower(1.0).compile(compiler_options={...})。从版本演进看,exec_time_optimization_effortmemory_fitting_effort旧 flags 已被移除,统一由EffortLevel枚举取代(见 CHANGELOG.md 0.11.1 的 breaking changes)。

7.2 进程级:XLA_FLAGS环境变量

要为整个进程(包括你不直接控制的编译)配置 XLA,设置XLA_FLAGS,各 flag 以空格分隔:

XLA_FLAGS='--flag1=value1 --flag2=value2' python3 source.py

XLA_FLAGS在 JAX 初始化后端时读取,因此必须在导入 JAX 之前设置,之后再改无效。XLA flags 的默认值定义于其debug_options_flags.cc、完整列表见xla.proto,JAX 侧则通过 jax/_src/config.py 等模块转发。

7.3 逐算子元数据:xla_metadata_call等 API

编译后 XLA 程序中的每个算子都可以携带元数据:字符串值的frontend_attributes不改变计算本身,但调试器、融合控制、调度提示等编译器级工具可以读取。JAX 的接口位于jax.experimental.xla_metadata(实验性,可能变动),推荐入口是xla_metadata_call

from jax.experimental.xla_metadata import xla_metadata_call @xla_metadata_call(tag="my_block") def block(x): y = jnp.sin(x) return y * jnp.cos(x) @jax.jit def f(x): return block(x) + 1. print(f.lower(1.0).as_text("hlo"))

被包装的函数会作为一个独立的子计算(subcomputation)staged 出来,调用点携带元数据(如frontend_attributes={tag="my_block"});XLA 优化内联该调用时会把属性传播到被内联的算子上。元数据值可以是字符串、bool、int 或 float,统一按字符串附加(bool 渲染为"true"/"false")。

由于元数据附着在函数而非环境式 trace 状态上,JAX 变换会保留它:jax.vmap下批量算子携带它,jax.grad下由该函数衍生的一切计算(前向与反向)都携带它——HLO 中可以看到前向残差计算与反向计算分别 staged 且都被打上标签。若想让反向不带标签或使用不同标签(例如自己的调度组),用xla_metadata_call2,它把元数据作为 dict 传入,并有ad_metadata选项:ad_metadata='drop'只标记前向,ad_metadata={"tag": "y"}为反向重新打标。一个基于此构建的应用是must_fuse_call:包装函数使 XLA 必须将其所有算子放进单一融合。

另有set_xla_metadata两种模式:包装只标记产生它的那一个算子(set_xla_metadata(y * z, breakpoint=True));不带值调用时作为上下文管理器/装饰器,标记其下 traced 的每个算子。两种模式各有文档明确指出的局限:上下文管理器通过环境式 trace 状态工作,是jit缓存键的一部分,其下任何 jitted 函数(包括内部 jit 的库代码)都会为每个不同元数据上下文重新 trace 与编译;值标记不经过自动微分传播(对g求导后反向算子无标签)。因此除非确实只需原地标记单个算子,否则优先用xla_metadata_call。最后,所有这类 API 都有一项共同注意点:XLA 意图在优化中保留frontend_attributes,但边缘情况可能丢弃它们——若某工具依赖元数据存活,请检查优化后的 HLO。


八、矩阵乘法精度控制:逐算子与全局

Matmul precision control(docs/201/precision.md)介绍precision参数——jax.lax.dot_generaljax.lax.dot以及基于它们的jax.numpy函数(jnp.dotjnp.matmul@jnp.einsum、卷积)都接受它。加速器硬件提供多种矩阵乘法实现,在精度与速度间权衡:真正的float32算术、NVIDIA tensor core 上的 TensorFloat32(TF32)、TPU 上一次或多次bfloat16pass、各种float8模式等。JAX 默认偏向速度(float32点积内部可能用降精度算术,TPU 上是 bf16、新 GPU 上是 TF32),但你可以逐算子与全局显式控制。

最直接的方式是传入jax.lax.DotAlgorithmPreset(或其字符串名)作为precision

y = jnp.dot(x, x, precision="F32_F32_F32") # 真正的 float32 y = jnp.dot(x, x, precision="BF16_BF16_F32") # bf16 输入、f32 累加 y = jnp.dot(x, x, precision=lax.DotAlgorithmPreset.TF32_TF32_F32) # TF32 tensor core

预设名遵循LHS_RHS_ACCUM模式:左右操作数被舍入到的元素类型,以及累加所用的类型。可用预设包括:

  • DEFAULT——根据输入与输出类型选择算法;
  • F32_F32_F32F64_F64_F64——普通全精度算术;
  • F16_F16_F16F16_F16_F32——半精度输入,半精度或单精度累加;
  • BF16_BF16_BF16BF16_BF16_F32——bfloat16同理;
  • BF16_BF16_F32_X3_X6_X9——_X后缀表示用多少次 bf16 运算模拟更高精度:_X3接近float32精度,_X6/_X9超过它,代价是相应更高的成本;
  • TF32_TF32_F32TF32_TF32_F32_X3——TensorFloat32 及其三次运算的高精度模拟;
  • ANY_F8_ANY_F8_F32ANY_F8_ANY_F8_F32_FAST_ACCUM——任意float8输入类型、float32累加;FAST_ACCUM变体使用更快但精度略低的累加(如 cuBLASLt 的快速累加模式);
  • ANY_F8_ANY_F8_ANYANY_F8_ANY_F8_ANY_FAST_ACCUM——同上,累加类型由preferred_element_type控制。

该接口的几个性质:接受任意输入 dtype(JAX 自动插入 cast 让操作数以算法的存储类型到达硬件);输出类型与输入一致(按通常的提升规则),与内部累加类型无关,因此切换算法不会在你的程序中引起类型涟漪——想保留累加器类型则用preferred_element_type自动微分把同一算法带到反向,梯度计算中的转置点积携带与 primal 相同的precision参数,可从梯度的 jaxpr 中直接看到。全局控制则通过jax.config中的全局精度配置(如jax_default_matmul_precision)实现,参见 docs/101/type_promotion.rst 相关章节。


九、如何跟进这批新文档

whats-new.md是"首次被文档化"特性的索引页,随 CHANGELOG.md 一起维护——后者按版本(当前仓库为 JAX 0.11.x,见 jax/version.py)记录新增特性、破坏性变更、弃用与 bug 修复。建议的跟进方式:

  1. 从本文按主题挑选最贴近自身场景的条目,直接阅读对应的完整文档(分片自动微分、VJP 对象、编译器控制、精度控制是四条实操性最强、最值得先读的);
  2. 关注 CHANGELOG 中对应的版本发布说明,确认 API 的稳定性与变动(例如 effort flags 已统一为CompilerEffortLevel枚举);
  3. 需要更底层验证时,以本文给出的源码路径(如 jax/_src/api.py 中的vjp/fwd_and_bwd、jax/_src/ref.py 中的new_ref)与测试文件(如 tests/state_test.py、tests/hijax_test.py)为证。

上述特性中,分片自动微分与微批量梯度累积、一等 VJP 对象与流水线调度、saveable_args与权重排除、逐函数编译选项与逐算子元数据,都是可以直接落地到训练与推理代码的实用能力;hijax 系列与 Refs 则面向希望扩展 JAX 语义的进阶开发者。

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询