import optax def create_optimizer( peak_learning_rate: float = 1e-3, warmup_steps: int = 1000, total_steps: int = 100000, weight_decay: float = 0.01, max_grad_norm: float = 1.0, end_learning_rate: float = 1e-5, ): """Create a production-ready optimizer with all the bells and whistles.""" # Learning rate schedule: warmup then cosine decay schedule = optax.warmup_cosine_decay_schedule( init_value=0.0, peak_value=peak_learning_rate, warmup_steps=warmup_steps, decay_steps=total_steps - warmup_steps, end_value=end_learning_rate, ) # Build the optimizer chain optimizer = optax.chain( # 1. Gradient clipping (prevent explosions) optax.clip_by_global_norm(max_grad_norm), # 2. Adam scaling optax.scale_by_adam(b1=0.9, b2=0.999, eps=1e-8), # 3. Weight decay (applied after Adam scaling) optax.add_decayed_weights(weight_decay), # 4. Scale by negative learning rate (makes it gradient descent) optax.scale_by_schedule(lambda step: -schedule(step)), ) return optimizer # Usage tx = create_optimizer( peak_learning_rate=1e-3, warmup_steps=1000, total_steps=50000, weight_decay=0.01, max_grad_norm=1.0, ) optimizer = nnx.Optimizer(model, tx, wrt=nnx.Param) __ __