import optax # Basic Optimizers optax.sgd(learning_rate, momentum=0.9) optax.adam(learning_rate) optax.adamw(learning_rate, weight_decay) # Gradient Clipping optax.clip_by_global_norm(max_norm) optax.clip(max_delta) # Learning Rate Schedules optax.warmup_cosine_decay_schedule(init, peak, warmup_steps, decay_steps) optax.exponential_decay(init, transition_steps, decay_rate) optax.piecewise_constant_schedule(init, boundaries_and_scales) # Composing optax.chain(transform1, transform2, transform3) # Building Blocks optax.scale_by_adam() # Adam's moment estimation optax.add_decayed_weights(wd) # Weight decay optax.scale(-lr) # Scale by learning rate optax.scale_by_schedule(fn) # Dynamic scaling # With NNX optimizer = nnx.Optimizer(model, tx, wrt=nnx.Param) optimizer.update(model, grads) __ __