import matplotlib.pyplot as plt import jax.numpy as jnp schedule = optax.warmup_cosine_decay_schedule( init_value=0.0, peak_value=0.001, warmup_steps=1000, decay_steps=9000 ) steps = jnp.arange(10000) lrs = [schedule(step) for step in steps] plt.figure(figsize=(10, 4)) plt.plot(steps, lrs) plt.xlabel('Step') plt.ylabel('Learning Rate') plt.title('Warmup Cosine Decay Schedule') plt.grid(True) plt.show() __ __