schedule = optax.warmup_cosine_decay_schedule( init_value=0.0, peak_value=1e-4, warmup_steps=2000, decay_steps=98000 ) tx = optax.chain( optax.clip_by_global_norm(1.0), optax.adamw(learning_rate=schedule, weight_decay=0.1) ) __ __