如何使用lax.scan更新Flax NNX中nnx.Module的权重?
使用NNX结合lax.scan高效训练神经网络的解决方案
问题描述
我使用Flax的NNX编写了一个基于nnx.Module的神经网络,希望用lax.scan替代for循环来高效训练该网络。但由于scan不允许原地修改,请问如何在每次迭代后更新nnx.Module的权重?
我期望实现类似如下的功能:
policy = neural_network() # <- nnx.Module optimizer = nnx.Optimizer(policy, adamw(0.001)) xs = jnp.zeros(1000, 1000) def backprop(_, x): grad_fn = value_and_grad(some_loss_function) loss, grads = grad_fn(policy, x) optimizer.update(grads) return _, loss _, losses = lax.scan(backprop, None, xs)
但上述方法无法实现权重更新,我希望用lax.scan快速训练神经网络,请问这是否可行?
解决方案
完全可行,核心是把nnx.Module和优化器的状态作为扫描的携带状态传入,而非在闭包中引用(scan不允许闭包内的可变状态)。NNX的模块和优化器支持通过状态提取与合并的方式,将不可变状态传递给scan的迭代函数。
具体实现步骤
- 提取初始状态:从
nnx.Module和优化器中分离出可训练参数与优化器状态,打包成scan的初始携带变量。 - 迭代函数内状态传递:在scan的step函数中,用当前携带的状态重建模块和优化器,完成梯度计算与更新后,将新状态作为返回值传递给下一次迭代。
- 最终状态合并:scan执行完成后,将最终状态合并回原模块和优化器。
完整示例代码
import jax import jax.numpy as jnp import nnx from nnx import optim # 自定义损失函数,根据任务需求替换 def some_loss_function(policy: nnx.Module, x: jnp.ndarray) -> jnp.ndarray: logits = policy(x) return jnp.mean(logits ** 2) # 定义NNX神经网络模块 class neural_network(nnx.Module): def __init__(self, rng: jax.Array): self.dense1 = nnx.Linear(1000, 512, rng=rng) self.dense2 = nnx.Linear(512, 10, rng=rng) def __call__(self, x: jnp.ndarray) -> jnp.ndarray: x = self.dense1(x) x = jax.nn.relu(x) return self.dense2(x) # 初始化模块与优化器 rng = jax.random.PRNGKey(42) policy = neural_network(rng) optimizer = optim.Adamw(learning_rate=0.001).init(policy) # 提取初始状态:模块可训练参数 + 优化器状态 def get_carry_state(policy, optimizer): params = nnx.state(policy, nnx.Param) opt_state = optimizer.state return (params, opt_state) initial_carry = get_carry_state(policy, optimizer) xs = jnp.zeros((1000, 1000)) # 示例输入数据 # 定义scan的迭代step函数 def backprop_step(carry, x): params, opt_state = carry # 用当前状态重建模块和优化器 policy = neural_network(rng) policy = policy.merge(params) optimizer = optim.Adamw(learning_rate=0.001).init(policy) optimizer = optimizer.merge(opt_state) # 计算损失与梯度 grad_fn = jax.value_and_grad(some_loss_function, argnums=0) loss, grads = grad_fn(policy, x) # 更新状态 optimizer.update(grads) new_params = nnx.state(policy, nnx.Param) new_opt_state = optimizer.state # 返回新状态与当前损失 return (new_params, new_opt_state), loss # 执行lax.scan final_carry, losses = jax.lax.scan(backprop_step, initial_carry, xs) # 将最终状态合并回原模块和优化器 final_params, final_opt_state = final_carry policy = policy.merge(final_params) optimizer = optimizer.merge(final_opt_state)
关键注意点
- 禁止闭包引用可变状态:scan的step函数不能依赖外部的模块或优化器实例,必须通过carry参数传递所有需要更新的状态。
- NNX状态操作:
nnx.state()用于提取指定类型的状态(如nnx.Param代表可训练参数),merge()用于将状态合并回模块/优化器,这是NNX处理不可变状态的核心方式。 - 性能优化:若模块结构固定,可提前用
nnx.split()分离模块的结构与状态,在step函数中直接用结构+状态重建模块,减少重复初始化的开销。
内容的提问来源于stack exchange,提问作者elderly
相关产品推荐
相关产品推荐

