"""Cosine annealing with warmup learning rate scheduler.""" import math import torch.optim as optim from torch.optim.lr_scheduler import LambdaLR from taoTrain.config import TrainingConfig from .registry import register_scheduler @register_scheduler("cosineWarmup") def create_cosine_warmup( optimizer: optim.Optimizer, config: TrainingConfig, num_training_steps: int, ) -> LambdaLR: """ Create a cosine annealing scheduler with optional linear warmup, steady phase, and decay. Three-phase schedule: 1. Linear warmup: 0 → 1.0 (warmup_steps) 2. Steady phase: 1.0 (plateau at peak LR) 3. Cosine decay: 1.0 → min_lr_ratio Args: optimizer: Optimizer instance config: TrainingConfig with scheduler configuration: - warmup_steps: linear warmup duration (overrides warmup_ratio if > 0) - warmup_ratio: warmup as fraction of total steps (default 0.1) - steady_ratio: steady phase as fraction of total steps (default 0.0) - min_lr_ratio: minimum LR at end as fraction of peak (default 0.0) num_training_steps: Total number of training steps Returns: LambdaLR scheduler instance """ scheduler_config = config.scheduler # Determine warmup steps if scheduler_config.warmup_steps > 0: warmup_steps = scheduler_config.warmup_steps else: warmup_steps = int(num_training_steps * scheduler_config.warmup_ratio) # Determine steady phase steps steady_steps = int(num_training_steps * scheduler_config.steady_ratio) # Remaining steps for cosine decay decay_steps = num_training_steps - warmup_steps - steady_steps min_lr_ratio = scheduler_config.min_lr_ratio num_cycles = scheduler_config.num_cycles print(f"✓ CosineWarmup scheduler: warmup={warmup_steps}, steady={steady_steps}, decay={decay_steps} (total={num_training_steps})") print(f" min_lr_ratio={min_lr_ratio}, num_cycles={num_cycles}") def lr_lambda(step): """Three-phase LR schedule: warmup → steady → cosine decay.""" if step < warmup_steps: # Phase 1: Linear warmup from 0 to 1.0 return float(step) / float(max(1, warmup_steps)) elif step < warmup_steps + steady_steps: # Phase 2: Steady at peak LR (1.0) return 1.0 else: # Phase 3: Cosine decay from 1.0 to min_lr_ratio decay_step = step - warmup_steps - steady_steps progress = float(decay_step) / float(max(1, decay_steps)) # Cosine annealing: 0.5 * (1 + cos(π * progress)) cosine_decay = 0.5 * (1.0 + math.cos(math.pi * progress)) # Scale to reach min_lr_ratio at the end return cosine_decay * (1.0 - min_lr_ratio) + min_lr_ratio return LambdaLR(optimizer, lr_lambda, last_epoch=scheduler_config.last_epoch)