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

Flax.nnx迁移学习冻结参数报错:参数不匹配问题求助

问题排查:Flax.nnx迁移学习中优化器参数不匹配错误

问题场景

尝试用Flax.nnx实现迁移学习,冻结nnx.Linear的kernel参数,仅优化bias。运行代码后出现参数不匹配错误。

原代码

from jax import numpy as jnp
from jax import random
from flax import nnx
import optax
from matplotlib import pyplot as plt

def f(x,m=2.234,b=-1.123):
    return m*x+b

def compute_loss(model, inputs, obs):
    prediction = model(inputs)
    error = obs - prediction
    loss = jnp.mean(error ** 2)
    mae = jnp.mean(jnp.abs(error ) )
    return loss, mae

if __name__ == '__main__':
    shape = (2,55,1)
    epochs = 123

    rngs = nnx.Rngs(123)
    model = nnx.Linear( 1, 1, rngs=rngs )

    model.kernel.value = jnp.array([[2.0]]) #load pretrained kernel  

    skey = rngs.params()
    xx = random.uniform( skey, shape, minval=-10, maxval=10 ) 
    obs1,obs2 = f(xx)
    x1,x2 = xx
    
    loss_grad = nnx.value_and_grad(compute_loss, has_aux = True)
    @nnx.scan(
        in_axes=(nnx.Carry,None,None,),
        out_axes=(nnx.Carry,0),
        length=epochs
    )
    def optimizer_scan( optimizer, x, obs ):
        (loss,mae), grads = loss_grad( optimizer.model, x, obs )        
        optimizer.update( grads )
        return optimizer, (loss,mae)

    transfer_params = nnx.All(nnx.PathContains("bias"))
    optimizer_transfer = nnx.Optimizer(model, optax.adam(learning_rate=1e-3), wrt = transfer_params)

    optimizer, (losses,maes) = optimizer_scan( optimizer_transfer, x1, obs1 )

    print( ' AFTER TRAINING' )
    print( 'training loss:', losses[-1] )

    y1,y2 = optimizer.model(xx)
    error = obs2-y2
    loss = jnp.mean( error*error )
    print( 'test loss:',loss )
    print( 'm approximation:', optimizer.model.kernel.value )
    print( 'b approximation:', optimizer.model.bias.value )

报错信息

ValueError: Mismatch custom node data: ('bias', 'kernel') != ('bias',); value: State({
  'bias': VariableState(
    type=Param,
    value=Traced<ShapedArray(float32[1])>with<DynamicJaxprTrace(level=1/0)>
  )
}).

报错原因

nnx.value_and_grad默认会计算模型中所有Param类型变量的梯度(包括kernel和bias),但配置的optimizer_transfer仅跟踪bias参数。调用optimizer.update(grads)时,梯度集合包含kernel和bias的梯度,与优化器跟踪的参数集合(仅bias)不匹配,导致参数不匹配错误。

解决步骤

修改nnx.value_and_grad的调用,传入wrt参数指定仅计算bias的梯度,确保梯度集合与优化器跟踪的参数一致:

将:

loss_grad = nnx.value_and_grad(compute_loss, has_aux = True)

替换为:

transfer_params = nnx.All(nnx.PathContains("bias"))
loss_grad = nnx.value_and_grad(compute_loss, has_aux = True, wrt=transfer_params)

修正后完整代码

from jax import numpy as jnp
from jax import random
from flax import nnx
import optax
from matplotlib import pyplot as plt

def f(x,m=2.234,b=-1.123):
    return m*x+b

def compute_loss(model, inputs, obs):
    prediction = model(inputs)
    error = obs - prediction
    loss = jnp.mean(error ** 2)
    mae = jnp.mean(jnp.abs(error ) )
    return loss, mae

if __name__ == '__main__':
    shape = (2,55,1)
    epochs = 123

    rngs = nnx.Rngs(123)
    model = nnx.Linear( 1, 1, rngs=rngs )

    model.kernel.value = jnp.array([[2.0]]) #load pretrained kernel  

    skey = rngs.params()
    xx = random.uniform( skey, shape, minval=-10, maxval=10 ) 
    obs1,obs2 = f(xx)
    x1,x2 = xx
    
    # 先定义参数选择器,再传给value_and_grad和optimizer
    transfer_params = nnx.All(nnx.PathContains("bias"))
    loss_grad = nnx.value_and_grad(compute_loss, has_aux = True, wrt=transfer_params)
    
    @nnx.scan(
        in_axes=(nnx.Carry,None,None,),
        out_axes=(nnx.Carry,0),
        length=epochs
    )
    def optimizer_scan( optimizer, x, obs ):
        (loss,mae), grads = loss_grad( optimizer.model, x, obs )        
        optimizer.update( grads )
        return optimizer, (loss,mae)

    optimizer_transfer = nnx.Optimizer(model, optax.adam(learning_rate=1e-3), wrt = transfer_params)

    optimizer, (losses,maes) = optimizer_scan( optimizer_transfer, x1, obs1 )

    print( ' AFTER TRAINING' )
    print( 'training loss:', losses[-1] )

    y1,y2 = optimizer.model(xx)
    error = obs2-y2
    loss = jnp.mean( error*error )
    print( 'test loss:',loss )
    print( 'm approximation:', optimizer.model.kernel.value )
    print( 'b approximation:', optimizer.model.bias.value )

内容的提问来源于stack exchange,提问作者jworrell

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 10:04:53