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 filters):
StateAxes、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 references);
- Module 视为无状态 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.*API |
graph_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_MODE=true/false # NNX_GRAPH_UPDATES=true/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_updates,get_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(args=args) @jax_transform def transformed_f(*args): current, prev = snapshot(labeled(args=args)) if graph: args = from_tree(args) out = f(*args) if graph: out = to_tree(out) check_no_aliases(**current, out=out) 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)在graph=True时把图对象转成树表示,之后统一交给 JAX 变换;同时用check_no_aliases检查输入间无共享引用; - 快照:
snapshot(labeled(args=args))记录变换入口处的 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模块的形式保留,实现为把graph=True、graph_updates=True固化的偏函数(partial):
nnx.compat.split = partial(nnx.split, graph=True) ... nnx.compat.jit = partial(nnx.jit, graph=True, graph_updates=True) ...移植存量代码只需机械替换:
nnx.split→nnx.compat.splitnnx.jit→nnx.compat.jit- …(其余变换同理)
这一设计已在仓库中实现:见 flax/nnx/compat.py,nnx.compat模块对 graphlib(split、state、clone、graphdef、flatten、iter_graph、recursive_map、cached_partial)、module 工具(view、iter_modules等)、rnglib(split_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固定为graph=True(必要时graph_updates=True)。其模块 docstring 也明确说明:"compat 模块提供了默认使用旧图模式实现的 NNX API 包装,通过把默认值改为graph=True与graph_updates=True实现"。
五、破坏性变更与改写指南
5.1 移除前缀过滤器(Prefix Filters)
依赖StateAxes、StateSharding、DiffState等前缀过滤器的代码需要重构——JAX 没有等价机制(这些过滤器当初是为了简化 Linen 迁移而引入的)。解决方案是用split/merge创建状态分组,再把每个分组以对应树前缀传给 JAX 变换。
旧代码:
state_axes = nnx.StateAxes({some_filter: 0, ...: None}) @nnx.vmap(in_axis=state_axes, graph=True, graph_updates=True) def f(model): ...新代码:先用之前的过滤器把 model 拆成两个状态组,一个向量化、一个广播,作为独立参数传入,再在变换内部用merge重建 model:
graphdef, 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): ... # 内部使用 argnums=nnx.DiffState(0, nnx.Param) grads = nnx.grad(loss_fn)(model)新代码:若想避免对不可微状态求梯度,必须显式split/merge:
graphdef, params, nondiff = nnx.split(model, nnx.Param, ...) def loss_fn(params, nondiff): model = nnx.merge(graphdef, params, nondiff) ... # 使用 argnums=0 grads = nnx.grad(loss_fn)(params, nondiff)如果不存在不可微状态,可以直接传入model,但梯度类型将与输入同型:
def loss_fn(model: Foo): ... # 使用 argnums=0 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_axes=0, out_axes=0, 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_axes=0, out_axes=0) @nnx.transform_metadata(in_axes=0, out_axes=0, partition='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')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, argnums=diff_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, init=perturbs)(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, ...)与argnums=0;梯度类型与输入同型 |
nnx.custom_vjp(State切向量 + updates 梯度) | 与jax.custom_vjp对齐:切向量与输入同型,不返回 updates 梯度,不可微 Variable 禁止在内部更新 |
变换的transform_metadata=参数 | 插入独立的nnx.transform_metadata(in_axes=..., out_axes=..., partition=...)变换 |
sow+nnx.pop提取中间值 | result, intermediates = nnx.capture(model, nnx.Intermediate)(x) |
perturb+DiffState提取中间值梯度 | 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),仅供参考