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
相关产品推荐
相关产品推荐

