Flax NNX Hijax 模式:在 JAX 变换中原地更新可变状态的实验性指南
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
Hijax 是 Flax NNX 为nnx.Variable提供的一种实验性状态传播机制,它借助 JAX 的 Hijax 框架(jax.experimental.hijax),让变量能够以引用(ref)的方式直接参与jax.jit、jax.grad等变换:在被 JIT 编译的函数内部,对变量的原地修改(如v[...] += 1、metrics['y_mean'] = ...)会被自动追踪并回写到外部的可变对象上。读完本文,你将掌握如何通过nnx.var_defaults(hijax=True)开启 Hijax 变量模式、如何在训练循环中结合nnx.split/nnx.merge与jax.value_and_grad使用它,并了解其可变性(mutability)、ref 支持、重复引用(aliasing)等边界与限制。
什么是 Hijax 变量模式
在 Flax NNX 中,nnx.Variable默认以"树模式"(tree-mode)参与 JAX 变换:变量是普通 pytree 叶子,跨函数边界的状态传递需要显式的nnx.split/nnx.merge。Hijax 模式则改变这一行为——它把变量包装成 JAX 的 Hijax 类型(HijaxVariable),使变量在 JIT 编译的函数内部表现为可变引用,任何原地写入都会经由 JAX 的get_variable/set_variable原语被捕获,并在函数执行结束后自动同步回原始变量。
在 flax/nnx/variablelib.py 中可以看到,Flax 直接复用了jax.experimental.hijax(源码第 34 行),并定义了VariableContext中独立的variable_hijax_stack(第 67 行)来管理 hijax 模式的默认开关。该开关的全局默认值来自配置项flax_hijax_variable,在 flax/configurations.py 第 304–308 行定义,默认False,即默认情况下变量不启用 Hijax 模式:
flax_hijax_variable = bool_flag( name='flax_hijax_variable', default=False, help='Whether to enable HiJAX support for `nnx.Variable`.', )也就是说,Hijax 是一个必须显式开启的实验特性,你可以通过环境变量FLAX_HIJAX_VARIABLE=True或运行时调用nnx.var_defaults(hijax=True)来启用。
快速上手:最小训练示例
docs_nnx/hijax/index.rst给出了一个完整的最小示例。开启 Hijax 模式后,模型内部的nnx.Param、nnx.BatchNorm的 running statistics、nnx.Dropout的 rng 状态等都会在 JIT 编译的train_step内被当作可变引用处理,训练循环因此可以写得更直接:
from flax import nnx import optax nnx.var_defaults(hijax=True) class Model(nnx.Module): def __init__(self, din, dmid, dout, rngs: nnx.Rngs): self.linear = nnx.Linear(din, dmid, rngs=rngs) self.bn = nnx.BatchNorm(dmid, rngs=rngs) self.dropout = nnx.Dropout(0.2) self.linear_out = nnx.Linear(dmid, dout, rngs=rngs) def __call__(self, x, rngs): x = nnx.relu(self.dropout(self.bn(self.linear(x)), rngs=rngs)) return self.linear_out(x) model = Model(2, 64, 3, rngs=nnx.Rngs(0)) # eager initialization optimizer = nnx.Optimizer(model, optax.adam(1e-3), wrt=nnx.Param) @jax.jit def train_step(model, optimizer, rngs, x, y): graphdef, params, nondiff = nnx.split(model, nnx.Param, ...) def loss_fn(params): model = nnx.merge(graphdef, params, nondiff) return ((model(x, rngs) - y) ** 2).mean() loss, grads = jax.value_and_grad(loss_fn)(nnx.vars_as(params, mutable=False)) optimizer.update(model, grads) # in-place updates return loss这里有几个值得注意的细节:
nnx.var_defaults(hijax=True)会修改进程内的默认行为,其返回值是VarDefaultsContext,支持作为上下文管理器使用(详见 variablelib.py 第 133–192 行的VarDefaults与VarDefaultsContext实现)。loss_fn内部用nnx.merge重建模型后求值,jax.value_and_grad需要参数不可变,因此用nnx.vars_as(params, mutable=False)传入(vars_as是with_vars的废弃别名,见 graphlib.py 第 2980 行)。- 优化器更新
optimizer.update(model, grads)是原地操作,这正是 Hijax 引用语义带来的便利。
docs_nnx/hijax/hijax.ipynb中还提供了一个等价的简化版本(不带 BatchNorm / Dropout),只用一个nnx.Linear(2, 3)模型即可验证:
nnx.var_defaults(hijax=True) rngs = nnx.Rngs(0) model = nnx.Linear(2, 3, rngs=rngs) optimizer = nnx.Optimizer(model, optax.adamw(1e-2), wrt=nnx.Param) @jax.jit def train_step(x, y): loss_fn = lambda m: jnp.mean((m(x) - y) ** 2) loss, grads = jax.value_and_grad(loss_fn)(nnx.with_vars(model, mutable=False)) optimizer.update(model, grads) return loss x, y = rngs.uniform((4, 2)), rngs.uniform((4, 3)) for _ in range(3): print(train_step(x, y))Hijax Variable 的状态传播语义
Hijax 模式下变量值在 JIT 边界上的传播由HijaxVariable类型承载,其实现集中在 flax/nnx/variablelib.py 第 270–740 行附近:_new_hijax_from_variable(第 286 行)、_set_hijax_state(第 343 行)与_get_hijax_state(第 392 行)负责在Variable与HijaxVariable之间双向转换,而_as_hijax_property/_as_hijax_attribute/_as_hijax_method(第 470–530 行)则把变量的属性、方法与运算符逐一代理到 Hijax 类型上。
状态传播:标量变量
最基础的用法是让一个标量变量在 JIT 函数内自增,并在函数外观察到更新:
v = nnx.Variable(jnp.array(0), hijax=True) @jax.jit def inc(v): v[...] += 1 print(v[...]); inc(v); print(v[...]) # 0, 1对同一个inc打印 jaxpr,可以看到 Hijax 变量的真实运行机制:JAX 为它插入了get_variable读取当前值、执行add之后再用set_variable写回,变量的元数据(如('hijax', True), ('mutable', True), ('ref', False))会被记录在 pytree 的 treedef 中:
{ lambda ; a:Variable(). let b:i32[] = get_variable[avals=..., treedef=PyTreeDef(CustomNode(Variable[..., ('hijax', True), ...]))] a c:i32[] = add b 1 _ = set_variable[...] a c in () }状态传播:pytree 值
变量的值本身可以是任意 pytree,Hijax 同样支持对其中元素的原地更新:
v = nnx.Variable({'a': jnp.array(0), 'b': jnp.array(2)}, hijax=True) @jax.jit def inc_and_double(v): v['a'] += 1 v['b'] *= 2 print(v) # {'a': 0, 'b': 2} inc_and_double(v) print(v) # {'a': 1, 'b': 4}动态状态结构:在 JIT 内写入新键
Hijax 的另一个独特能力是可以在 JIT 编译的函数内动态扩展状态结构——比如收集指标时,往一个空的 metrics 字典里写入键:
rngs = nnx.Rngs(0) x = rngs.uniform((4, 5)) w = rngs.normal((5, 3)) metrics = nnx.Variable({}, hijax=True) @jax.jit def linear(x, w, metrics: nnx.Variable): y = x @ w metrics['y_mean'] = jnp.mean(y) return y print("Before:", metrics) # {} y = linear(x, w, metrics) print("After:", metrics) # {'y_mean': -1.178...}函数执行后,metrics变量自动获得了新键'y_mean'。这种能力在传统的树模式 NNX 中需要额外的图结构操作才能实现,是 Hijax 引用语义的典型优势。
可变性(Mutability)控制
Hijax 变量与普通 NNX 变量一样受mutable标志约束。nnx.with_vars可以在不改动原始对象的前提下,返回一份带有新元数据的变量副本;其签名(graphlib.py 第 2932–2941 行)支持hijax、ref、mutable、only(过滤条件)等参数:
def with_vars( node: A, /, *, hijax: bool | None = None, ref: bool | None = None, mutable: bool | None = None, only: filterlib.Filter = ..., allow_duplicates: bool = False, ) -> A:对模型整体设置mutable,可以批量切换其内部所有变量的可变性:
class Linear(nnx.Module): def __init__(self, in_features, out_features, rngs: nnx.Rngs): self.kernel = nnx.Param(rngs.normal((in_features, out_features))) def __call__(self, x): return x @ self.kernel model = Linear(1, 3, rngs=nnx.Rngs(0)) print(nnx.with_vars(model, mutable=False)) # kernel 的 mutable=False print(nnx.with_vars(model, mutable=True)) # kernel 的 mutable=True尝试对不可变变量进行原地写操作会抛出ImmutableVariableError:
v = nnx.Variable(jnp.array(0)) v_immut = nnx.with_vars(v, mutable=False) assert not v_immut.mutable try: v_immut[...] += 1 # raises an error except Exception as e: print(f"{type(e).__name__}: {e}") # ImmutableVariableError: Cannot mutate Variable as it is marked as immutable.这正是训练循环里给jax.grad传入nnx.with_vars(params, mutable=False)的原因:Hijax 可变的引用(ref)语义与jax.grad对纯函数的求值要求冲突,必须先切换为不可变视图。
Ref 支持与原始值访问
除了mutable,Hijax 变量还支持ref标志。当ref=True时,变量在 JIT 内部表现为 JAX 的Ref对象,可以直接用get_raw_value()读取未被追踪的原始值:
v = nnx.Variable(jnp.array(0)) v_ref = nnx.with_vars(v, ref=True) assert v_ref.ref print(v_ref) # Variable(value=Array(0, dtype=int32, weak_type=True), hijax=True, ref=True) print(v_ref.get_raw_value()) # Ref(0, dtype=int32, weak_type=True)ref与mutable可以相互转换:把ref=True的变量切回mutable=False会失去 ref 语义(打印时出现had_ref=True),再次切回mutable=True则恢复 ref:
v_immut = nnx.with_vars(v_ref, mutable=False) assert not v_immut.ref print("immutable =", v_immut) # had_ref=True, mutable=False v_ref = nnx.with_vars(v_immut, mutable=True) assert v_ref.ref print("mutable =", v_ref) # ref=True, mutable=True从实现上看,HijaxVariable内部保存_treedef、_leaves、_var_type、_ref四个字段(variablelib.py 第 582–587 行),ref语义正是由_ref标志决定的;而Variable的hijax/ref/mutable元数据由 variablelib.py 第 1097–1106 行附近的hijax属性统一暴露。
综合示例:带 BatchNorm 与 Dropout 的训练循环
docs_nnx/hijax/hijax.ipynb的 "Examples" 小节给出了一个更贴近真实场景的完整示例,把上述机制组合在一起。定义一个由Linear+BatchNorm+Dropout构成的Block:
class Block(nnx.Module): def __init__(self, din, dmid, dout, rngs: nnx.Rngs): self.linear = Linear(din, dmid, rngs=rngs) self.bn = nnx.BatchNorm(dmid, use_running_average=False, rngs=rngs) self.dropout = nnx.Dropout(0.1, deterministic=False, rngs=rngs) self.linear_out = Linear(dmid, dout, rngs=rngs) def __call__(self, x): x = nnx.gelu(self.dropout(self.bn(self.linear(x)))) return self.linear_out(x)训练循环与前面一致:开启 hijax 默认模式后,BatchNorm的 running statistics 与Dropout的 rng 状态会在 JIT 编译函数内被自动原地更新,而参数通过nnx.split分离后以不可变视图交给jax.value_and_grad求梯度,再由optimizer.update原地更新:
# hijax Variables by default model = Block(2, 64, 3, rngs=nnx.Rngs(0)) optimizer = nnx.Optimizer(model, optax.adam(1e-3), wrt=nnx.Param) @jax.jit def train_step(model, optimizer, x, y): graphdef, params, nondiff = nnx.split(model, nnx.Param, ...) def loss_fn(params): model = nnx.merge(graphdef, params, nondiff) return ((model(x) - y) ** 2).mean() loss, grads = jax.value_and_grad(loss_fn)(nnx.with_vars(params, mutable=False)) optimizer.update(model, grads) return loss for _ in range(3): loss = train_step(model, optimizer, x=jnp.ones((10, 2)), y=jnp.ones((10, 3))) print(f"{loss = !s}")局限:Scan Over Layers 尚不支持
原文档的 "Scan Over Layers" 小节以注释形式保留了一个尚未可用的示例:用jax.vmap+nnx.as_immutable_vars/nnx.as_mutable_vars创建层栈,再用jax.lax.scan逐层前向传播的写法,目前TODO 标注为 "does not work with hijax yet"(尚不支持 Hijax)。如果你需要按层堆叠模型,当前仍需退回传统的树模式 NNX 变换。
已知限制与注意事项
Hijax 是实验特性,docs_nnx/hijax/index.rst 与 docs_nnx/hijax/hijax.ipynb 明确列出以下限制,使用时需格外注意。
可变的函数输出
JIT 函数的返回值不能包含可变(mutable)的 Hijax 变量。直接在@jax.jit函数内创建并返回模型会报错("mutable hitypes should use lo_ty_qdd instead"):
@jax.jit def create_model(rngs): return Block(2, 64, 3, rngs=rngs) try: model = create_model(nnx.Rngs(0)) except Exception as e: print(f"Error:", e) # Error: mutable hitypes should use lo_ty_qdd instead绕过方式是在 JIT 内部创建时临时关闭 hijax(返回不可变视图),再在外部重新开启:
@jax.jit def create_model(rngs): return nnx.with_vars((Block(2, 64, 3, rngs=rngs)), hijax=False) model = nnx.with_vars(create_model(nnx.Rngs(0)), hijax=True) print("model.linear =", model.linear)引用共享(Aliasing)不传播
如果多个属性引用同一个Variable对象(即存在别名),Hijax 模式不会把 JIT 内的写入传播回所有别名。以下代码虽然在 JAX 侧不会报错,但has_shared.b的更新不会生效:
class HasShared(nnx.Pytree): def __init__(self): self.a = nnx.Variable(jnp.array(0)) self.b = self.a @jax.jit def g(has_shared): has_shared.a[...] = 5 has_shared = HasShared() print(get_error(g, has_shared)) # 无异常(JAX 侧暂不报错) print(has_shared) # a 与 b 的更新都不会传播可以用nnx.find_duplicates检测这种别名问题,它会返回重复引用的路径对:
print("Duplicates found:") if (all_duplicates := nnx.find_duplicates(has_shared)): for duplicates in all_duplicates: print("-", duplicates) # - [('a',), ('b',)]find_duplicates的实现位于 flax/nnx/graphlib.py 第 3556 行,它遍历节点图找出指向同一对象的多条路径。值得注意的是,nnx.with_vars在默认情况下也会先调用find_duplicates检查重复引用,发现别名时抛出ValueError(见 graphlib.py 第 2959–2968 行),除非显式传入allow_duplicates=True。
不过,若别名结构先经过nnx.split/nnx.merge重建(即先序列化成 graphdef + state 再合并),引用关系会被重新建立,更新可以正常传播:
@jax.jit def h(graphdef, state): has_shared = nnx.merge(graphdef, state) has_shared.a[...] = 5 graphdef, state = nnx.split(has_shared) h(graphdef, state) print(has_shared) # a 与 b 均为 5,更新正常传播全局开关的清理
由于nnx.var_defaults(hijax=True)修改的是全局默认值,在多测试用例或 CI 环境中运行后应当恢复原值。文档中的示例使用current_mode = nnx.var_defaults().hijax保存旧值,结束时用nnx.var_defaults(hijax=current_mode)还原:
current_mode = nnx.var_defaults().hijax # 保存当前模式 nnx.var_defaults(hijax=True) # ... 业务代码 ... nnx.var_defaults(hijax=current_mode) # 清理,恢复原模式更推荐的做法是直接利用VarDefaultsContext作为上下文管理器,让模式在with块结束后自动恢复。
总结与适用场景
综合来看,Hijax 变量模式为 Flax NNX 提供了一条"把可变引用直接带进 JIT"的新路径,其核心价值与边界可以归纳如下:
- 开启方式:
nnx.var_defaults(hijax=True)(进程级默认),或全局配置flax_hijax_variable(默认False,见 flax/configurations.py)。 - 核心能力:在
@jax.jit函数内原地更新变量(含 pytree 值、动态新增键),函数外自动同步;配合nnx.with_vars(..., mutable=False)可安全对接jax.value_and_grad。 - 配套 API:
nnx.with_vars(及废弃别名nnx.vars_as)用于切换hijax/ref/mutable元数据;nnx.find_duplicates用于排查引用别名;HijaxVariable的底层实现在 flax/nnx/variablelib.py。 - 限制:JIT 函数不能直接返回可变 Hijax 对象;引用共享(aliasing)在直接 JIT 调用下不传播;
jax.lax.scan按层堆叠暂不支持;所有 API 与行为均以仓库当前实现为准,属于实验特性,后续版本可能调整。
如果需要编写带 BatchNorm / Dropout 的紧凑训练循环,或在 JIT 内收集动态指标,Hijax 模式是一个值得尝试的方向;而在涉及模型输出、层堆叠或对象别名等复杂结构时,建议先用find_duplicates自检,或继续使用传统的树模式 NNX 变换。
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考