import optax # Each transformation modifies the gradients, then passes them to the next optimizer = optax.chain( optax.clip_by_global_norm(1.0), # First: clip gradients optax.scale_by_adam(), # Second: compute Adam scaling optax.add_decayed_weights(0.01), # Third: add weight decay optax.scale(-0.001) # Fourth: scale by -learning_rate ) __ __