File size: 2,365 Bytes
c1a46f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
优化器与学习率调度器 — 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, config: dict) -> torch.optim.Optimizer:
    """
    根据配置创建优化器。

    TODO [Person C]: 实现以下逻辑:
    1. 从 config 中读取 optimizer type, lr, weight_decay, betas, eps
    2. 根据 type 创建 Adam / AdamW / Adafactor
    3. (可选) 对不同参数组设置不同学习率:
       - embedding 层可以用较小的 lr
       - LayerNorm 的 bias 不加 weight_decay
    """
    raise NotImplementedError("TODO: Person C 实现 build_optimizer")


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

    TODO [Person C]: 实现以下逻辑:
    1. 从 config 中读取 scheduler type, warmup_steps, min_lr
    2. type == "cosine_with_warmup":
       使用 get_cosine_schedule_with_warmup (transformers 库)
    3. type == "inverse_sqrt":
       实现经典的 lr = d_model^(-0.5) * min(step^(-0.5), step * warmup^(-1.5))
    4. type == "linear":
       使用 get_linear_schedule_with_warmup

    参考: Attention Is All You Need, Section 5.3
    """
    raise NotImplementedError("TODO: Person C 实现 build_scheduler")


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

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

    这是原始 Transformer 论文使用的调度策略。

    TODO [Person C]:
    1. 实现 lr_lambda 函数
    2. warmup 阶段线性增长
    3. warmup 后按 step^{-0.5} 衰减
    """

    def __init__(self, optimizer, warmup_steps: int = 4000):
        raise NotImplementedError("TODO: Person C 实现 InverseSqrtScheduler.__init__")