from flax import nnx import jax class FineTuneModel(nnx.Module): def __init__(self, *, rngs: nnx.Rngs): # Pretrained backbone (we want small LR) self.backbone = nnx.Linear(784, 256, rngs=rngs) # New classification head (we want larger LR) self.head = nnx.Linear(256, 10, rngs=rngs) def __call__(self, x): x = nnx.relu(self.backbone(x)) return self.head(x) model = FineTuneModel(rngs=nnx.Rngs(0)) # Create parameter labels def label_fn(path, _): """Assign labels based on parameter path.""" path_str = '/'.join(str(p) for p in path) if 'backbone' in path_str: return 'backbone' else: return 'head' # Get params and create labels params = nnx.state(model, nnx.Param) param_labels = jax.tree_util.tree_map_with_path(label_fn, params) # Create the multi-transform optimizer tx = optax.multi_transform( transforms={ 'backbone': optax.adam(learning_rate=1e-5), 'head': optax.adam(learning_rate=1e-3), }, param_labels=param_labels ) optimizer = nnx.Optimizer(model, tx, wrt=nnx.Param) __ __