"""Constant learning rate scheduler with optional warmup.""" import torch.optim as optim from torch.optim.lr_scheduler import LambdaLR from taoTrain.config import TrainingConfig from .registry import register_scheduler @register_scheduler("constant") def create_constant( optimizer: optim.Optimizer, config: TrainingConfig, num_training_steps: int, ) -> LambdaLR: """ Create a constant learning rate scheduler with optional linear warmup. Linearly increases learning rate from 0 to peak over warmup steps, then keeps it constant for the rest of training. Args: optimizer: Optimizer instance config: TrainingConfig with scheduler configuration 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) def lr_lambda(step): """Constant learning rate with optional warmup.""" if step < warmup_steps: # Linear warmup return float(step) / float(max(1, warmup_steps)) return 1.0 return LambdaLR(optimizer, lr_lambda, last_epoch=scheduler_config.last_epoch)