from flax import nnx import optax # Create a mask that's True for kernel/weight parameters, False for biases def create_weight_decay_mask(params): def should_decay(path, _): # path is a tuple like ('linear1', 'kernel') or ('linear1', 'bias') return 'kernel' in path or 'weight' in path return jax.tree_util.tree_map_with_path( lambda path, _: should_decay(path, _), params ) # Use optax.masked to apply weight decay selectively optimizer = optax.chain( optax.clip_by_global_norm(1.0), optax.adamw(learning_rate=0.001, weight_decay=0.01, mask=create_weight_decay_mask) ) __ __