You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.13 09:01:21