# AdamW: Adam with decoupled weight decay optimizer = optax.adamw(learning_rate=0.001, weight_decay=0.01) # Or build it manually: optimizer = optax.chain( optax.scale_by_adam(), optax.add_decayed_weights(weight_decay=0.01), optax.scale(-0.001) ) __ __