
Flax NNX Tree Mode以 JAX 树语义重构 NNX 变换体系的设计与迁移指南【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本篇文章基于 Flax 官方设计提案FLIP《Tree Mode NNX》展开系统讲解 Flax NNX 即将推出的 Tree Mode一套只处理 pytree、假定引用透明性、与 JAX 变换完全对齐的 NNX API 重新实现。文章覆盖 Tree Mode 的动机、graph/graph_updates参数与配置开关、简化后的变换内部模式、向后兼容方案以及 prefix filters、nnx.grad、nnx.custom_vjp、transform_metadata、Module.sow/Module.perturb等六大破坏性变更的具体改写方式。读完本文你将理解 NNX 从图模式走向树模式的设计取舍并能把存量 NNX 代码迁移到新的 Tree Mode 语义。一、背景与动机NNX 现有能力为何需要简化当前 NNX API 支持通用的图结构与图变换涵盖四类能力追踪 Variable 状态的更新处理共享引用即图结构支持前缀过滤器prefix filtersStateAxes、DiffState、StateSharding传播图更新静态状态与结构变化。这四种能力中第 3 项与第 4 项超出了 JAX 变换 API 的能力范围。支撑它们带来了三重代价内部复杂度高需要专门维护图的遍历、别名检测与结构更新传播代码难以推理共享引用使得一个对象被多处修改的行为难以追踪API 学习负担大用户必须额外掌握前缀过滤器等一套 JAX 之外的抽象。FLIP 5310 的初衷就是通过简化 NNX 来同时解决以上问题让 NNX 变换与 JAX 原生变换在语义上对齐。二、核心提案Tree Mode NNX 与图支持的精简提案包含两条主线。2.1 Tree Mode只处理树的 NNX 重实现Tree Mode NNX是 NNX API 的一次重新实现其核心约束是自动状态更新仅限 NNX 变换中的 Variable不再有通用图更新机制状态更新只发生在显式变换边界内所有 API 假定并强制树结构取消共享引用shared referencesModule 视为无状态 pytree不再传播图结构更新完整兼容 JAX 变换移除前缀过滤器StateAxes、DiffState、StateSharding。这意味着进入 Tree Mode 后一个nnx.Module就是一棵普通的 JAX pytree可以直接被jax.jit、jax.vmap、jax.grad等原生变换处理NNX 与 JAX 的边界被大幅抹平。2.2 图支持的保留范围图graph对部分 NNX 用户仍是重要特性因此提案保留能力 1追踪状态更新与能力 2共享引用而放弃前缀过滤器与图更新传播能力 3、4。经过裁剪后树与图两种变换可以共享同一套底层实现与语义同时保持足够的表达力。三、实现方案graph 与 graph_updates 参数Tree Mode 在现有 API 之上实现引入两个新参数def split(..., graph: bool | None None) ... def jit(..., graph: bool | None None, graph_updates: bool | None None) ...参数语义参数True图模式False树模式graph启用图支持内部走图协议只支持树内部依赖jax.tree.*APIgraph_updates传播图结构更新能力 4支持前缀过滤器能力 3变换不再传播图结构更新也不支持前缀过滤器当graph或graph_updates未显式给出时其默认值取自配置标志nnx_graph_mode与nnx_graph_updates。3.1 配置标志与运行时开关提案目标是将nnx_graph_mode与nnx_graph_updates的默认值设为False从而让新项目默认进入 Tree Mode。在仓库当前实现中这两个标志定义于 flax/configurations.py通过bool_flag声明并可通过环境变量覆盖# 查看当前状态 print(nnx.set_graph_mode.current_value()) print(nnx.set_graph_updates.current_value()) # 设置值全局 nnx.set_graph_mode(True/False) nnx.set_graph_updates(True/False) # 环境变量 # NNX_GRAPH_MODEtrue/false # NNX_GRAPH_UPDATEStrue/false # 上下文管理器局部生效 with nnx.set_graph_mode(True/False): ... with nnx.set_graph_updates(True/False): ...从源码看set_graph_mode与set_graph_updates定义于 flax/nnx/graphlib.py继承自BaseConfigContext其get_default分别绑定到config.nnx_graph_mode与config.nnx_graph_updatesget_stack对应GRAPH_CONTEXT中独立的 mode 栈与 updates 栈。这意味着它们天然支持设置/回退/上下文局部覆盖三种用法而with块内临时切换也是线程/上下文安全的栈式管理。需要说明的是这是 FLIP 提案的目标状态。截至当前仓库代码nnx_graph_mode与nnx_graph_updates的默认值仍为True见 flax/configurations.py即现阶段图模式仍是默认行为提案规划的未来版本将把默认值翻转。3.2 简化后的变换内部模式新的变换实现相比现有版本大幅简化且同时支持树与图。给定用户函数f大多数简化变换遵循如下模式def transform_wrapper(*args): if graph: args to_tree(args) variables check_no_aliases(argsargs) jax_transform def transformed_f(*args): current, prev snapshot(labeled(argsargs)) if graph: args from_tree(args) out f(*args) if graph: out to_tree(out) check_no_aliases(**current, outout) updates get_updates(current, prev) return out, updates out, updates transformed_f(*args) apply_updates(variables, updates) if graph: out from_tree(out) return out分步解读这个伪代码入口转换to_tree(args)在graphTrue时把图对象转成树表示之后统一交给 JAX 变换同时用check_no_aliases检查输入间无共享引用快照snapshot(labeled(argsargs))记录变换入口处的 Variable 状态基线执行用户函数f(*args)在 JAX 变换内部运行别名检查对输出再做一次check_no_aliases确保输入输出之间无共享引用计算更新get_updates(current, prev)只产生实际发生变化的 Variable 的更新未变的 Variable 不产生更新回写apply_updates(variables, updates)把更新应用到输入 Variable 上最后返回用户输出。支持图的方式很简单进入 JAX 变换前把对象转成树交给用户代码前再从树还原成图。这样 JAX 永远只看到普通 pytree而图的共享语义在边界处被扁平化/还原。该模式与 flax/nnx/transforms/transforms.py 中checkify、cond、switch等变换的实际实现一致内部确实使用了extract.to_tree2/extract.from_tree2、extract.snapshot、extract.check_no_aliases、extract.get_updates、extract.apply_updates这一组工具函数说明提案描述的简化变换骨架已在仓库中落地。四、向后兼容两种迁移路径当 Tree Mode 成为默认行为后依赖图、图更新与前缀过滤器的存量代码将停止工作。提案给出两种移植方式。4.1 路径一回退默认配置在 import 之后显式恢复图模式from flax import nnx ... nnx.set_graph_mode(True) nnx.set_graph_updates(True)4.2 路径二使用 nnx.compat 兼容模块旧版变换 API 会以nnx.compat模块的形式保留实现为把graphTrue、graph_updatesTrue固化的偏函数partialnnx.compat.split partial(nnx.split, graphTrue) ... nnx.compat.jit partial(nnx.jit, graphTrue, graph_updatesTrue) ...移植存量代码只需机械替换nnx.split→nnx.compat.splitnnx.jit→nnx.compat.jit…其余变换同理这一设计已在仓库中实现见 flax/nnx/compat.pynnx.compat模块对 graphlibsplit、state、clone、graphdef、flatten、iter_graph、recursive_map、cached_partial、module 工具view、iter_modules等、rnglibsplit_rngs、fork_rngs、reseed、backup_keys、以及全部变换jit、shard_map、grad、value_and_grad、custom_vjp、vjp、jvp、remat、vmap、scan、pmap、while_loop、fori_loop、eval_shape、cond、switch、checkify、get_abstract_model都用functools.partial固定为graphTrue必要时graph_updatesTrue。其模块 docstring 也明确说明compat 模块提供了默认使用旧图模式实现的 NNX API 包装通过把默认值改为graphTrue与graph_updatesTrue实现。五、破坏性变更与改写指南5.1 移除前缀过滤器Prefix Filters依赖StateAxes、StateSharding、DiffState等前缀过滤器的代码需要重构——JAX 没有等价机制这些过滤器当初是为了简化 Linen 迁移而引入的。解决方案是用split/merge创建状态分组再把每个分组以对应树前缀传给 JAX 变换。旧代码state_axes nnx.StateAxes({some_filter: 0, ...: None}) nnx.vmap(in_axisstate_axes, graphTrue, graph_updatesTrue) def f(model): ...新代码先用之前的过滤器把 model 拆成两个状态组一个向量化、一个广播作为独立参数传入再在变换内部用merge重建 modelgraphdef, vectorized, broadcasted nnx.split(model, some_filter, ...) nnx.vmap(in_axis(0, None)) def f(vectorized, broadcasted): model nnx.merge(graphdef, vectorized, broadcasted) ...这正是前缀过滤器在底层的大致实现方式——拆组、分别映射、再合并现在它被显式化到用户代码里。5.2 nnx.grad 的两处变化nnx.grad的语义将改变两点第一个参数不再默认只对Param求导旧实现默认使用前缀过滤器DiffState(0, Param)NNX Pytree/Module 类型的梯度不再返回State现在遵循 JAX 惯例返回与输入相同的类型。旧代码隐式依赖默认过滤器def loss_fn(model: Foo): ... # 内部使用 argnumsnnx.DiffState(0, nnx.Param) grads nnx.grad(loss_fn)(model)新代码若想避免对不可微状态求梯度必须显式split/mergegraphdef, params, nondiff nnx.split(model, nnx.Param, ...) def loss_fn(params, nondiff): model nnx.merge(graphdef, params, nondiff) ... # 使用 argnums0 grads nnx.grad(loss_fn)(params, nondiff)如果不存在不可微状态可以直接传入model但梯度类型将与输入同型def loss_fn(model: Foo): ... # 使用 argnums0 grads: Foo nnx.grad(loss_fn)(model)5.3 nnx.custom_vjp 语义对齐 JAX旧版nnx.custom_vjp有两个特殊行为backward 函数返回Variable 更新梯度m_updates_g与输出梯度nnx.Pytree/Module对象的切向量tangent类型为nnx.State。以拥有x: Param、y: Param两个属性的FooModule 为例旧代码nnx.custom_vjp def f(m: Foo): return jnp.sin(m.x) * m.y def f_fwd(m: Foo): return f(m), (jnp.cos(m.x), jnp.sin(m.x), m) def f_bwd(res, g): (m_updates_g,), out_g g cos_x, sin_x, m res m_g: nnx.State nnx.clone(m_updates_g) # 创建副本 m_g[x][...] cos_x * out_g * m.y m_g[y][...] sin_x * out_g return (m_g,) # State 梯度新代码不再返回 Variable 更新的梯度切向量类型与输入类型相同Foo与jax.custom_vjp行为一致nnx.custom_vjp def f(m: Foo): return jnp.sin(m.x) * m.y def f_fwd(m: Foo): return f(m), (jnp.cos(m.x), jnp.sin(m.x), m) def f_bwd(res, g): # 不再有 updates 的梯度 cos_x, sin_x, m res m_g: Foo nnx.clone(m) # 创建副本 m_g.x[...] cos_x * g * m.y m_g.y[...] sin_x * g return (m_g,) # Foo 梯度注意为避免信息丢失新版nnx.custom_vjp内不允许更新可微的 Variable。5.4 transform_metadata 迁移为独立变换旧版 NNX 变换如vmap、scan带有transform_metadata元数据参数用于更新分片sharding元数据。新的简化实现不再支持该参数改为引入独立的nnx.transform_metadata变换来保持同样的行为。旧代码nnx.split_rngs(8) nnx.vmap(in_axes0, out_axes0, transform_metadata{nnx.PARTITION_NAME: din}) class create_stack(rngs): # din 被加入 out_sharding 元数据 return nnx.Variable(rngs.uniform((16,)), out_sharding(dout,)) v_stack create_stack(nnx.Rngs(0)) assert v_stack.shape (8, 16) assert v_stack.out_shardings (din, dout)新代码把transform_metadata抽成独立的、可插入的变换层nnx.split_rngs(8) nnx.vmap(in_axes0, out_axes0) nnx.transform_metadata(in_axes0, out_axes0, partitiondin) class create_stack(rngs): # din 被加入 out_sharding 元数据 return nnx.Variable(rngs.uniform((16,)), out_sharding(dout,)) v_stack create_stack(nnx.Rngs(0)) assert v_stack.shape (8, 16) assert v_stack.out_shardings (din, dout)nnx.transform_metadata接受in_axes与out_axes它们必须与对应变换如nnx.vmap传入的轴值保持一致。该变换已存在于仓库中见 flax/nnx/transforms/iteration.py。5.5 Module.sow改用 nnx.capture 提取中间值旧版Module.sow依赖图更新在计算过程中捕获中间值并传播到外部常与nnx.pop配合提取中间结果class Foo(nnx.Module): def __call__(self, x): self.sow(nnx.Intermediate, y_mean, jnp.mean(x)) return x model Foo() result model(x) intermediates nnx.pop(model, nnx.Intermediate) # 提取中间值在不使用图更新的前提下提案新增了nnx.captureAPI提供类似的工作流class Foo(nnx.Module): def __call__(self, x): self.sow(nnx.Intermediate, y_mean, jnp.mean(x)) return x model Foo() result, intermediates nnx.capture(model, nnx.Intermediate)(x)一般地nnx.capture接受一个函数或 Module 作为被变换对象、一个要收集的nnx.Variable子类以及可选的init参数用于初始化被收集的状态该状态存放在nnx.Variable对象内。nnx.capture会在每个Module实例上创建__captures__: tuple[Variable, ...]属性其中的每个 Variable 都含一个字典由sow与perturb填充。从源码看capture定义于 flax/nnx/module.py签名含fn函数、Module 实例或绑定方法、*var_types、init、method_outputs返回包装后的函数其结果为(result, *intermediates)元组若method_outputs提供还会自动以指定 Variable 类型 sow 每个方法含子模块的输出。pop工具则仍保留在 flax/nnx/graphlib.py。5.6 Module.perturb中间值梯度提取的新写法旧版Module.perturb用于提取中间值的梯度分两步先运行一次模块初始化扰动perturbation状态再把扰动状态作为可微目标传给grad。class Model(nnx.Module): def __call__(self, x): x self.perturb(grad_of_x, x) ... return y # 旧代码 nnx.jit def train_step(model, optimizer, x, y): model(x) # 初始化扰动状态 def loss_fn(model): y_pred model(x) return jnp.mean((y_pred - y) ** 2) diff_state nnx.DiffState(0, (nnx.Param, nnx.Perturbation)) grads nnx.grad(loss_fn, argnumsdiff_state)(model) grads, interm_grads nnx.state(grads, nnx.Param, nnx.Perturbation) optimizer.update(model, grads) nnx.pop(model, nnx.Perturbation) # 清理扰动 return interm_grads新模式可以在扰动初始化和前向传播两处都使用nnx.capture把perturbs状态作为独立参数显式传递并用argnums指明两个参数都可微# 新代码 nnx.jit def train_step(model, optimizer, x, y): _, perturbs nnx.capture(model, nnx.Perturbation)(x) # 初始化扰动 def loss_fn(model, perturbs): y_pred nnx.capture(model, initperturbs)(x) return jnp.mean((y_pred - y) ** 2) grads, interm_grads nnx.grad(loss_fn, argnums(0, 1))(model, perturbs) optimizer.update(model, grads) return interm_grads关键差异在于新写法不再依赖DiffState前缀过滤器与nnx.pop清理perturbs成为一等参数可微性由argnums显式控制中间值梯度与参数梯度由nnx.grad一次性返回。六、迁移速查表旧用法新用法nnx.split/nnx.jit等依赖图模式nnx.compat.split/nnx.compat.jit等或启动时nnx.set_graph_mode(True)nnx.set_graph_updates(True)nnx.StateAxes/StateSharding/DiffState前缀过滤器用nnx.split拆组 nnx.merge重组各组独立传参nnx.grad默认DiffState(0, Param)显式split(model, nnx.Param, ...)与argnums0梯度类型与输入同型nnx.custom_vjpState切向量 updates 梯度与jax.custom_vjp对齐切向量与输入同型不返回 updates 梯度不可微 Variable 禁止在内部更新变换的transform_metadata参数插入独立的nnx.transform_metadata(in_axes..., out_axes..., partition...)变换sownnx.pop提取中间值result, intermediates nnx.capture(model, nnx.Intermediate)(x)perturbDiffState提取中间值梯度nnx.capture初始化perturbs状态nnx.grad(loss_fn, argnums(0, 1))七、结语Tree Mode NNX 是 NNX 走向与 JAX 同构的关键一步通过把 Module 视为无状态 pytree、用split/merge显式化状态分组、用nnx.capture替代图更新传播NNX 在保留图模式重要能力状态追踪与共享引用的同时大幅收敛了 API 面与内部复杂度。对普通用户而言Tree Mode 意味着更少的概念、更透明的变换语义与更接近原生 JAX 的开发体验对库维护者而言树与图共享同一套实现骨架也让未来优化如利用jax.tree.*的高效遍历成为可能。如果你正在维护存量 NNX 代码建议优先按上文速查表逐一替换需要图语义时走nnx.compat需要新语义时改写为split/merge/capture组合并以nnx.set_graph_mode/nnx.set_graph_updates或环境变量NNX_GRAPH_MODE/NNX_GRAPH_UPDATES控制全局默认行为平滑过渡到 Tree Mode。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考