| """ |
| 优化器与学习率调度器 — 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) |
|
|