import optax # A simple Adam optimizer optimizer = optax.adam(learning_rate=0.001) # Adam with gradient clipping optimizer = optax.chain( optax.clip_by_global_norm(1.0), # First, clip gradients optax.adam(learning_rate=0.001) # Then, apply Adam ) # Adam with weight decay (AdamW) optimizer = optax.adamw(learning_rate=0.001, weight_decay=0.01) __ __