import optax from flax import nnx # Define different optimizers for different parameter groups optimizer = optax.multi_transform( transforms={ 'backbone': optax.adam(learning_rate=1e-5), # Small LR for pretrained 'head': optax.adam(learning_rate=1e-3), # Larger LR for new layers }, param_labels=param_labels # A pytree matching params, with string labels ) __ __