TaoNet-mini-A2 / src /taoTrain /schedulers /cosine_warmup.py
Lobakkang's picture
Upload folder using huggingface_hub
fd448dd verified
Raw
History Blame
3.05 kB
"""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)