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