| """优化器:AdamW / 学习率调度"""
|
|
|
|
|
| import math
|
| import torch.optim as optim
|
| from torch.optim.lr_scheduler import LambdaLR
|
|
|
|
|
| def get_optimizer(model, config):
|
| """
|
| 获取优化器(AdamW)
|
|
|
| 参数:
|
| model: 模型
|
| config: 训练配置字典,应包含以下键:
|
| - learning_rate: 学习率(默认: 1e-4)
|
| - weight_decay: 权重衰减(默认: 0.01)
|
| - beta1: Adam beta1(默认: 0.9)
|
| - beta2: Adam beta2(默认: 0.999)
|
| - eps: Adam epsilon(默认: 1e-8)
|
|
|
| 返回:
|
| AdamW 优化器
|
| """
|
|
|
| lr = float(config.get("learning_rate", 1e-4))
|
| weight_decay = float(config.get("weight_decay", 0.01))
|
| beta1 = float(config.get("beta1", 0.9))
|
| beta2 = float(config.get("beta2", 0.999))
|
| eps = float(config.get("eps", 1e-8))
|
|
|
| optimizer = optim.AdamW(
|
| model.parameters(),
|
| lr=lr,
|
| weight_decay=weight_decay,
|
| betas=(beta1, beta2),
|
| eps=eps,
|
| )
|
|
|
| return optimizer
|
|
|
|
|
| def get_lr_scheduler(optimizer, config):
|
| """
|
| 获取学习率调度器(支持预热)
|
|
|
| 参数:
|
| optimizer: 优化器
|
| config: 训练配置字典,应包含以下键:
|
| - lr_scheduler: 调度器类型,可选值:
|
| - "cosine": 余弦退火(带预热)
|
| - "linear": 线性衰减(带预热)
|
| - "constant": 常数学习率
|
| - warmup_steps: 预热步数(默认: 100)
|
| - max_steps: 最大训练步数(默认: 10000)
|
|
|
| 返回:
|
| 学习率调度器(如果为 constant,返回 None)
|
| """
|
| scheduler_type = config.get("lr_scheduler", "cosine")
|
| max_steps = config.get("max_steps", 10000)
|
| warmup_steps = config.get("warmup_steps", 100)
|
|
|
| if scheduler_type == "cosine":
|
|
|
| def lr_lambda(step):
|
| if step < warmup_steps:
|
|
|
| return step / warmup_steps if warmup_steps > 0 else 1.0
|
| else:
|
|
|
|
|
| if max_steps <= warmup_steps:
|
| return 0.0
|
| progress = (step - warmup_steps) / (max_steps - warmup_steps)
|
|
|
| progress = min(progress, 1.0)
|
|
|
| return 0.5 * (1.0 + math.cos(progress * math.pi))
|
|
|
| scheduler = LambdaLR(optimizer, lr_lambda)
|
|
|
| elif scheduler_type == "linear":
|
|
|
| def lr_lambda(step):
|
| if step < warmup_steps:
|
|
|
| return step / warmup_steps if warmup_steps > 0 else 1.0
|
| else:
|
|
|
|
|
| if max_steps <= warmup_steps:
|
| return 0.1
|
| progress = (step - warmup_steps) / (max_steps - warmup_steps)
|
|
|
| progress = min(progress, 1.0)
|
|
|
| return 1.0 - 0.9 * progress
|
|
|
| scheduler = LambdaLR(optimizer, lr_lambda)
|
|
|
| elif scheduler_type == "constant":
|
|
|
| scheduler = None
|
|
|
| else:
|
| raise ValueError(
|
| f"未知的学习率调度器类型: {scheduler_type}。"
|
| f"支持的类型: cosine, linear, constant"
|
| )
|
|
|
| return scheduler
|
|
|
|
|
| if __name__ == "__main__":
|
| import sys
|
| import io
|
|
|
|
|
| if sys.platform == "win32":
|
| sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8")
|
|
|
| print("=" * 60)
|
| print("优化器和学习率调度器测试")
|
| print("=" * 60)
|
|
|
|
|
| import torch.nn as nn
|
|
|
| class SimpleModel(nn.Module):
|
| def __init__(self):
|
| super().__init__()
|
| self.linear = nn.Linear(10, 1)
|
|
|
| def forward(self, x):
|
| return self.linear(x)
|
|
|
| model = SimpleModel()
|
|
|
|
|
| config = {
|
| "learning_rate": 1e-3,
|
| "weight_decay": 0.01,
|
| "beta1": 0.9,
|
| "beta2": 0.999,
|
| "eps": 1e-8,
|
| "lr_scheduler": "cosine",
|
| "warmup_steps": 10,
|
| "max_steps": 100,
|
| }
|
|
|
|
|
| print("\n1. 测试优化器创建")
|
| optimizer = get_optimizer(model, config)
|
| print(f" 优化器类型: {type(optimizer).__name__}")
|
| print(f" 学习率: {optimizer.param_groups[0]['lr']}")
|
| print(f" 权重衰减: {optimizer.param_groups[0]['weight_decay']}")
|
| print(f" Beta1: {optimizer.param_groups[0]['betas'][0]}")
|
| print(f" Beta2: {optimizer.param_groups[0]['betas'][1]}")
|
|
|
|
|
| print("\n2. 测试余弦退火学习率调度器(带预热)")
|
| scheduler = get_lr_scheduler(optimizer, config)
|
| print(f" 调度器类型: {type(scheduler).__name__}")
|
| print(f" 初始学习率: {optimizer.param_groups[0]['lr']:.6f}")
|
|
|
|
|
| print("\n3. 模拟训练步骤,观察学习率变化(cosine)")
|
| lrs = []
|
| for step in range(0, 101, 10):
|
| current_lr = optimizer.param_groups[0]["lr"]
|
| lrs.append(current_lr)
|
| print(f" 步数 {step:3d}: 学习率 = {current_lr:.6f}")
|
|
|
|
|
| dummy_loss = sum(p.sum() for p in model.parameters())
|
| dummy_loss.backward()
|
| optimizer.step()
|
| optimizer.zero_grad()
|
| if scheduler is not None:
|
| scheduler.step()
|
|
|
|
|
| print("\n4. 测试线性学习率调度器(带预热)")
|
| config_linear = config.copy()
|
| config_linear["lr_scheduler"] = "linear"
|
| optimizer_linear = get_optimizer(model, config_linear)
|
| scheduler_linear = get_lr_scheduler(optimizer_linear, config_linear)
|
| print(f" 调度器类型: {type(scheduler_linear).__name__}")
|
|
|
| print("\n5. 模拟训练步骤,观察学习率变化(linear)")
|
| for step in range(0, 101, 10):
|
| current_lr = optimizer_linear.param_groups[0]["lr"]
|
| print(f" 步数 {step:3d}: 学习率 = {current_lr:.6f}")
|
|
|
| dummy_loss = sum(p.sum() for p in model.parameters())
|
| dummy_loss.backward()
|
| optimizer_linear.step()
|
| optimizer_linear.zero_grad()
|
| if scheduler_linear is not None:
|
| scheduler_linear.step()
|
|
|
|
|
| print("\n6. 测试常数学习率")
|
| config_constant = config.copy()
|
| config_constant["lr_scheduler"] = "constant"
|
| optimizer_constant = get_optimizer(model, config_constant)
|
| scheduler_constant = get_lr_scheduler(optimizer_constant, config_constant)
|
| print(f" 调度器类型: {scheduler_constant}")
|
| print(f" 学习率: {optimizer_constant.param_groups[0]['lr']:.6f}")
|
|
|
|
|
| print("\n7. 测试错误类型处理")
|
| try:
|
| config_error = config.copy()
|
| config_error["lr_scheduler"] = "invalid"
|
| scheduler_error = get_lr_scheduler(optimizer, config_error)
|
| except ValueError as e:
|
| print(f" 正确捕获错误: {e}")
|
|
|
| print("\n" + "=" * 60)
|
| print("所有测试完成!")
|
| print("=" * 60)
|
|
|