"""Scheduler registry and factory for instantiating learning rate schedulers.""" from typing import Dict, Callable, Optional import torch.optim as optim from torch.optim.lr_scheduler import LambdaLR from taoTrain.config import TrainingConfig, SchedulerEnum # Global registry for schedulers _SCHEDULER_REGISTRY: Dict[str, Callable] = {} def register_scheduler(name: str): """ Decorator to register a custom scheduler factory function. Args: name: Name of the scheduler (e.g., 'linearWarmup', 'cosineWarmup', 'constant') """ def decorator(fn: Callable) -> Callable: if name in _SCHEDULER_REGISTRY: raise ValueError(f"Scheduler '{name}' is already registered") _SCHEDULER_REGISTRY[name] = fn return fn return decorator def get_registered_schedulers() -> Dict[str, Callable]: """Get all registered scheduler factory functions.""" return _SCHEDULER_REGISTRY.copy() def get_scheduler( optimizer: optim.Optimizer, config: TrainingConfig, num_training_steps: int, ) -> LambdaLR: """ Create a learning rate scheduler instance from config. Args: optimizer: Optimizer to schedule learning rate for config: TrainingConfig with scheduler configuration num_training_steps: Total number of training steps Returns: Learning rate scheduler instance Raises: ValueError: If scheduler type is not registered """ # Handle both enum and string values scheduler_type = config.scheduler.scheduler_type if isinstance(scheduler_type, str): scheduler_name = scheduler_type else: scheduler_name = scheduler_type.value if scheduler_name not in _SCHEDULER_REGISTRY: raise ValueError( f"Unknown scheduler: {scheduler_name}. " f"Available: {list(_SCHEDULER_REGISTRY.keys())}" ) factory_fn = _SCHEDULER_REGISTRY[scheduler_name] return factory_fn(optimizer, config, num_training_steps) def register_builtin_schedulers(): """Register all built-in schedulers.""" # Import here to trigger decorator registration (avoid circular imports) from . import linear_warmup # noqa: F401 from . import cosine_warmup # noqa: F401 from . import constant # noqa: F401 # Auto-register built-in schedulers when module is imported register_builtin_schedulers()