File size: 2,128 Bytes
228add1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Learning rate scheduler factory.
"""

import torch.optim as optim
from torch.optim.lr_scheduler import (
    CosineAnnealingLR,
    StepLR,
    ReduceLROnPlateau,
    CosineAnnealingWarmRestarts,
    LinearLR,
    SequentialLR,
)


def create_scheduler(
    optimizer: optim.Optimizer,
    scheduler_name: str = 'cosine',
    num_epochs: int = 30,
    step_size: int = 10,
    gamma: float = 0.1,
    warmup_epochs: int = 0,
    min_lr: float = 1e-6,
) -> optim.lr_scheduler.LRScheduler | None:
    """
    Factory function to create learning rate scheduler.

    Args:
        optimizer: optimizer instance
        scheduler_name: 'cosine', 'step', 'plateau', 'cosine_restart', or 'none'
        num_epochs: total number of training epochs
        step_size: step size for StepLR
        gamma: decay factor for StepLR
        warmup_epochs: number of warmup epochs (linear warmup)
        min_lr: minimum learning rate for cosine schedulers
    """
    name = scheduler_name.lower()

    if name == 'none':
        return None

    # Create the main scheduler
    if name == 'cosine':
        main_scheduler = CosineAnnealingLR(
            optimizer, T_max=num_epochs - warmup_epochs, eta_min=min_lr,
        )
    elif name == 'step':
        main_scheduler = StepLR(optimizer, step_size=step_size, gamma=gamma)
    elif name == 'plateau':
        main_scheduler = ReduceLROnPlateau(
            optimizer, mode='min', factor=gamma, patience=5, min_lr=min_lr,
        )
    elif name == 'cosine_restart':
        main_scheduler = CosineAnnealingWarmRestarts(
            optimizer, T_0=10, T_mult=2, eta_min=min_lr,
        )
    else:
        raise ValueError(f"Unknown scheduler: {name}")

    # Add warmup if requested
    if warmup_epochs > 0 and name != 'plateau':
        warmup_scheduler = LinearLR(
            optimizer, start_factor=0.01, end_factor=1.0, total_iters=warmup_epochs,
        )
        scheduler = SequentialLR(
            optimizer,
            schedulers=[warmup_scheduler, main_scheduler],
            milestones=[warmup_epochs],
        )
        return scheduler

    return main_scheduler