lijn14
完成C部分内容
d572bbd
Raw
History Blame
5.91 kB
"""
优化器与学习率调度器 — 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)