ARTICLE DETAIL

资讯详情

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

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

JAX 新特性文档全景解读:分片自动微分、一等 VJP 对象与编译器控制 JAX 新特性文档全景解读分片自动微分、一等 VJP 对象与编译器控制【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 项目在 2026 年 8 月集中文档化了一批新特性覆盖显式分片模式下的自动微分Autodiff and sharding、可被当作 pytree 传递与拆分的 VJP 对象、基于 hijax 的自定义微分与类型系统、多主机容错、以及逐函数/逐算子级别的编译器控制。本文以 docs/whats-new.md 的十二条条目为骨架逐条展开其核心概念、类型规则与可复跑的代码示例并补充仓库源码与测试中的实现依据帮助你快速判断哪些能力可以直接引入自己的训练与推理代码。一、分片与自动微分cotangent 类型是 primal 类型的函数在显式分片模式explicit sharding mode下分片信息是 JAX 类型的一部分jax.typeof(x)可能打印float32[8X,4]表示前导轴沿网格轴X分片。新的 Autodiff and sharding 文档回答了对这样的程序求导会发生什么这一核心问题其主线思想是cotangent梯度类型由 primal前向值类型决定——反向传播中每个梯度值的形状、dtype 与分片方式都完全由前向传播中对应值的类型决定。因此只要确定了前向类型反向类型以及反向通信发生在哪里就完全可预测。1.1 为什么必须如此设计文档给出了两条理由用户控制。显式分片模式的目标是让用户代码以可预测、局部的方式决定计算中的全部分片。反向传播虽然是自动微分生成的但让 cotangent 分片成为 primal 分片的函数意味着你对前向分片的决定就是你对反向分片的决定。编译器自动分片模式没有这一保证。排除歧义。当前向变量被多次使用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_sharding与reshard中的关键字参数。例如一个收缩维被分片的矩阵乘可以在归约前停下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_shardingjax.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, batchX, 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[8X]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()绑定为 invaryingP(unreduced{X})/P(reduced{X})原样传入。cotangent 映射同样遵循上文表格varying 与 invarying 各自是自己的 cotangent 类型unreduced 与 reduced 互换。jax.lax.pcast与jax.lax.psum在转置下配对前向类型转置类型jax.lax.psumvarying → invaryingpcast(..., tovarying)invarying → varyingpcast(..., toreduced)invarying → reduced对 unreduced 值做psumunreduced → invaryingpcast(..., tounreduced)varying → unreducedpcast(..., tovarying)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.grad和jax.vjp默认把前向与反向打包在一起在jax.jit下编译为单个程序。用jax.vjp可直接构造分离版本其可行性正源于 VJP 对象是 pytreefrom 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_aux、jitted等选项其实现位于 jax/_src/api.py 的fwd_and_bwd定义中并默认开启 jit。典型应用是在流水线或微批量调度中交错不同 micro-batch / 不同流水级的前向与反向每函数只编译一次、多次复用。2.2 VJP 对象保存什么VJP 对象通过三个属性暴露保存状态args_res反向所需原样保留的参数值排列与参数镜像对应反向不需要的参数以NotNeeded()哨兵出现。opaque_residuals前向过程中计算出的值。structured_residuals第三条通道保留用户可理解结构的残差详见下文。对于layerx W的反向需要原样的x与W。若很多微批量经过同一层后才做反向或每个 VJP 对象都要序列化/卸载则每个对象都复制一份权重——权重通常是保存状态中最大的部分且我们本已持有。这正是saveable_args的用武之地。structured_residuals保存的是保持用户语义结构的残差自定义微分规则可以把命名的 pytree 残差存于此并在变换中保持结构——scan跨迭代堆叠条目、cond记录标记了哪个分支运行的和、shard_map沿前导网格轴堆叠各分片条目。JAX 还会对保存内容去重一个值出现在多个名字下只存一次这是优化而非保证。对典型程序它通常为空填充它属于 hijax primitive 规则的工作。2.3 saveable_args把权重排除在保存状态之外saveable_args是jax.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 而非 4W1 不在保存状态中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负责在未恢复时给出指明参数位置的错误信息。两个细节值得注意saveable_args只需是参数的宽松树前缀容器仅按子节点数量匹配tuple 条目可对齐 dict 参数单个 bool 会广播覆盖整个参数子树默认True即如此。恢复时可用原始 pytree 结构整体赋值如g_vjp.args_res [d]。只有原样保存的参数值受影响。从参数计算出的残差照常存入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 提供了大量行为测试包括与jit、grad、vmap等变换的组合。适合用 refs 表达的场景包括状态化的内核如 Pallas/手动模式、需要原地累积缓冲区的循环体等。四、hijax新一代自定义微分、类型与残差机制whats-new 中数条条目围绕 hijaxJAX 的原始 primitive 扩展机制展开彼此配套Custom derivatives with hijax primitivesdocs/301/custom-derivatives.md一个 primitive 可以同时携带两类微分规则与 batching 规则是jax.custom_vjp/jax.custom_jvp之外更强大的替代方案——后者只能为整个函数指定一种微分方式而 hijax primitive 可以精细到算子级别同时定义 JVP、transpose 与 batching。New JAX types with hijaxdocs/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 hijaxdocs/401/ffi.md把外部函数接口foreign function interface文档围绕 hijax primitives 重写外部调用现在可以携带自己的 batching、微分与分片规则从而与vmap、grad以及分片输入自然组合。对性能敏感的算子如自定义 CUDA kernel这意味着不必再在外部调用与可微分/可向量化之间二选一。仓库中配套的示例见 examples/ffi包含 Python 绑定与 C/CUDA 实现的完整工程骨架。六、容错多主机作业中的设备故障恢复Fault tolerancedocs/501/fault-tolerance.rst关注多主机multi-host训练作业中设备故障的存活问题核心 API 是jax.live_devices。在多机训练里某台设备故障可能导致整个作业崩溃容错机制允许作业探测当前仍然存活的设备集合据此调整数据分片、shard_map网格或检查点恢复策略从而把故障影响限制在可恢复的范围内。配套的实现细节与示例位于 docs/_static/fault_tolerance 下的 Python 演示文件中。七、编译器控制逐函数 flags 与逐算子元数据Compiler controldocs/201/controlling-xla.md给出两层 XLA 控制手段编译器 flags全局或逐函数地引导 XLA 如何编译与XLA metadata为编译后程序中的单个算子附加注解供调试器与调度提示等编译器级工具读取。7.1 逐函数jit 的compiler_optionsjax.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 flagscompiler_options还接受 XLA 的编译投入度旋钮optimization_level与memory_fitting_level取值为jax.CompilerEffortLevel成员或其字符串名O0至O3g 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_effort与memory_fitting_effort旧 flags 已被移除统一由EffortLevel枚举取代见 CHANGELOG.md 0.11.1 的 breaking changes。7.2 进程级XLA_FLAGS环境变量要为整个进程包括你不直接控制的编译配置 XLA设置XLA_FLAGS各 flag 以空格分隔XLA_FLAGS--flag1value1 --flag2value2 python3 source.pyXLA_FLAGS在 JAX 初始化后端时读取因此必须在导入 JAX 之前设置之后再改无效。XLA flags 的默认值定义于其debug_options_flags.cc、完整列表见xla.protoJAX 侧则通过 jax/_src/config.py 等模块转发。7.3 逐算子元数据xla_metadata_call等 API编译后 XLA 程序中的每个算子都可以携带元数据字符串值的frontend_attributes不改变计算本身但调试器、融合控制、调度提示等编译器级工具可以读取。JAX 的接口位于jax.experimental.xla_metadata实验性可能变动推荐入口是xla_metadata_callfrom jax.experimental.xla_metadata import xla_metadata_call xla_metadata_call(tagmy_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))被包装的函数会作为一个独立的子计算subcomputationstaged 出来调用点携带元数据如frontend_attributes{tagmy_block}XLA 优化内联该调用时会把属性传播到被内联的算子上。元数据值可以是字符串、bool、int 或 float统一按字符串附加bool 渲染为true/false。由于元数据附着在函数而非环境式 trace 状态上JAX 变换会保留它jax.vmap下批量算子携带它jax.grad下由该函数衍生的一切计算前向与反向都携带它——HLO 中可以看到前向残差计算与反向计算分别 staged 且都被打上标签。若想让反向不带标签或使用不同标签例如自己的调度组用xla_metadata_call2它把元数据作为 dict 传入并有ad_metadata选项ad_metadatadrop只标记前向ad_metadata{tag: y}为反向重新打标。一个基于此构建的应用是must_fuse_call包装函数使 XLA 必须将其所有算子放进单一融合。另有set_xla_metadata两种模式包装值只标记产生它的那一个算子set_xla_metadata(y * z, breakpointTrue)不带值调用时作为上下文管理器/装饰器标记其下 traced 的每个算子。两种模式各有文档明确指出的局限上下文管理器通过环境式 trace 状态工作是jit缓存键的一部分其下任何 jitted 函数包括内部 jit 的库代码都会为每个不同元数据上下文重新 trace 与编译值标记不经过自动微分传播对g求导后反向算子无标签。因此除非确实只需原地标记单个算子否则优先用xla_metadata_call。最后所有这类 API 都有一项共同注意点XLA 意图在优化中保留frontend_attributes但边缘情况可能丢弃它们——若某工具依赖元数据存活请检查优化后的 HLO。八、矩阵乘法精度控制逐算子与全局Matmul precision controldocs/201/precision.md介绍precision参数——jax.lax.dot_general、jax.lax.dot以及基于它们的jax.numpy函数jnp.dot、jnp.matmul、、jnp.einsum、卷积都接受它。加速器硬件提供多种矩阵乘法实现在精度与速度间权衡真正的float32算术、NVIDIA tensor core 上的 TensorFloat32TF32、TPU 上一次或多次bfloat16pass、各种float8模式等。JAX 默认偏向速度float32点积内部可能用降精度算术TPU 上是 bf16、新 GPU 上是 TF32但你可以逐算子与全局显式控制。最直接的方式是传入jax.lax.DotAlgorithmPreset或其字符串名作为precisiony jnp.dot(x, x, precisionF32_F32_F32) # 真正的 float32 y jnp.dot(x, x, precisionBF16_BF16_F32) # bf16 输入、f32 累加 y jnp.dot(x, x, precisionlax.DotAlgorithmPreset.TF32_TF32_F32) # TF32 tensor core预设名遵循LHS_RHS_ACCUM模式左右操作数被舍入到的元素类型以及累加所用的类型。可用预设包括DEFAULT——根据输入与输出类型选择算法F32_F32_F32、F64_F64_F64——普通全精度算术F16_F16_F16、F16_F16_F32——半精度输入半精度或单精度累加BF16_BF16_BF16、BF16_BF16_F32——bfloat16同理BF16_BF16_F32_X3、_X6、_X9——_X后缀表示用多少次 bf16 运算模拟更高精度_X3接近float32精度_X6/_X9超过它代价是相应更高的成本TF32_TF32_F32、TF32_TF32_F32_X3——TensorFloat32 及其三次运算的高精度模拟ANY_F8_ANY_F8_F32、ANY_F8_ANY_F8_F32_FAST_ACCUM——任意float8输入类型、float32累加FAST_ACCUM变体使用更快但精度略低的累加如 cuBLASLt 的快速累加模式ANY_F8_ANY_F8_ANY、ANY_F8_ANY_F8_ANY_FAST_ACCUM——同上累加类型由preferred_element_type控制。该接口的几个性质接受任意输入 dtypeJAX 自动插入 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 修复。建议的跟进方式从本文按主题挑选最贴近自身场景的条目直接阅读对应的完整文档分片自动微分、VJP 对象、编译器控制、精度控制是四条实操性最强、最值得先读的关注 CHANGELOG 中对应的版本发布说明确认 API 的稳定性与变动例如 effort flags 已统一为CompilerEffortLevel枚举需要更底层验证时以本文给出的源码路径如 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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表