File size: 5,908 Bytes
c1a46f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d572bbd
c1a46f7
 
 
d572bbd
 
 
c1a46f7
d572bbd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c1a46f7
 
 
 
 
 
 
 
 
 
d572bbd
 
 
 
c1a46f7
d572bbd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c1a46f7
 
 
 
 
 
 
 
d572bbd
 
c1a46f7
 
 
d572bbd
 
 
 
 
 
 
 
 
 
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
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
"""
优化器与学习率调度器 — Person C 负责实现

功能要求:
1. build_optimizer: 根据配置创建优化器
2. build_scheduler: 根据配置创建学习率调度器
3. InverseSqrtScheduler: 自定义 inverse square root 调度器

技术要点:
- AdamW 是 Transformer 训练的标准优化器
- Cosine with warmup 是目前最流行的调度策略
- Inverse sqrt 是经典 Transformer 论文使用的调度策略
"""

from __future__ import annotations

import math
from typing import Optional

import torch
from torch.optim import Adam, AdamW
from torch.optim.lr_scheduler import LambdaLR


def build_optimizer(model: torch.nn.Module, config: dict) -> torch.optim.Optimizer:
    """
    根据配置创建优化器。

    支持参数分组:
    - LayerNorm 和 bias 参数不施加 weight_decay
    - embedding 层可使用较小的学习率
    """
    opt_config = config.get("optimizer", {})
    opt_type = opt_config.get("type", "adamw")
    lr = float(opt_config.get("lr", 3e-4))
    weight_decay = float(opt_config.get("weight_decay", 0.01))
    betas = tuple(opt_config.get("betas", [0.9, 0.98]))
    eps = float(opt_config.get("eps", 1e-8))

    no_decay = ["bias", "LayerNorm.weight", "layer_norm.weight"]
    optimizer_grouped_parameters = [
        {
            "params": [
                p for n, p in model.named_parameters()
                if p.requires_grad and not any(nd in n for nd in no_decay)
            ],
            "weight_decay": weight_decay,
        },
        {
            "params": [
                p for n, p in model.named_parameters()
                if p.requires_grad and any(nd in n for nd in no_decay)
            ],
            "weight_decay": 0.0,
        },
    ]

    if opt_type == "adam":
        return Adam(optimizer_grouped_parameters, lr=lr, betas=betas, eps=eps)
    elif opt_type == "adamw":
        return AdamW(optimizer_grouped_parameters, lr=lr, betas=betas, eps=eps)
    elif opt_type == "adafactor":
        try:
            from transformers.optimization import Adafactor
            return Adafactor(
                optimizer_grouped_parameters,
                lr=lr,
                scale_parameter=False,
                relative_step=False,
            )
        except ImportError:
            raise ImportError("Adafactor requires transformers library")
    else:
        raise ValueError(f"Unsupported optimizer type: {opt_type}")


def build_scheduler(
    optimizer: torch.optim.Optimizer,
    config: dict,
    num_training_steps: Optional[int] = None,
) -> torch.optim.lr_scheduler._LRScheduler:
    """
    根据配置创建学习率调度器。

    支持:
    - cosine_with_warmup: Cosine 衰减 + 线性 warmup
    - inverse_sqrt: 经典 Transformer 调度策略
    - linear: 线性衰减 + warmup
    """
    sched_config = config.get("scheduler", {})
    sched_type = sched_config.get("type", "cosine_with_warmup")
    warmup_steps = int(sched_config.get("warmup_steps", 4000))
    min_lr = float(sched_config.get("min_lr", 1e-6))

    if sched_type == "cosine_with_warmup":
        try:
            from transformers import get_cosine_schedule_with_warmup
            return get_cosine_schedule_with_warmup(
                optimizer,
                num_warmup_steps=warmup_steps,
                num_training_steps=num_training_steps or 100000,
            )
        except ImportError:
            return _cosine_with_warmup(optimizer, warmup_steps, num_training_steps or 100000, min_lr)

    elif sched_type == "inverse_sqrt":
        return InverseSqrtScheduler(optimizer, warmup_steps=warmup_steps)

    elif sched_type == "linear":
        try:
            from transformers import get_linear_schedule_with_warmup
            return get_linear_schedule_with_warmup(
                optimizer,
                num_warmup_steps=warmup_steps,
                num_training_steps=num_training_steps or 100000,
            )
        except ImportError:
            return _linear_with_warmup(optimizer, warmup_steps, num_training_steps or 100000, min_lr)

    else:
        raise ValueError(f"Unsupported scheduler type: {sched_type}")


def _cosine_with_warmup(optimizer, warmup_steps, total_steps, min_lr=0.0):
    def lr_lambda(current_step):
        if current_step < warmup_steps:
            return float(current_step) / float(max(1, warmup_steps))
        progress = float(current_step - warmup_steps) / float(max(1, total_steps - warmup_steps))
        cosine_decay = 0.5 * (1.0 + math.cos(math.pi * progress))
        return max(min_lr / _get_base_lr(optimizer), cosine_decay)
    return LambdaLR(optimizer, lr_lambda)


def _linear_with_warmup(optimizer, warmup_steps, total_steps, min_lr=0.0):
    def lr_lambda(current_step):
        if current_step < warmup_steps:
            return float(current_step) / float(max(1, warmup_steps))
        progress = float(current_step - warmup_steps) / float(max(1, total_steps - warmup_steps))
        return max(min_lr / _get_base_lr(optimizer), 1.0 - progress)
    return LambdaLR(optimizer, lr_lambda)


def _get_base_lr(optimizer):
    for param_group in optimizer.param_groups:
        return param_group.get("lr", param_group.get("initial_lr", 1e-3))
    return 1e-3


class InverseSqrtScheduler(LambdaLR):
    """
    Inverse Square Root 学习率调度器。

    lr = base_lr * min(step^{-0.5}, step * warmup_steps^{-1.5})

    这是原始 Transformer 论文 (Vaswani et al., 2017) 使用的调度策略。
    warmup 阶段线性增长,warmup 后按 step^{-0.5} 衰减。
    """

    def __init__(self, optimizer, warmup_steps: int = 4000):
        self.warmup_steps = warmup_steps
        warmup_factor = warmup_steps ** (-1.5)

        def lr_lambda(step):
            step += 1
            arg1 = step ** (-0.5)
            arg2 = step * warmup_factor
            return min(arg1, arg2)

        super().__init__(optimizer, lr_lambda)