File size: 2,488 Bytes
fd448dd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""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()