第 5 章 训练循环 · 上:流程总览与完整源码 本章最重,故拆成上下两篇。上篇给出训练 8 阶段全景与完整源码,下篇逐行讲解重难点。一个工业级训练循环涉及配置、设备、优化器、调度器、梯度裁剪、断点续训、日志、checkpoint——全部串起来。 5.1 训练流程全景 训练主函数 的 8 个阶段: 下面给出完整源码。把它与第 2 章的配置、第 3 章的数据、第 4 章的模型对照着读,你会发现整个文件就是把那三章的能力串成一个循环。 5.2 工具函数 5.2.1 可复现性:setseed 固定四个随机源。为什么是四个?
本章最重,故拆成上下两篇。上篇给出训练 8 阶段全景与完整源码,下篇逐行讲解重难点。一个工业级训练循环涉及配置、设备、优化器、调度器、梯度裁剪、断点续训、日志、checkpoint——全部串起来。
训练主函数 train() 的 8 个阶段:
1. 解析参数,覆盖默认配置 2. 设置随机种子 + 选择设备 3. 加载数据 + 构建 DataLoader 4. 构建模型 + 迁到设备 5. 构建优化器 (AdamW) + 调度器 (余弦退火) 6.(可选)从 checkpoint 恢复 7. 训练循环(前向 → 反向 → 裁剪 → 步进 → 日志 → 存盘) 8. 训练结束保存 final 模型
下面给出完整源码。把它与第 2 章的配置、第 3 章的数据、第 4 章的模型对照着读,你会发现整个文件就是把那三章的能力串成一个循环。
def set_seed(seed: int) -> None: """固定 random / numpy / torch 的随机种子,保证实验可复现。""" random.seed(seed) # Python random np.random.seed(seed) # NumPy torch.manual_seed(seed) # CPU torch torch.cuda.manual_seed_all(seed) # 所有 GPU torch
固定四个随机源。为什么是四个?因为 PyTorch 项目的随机性来自:数据增强(Python random)、NumPy 操作(如 shuffle)、模型初始化与 dropout(CPU/GPU torch)。少固定任何一个,结果就有偏差。
同种子 + 同硬件 + 同代码 → 近似可复现,但不完全。原因:cuDNN 的某些卷积/attention 算子默认非确定;多卡 all-reduce 浮点累加顺序不固定。完全复现要额外加:
torch.use_deterministic_algorithms(True) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False
但会牺牲性能,通常只在调试/学术复现时开启。
def select_device(name: str) -> torch.device: """根据配置字符串选择训练设备。""" if name == "auto": return torch.device("cuda" if torch.cuda.is_available() else "cpu") return torch.device(name)
简洁的「auto + 显式」二选一。生产中还会区分多卡(cuda:0、cuda:1),本项目单卡够用。
def get_cosine_schedule_with_warmup( optimizer, num_warmup_steps, num_training_steps, min_lr_ratio=0.1, ) -> LambdaLR: """自定义"线性预热 + 余弦退火"学习率调度器。""" min_lr_ratio = max(0.0, min(1.0, min_lr_ratio)) # 钳制到 [0, 1] def lr_lambda(current_step: int) -> float: # 1) 预热阶段:线性升温 if current_step < num_warmup_steps: return float(current_step) / float(max(1, num_warmup_steps)) # 2) 余弦退火阶段:从 1.0 平滑下降到 min_lr_ratio progress = float(current_step - num_warmup_steps) / float( max(1, num_training_steps - num_warmup_steps) ) cosine = 0.5 * (1.0 + math.cos(math.pi * progress)) return min_lr_ratio + (1.0 - min_lr_ratio) * cosine return LambdaLR(optimizer, lr_lambda)
💡 为什么不直接用框架自带的同名函数?因为框架版本的终点固定是 0,不支持
min_lr_ratio。本项目自己实现一份,让退火终点可配——这是工业训练常用的小改进(终点留一点 lr,避免末期梯度完全消失)。
def save_checkpoint(model, optimizer, scheduler, gpt_config, train_config, step, loss, checkpoint_dir) -> str: """将模型权重、优化器状态与配置保存到磁盘,返回保存路径。""" ckpt_dir = Path(checkpoint_dir) ckpt_dir.mkdir(parents=True, exist_ok=True) path = ckpt_dir / f"gpt_step{step}.pt" torch.save({ "step": step, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "scheduler_state_dict": scheduler.state_dict(), "loss": loss, "gpt_config": gpt_config.__dict__, "train_config": train_config.__dict__, }, path) return str(path)
注意存了优化器状态和调度器状态——这是断点续训能无缝衔接的关键(见下篇 5.7)。还存了 gpt_config 和 train_config,这样推理加载时能直接读出当时的架构,无需手动传配置。
def parse_args() -> argparse.Namespace: """解析命令行参数,未指定的项保持 TrainConfig / GPTConfig 的默认值。""" p = argparse.ArgumentParser(description="GPT 模型训练脚本") # 训练相关 p.add_argument("--batch-size", type=int, default=None) p.add_argument("--learning-rate", type=float, default=None) p.add_argument("--max-iters", type=int, default=None) p.add_argument("--warmup-iters", type=int, default=None) p.add_argument("--save-iter", type=int, default=None) p.add_argument("--log-iter", type=int, default=None) p.add_argument("--num-workers", type=int, default=None) p.add_argument("--device", type=str, default=None) p.add_argument("--seed", type=int, default=None) # 模型相关 p.add_argument("--n-layer", type=int, default=None) p.add_argument("--n-head", type=int, default=None) p.add_argument("--n-embd", type=int, default=None) p.add_argument("--block-size", type=int, default=None) p.add_argument("--dropout", type=float, default=None) # 从 checkpoint 恢复训练 p.add_argument("--resume", type=str, default=None) return p.parse_args()
关键设计:所有参数 default=None。这样能区分「用户没传」(保持 config 默认)和「用户传了 0」(覆盖为 0)。
apply_args_to_configs 用 is not None 判断用户是否传了参数,没传就保持 dataclass 默认值。
# 默认 python train.py # 快速冒烟测试 python train.py --max-iters 50 --n-layer 2 --n-embd 64 --n-head 2 --block-size 32 # 调模型规模(接近 GPT-2 small) python train.py --n-layer 12 --n-embd 768 --n-head 12 --block-size 256 # CPU 训练 python train.py --device cpu # 断点续训 python train.py --resume checkpoints/gpt_step500.pt
完整源码较长,这里先给骨架(带行内注释),下篇再逐段深挖。
def train(): args = parse_args() # 1) 初始化配置 gpt_config = GPTConfig() train_config = TrainConfig() apply_args_to_configs(args, gpt_config, train_config) # 2) 随机种子 + 设备 set_seed(train_config.seed) device = select_device(train_config.device) # 3) 数据 text = get_dataset() dataloader = build_dataloader( text=text, block_size=gpt_config.block_size, batch_size=train_config.batch_size, num_workers=train_config.num_workers, shuffle=True, ) # 4) 模型 model = build_model(gpt_config) model.to(device) # 5) 优化器 + 调度器 optimizer = AdamW(model.parameters(), lr=train_config.learning_rate, betas=train_config.betas, weight_decay=train_config.weight_decay) scheduler = get_cosine_schedule_with_warmup( optimizer, num_warmup_steps=train_config.warmup_iters, num_training_steps=train_config.max_iters, min_lr_ratio=train_config.min_lr_ratio, ) # 6) 可选:从 checkpoint 恢复 start_step = 0 if args.resume: ckpt = torch.load(args.resume, map_location=device) model.load_state_dict(ckpt["model_state_dict"]) optimizer.load_state_dict(ckpt["optimizer_state_dict"]) scheduler.load_state_dict(ckpt["scheduler_state_dict"]) start_step = ckpt["step"] + 1 # 7) 训练循环 model.train() step = start_step data_iter = iter(dataloader) while step < train_config.max_iters: try: x, y = next(data_iter) except StopIteration: # epoch 用完自动重启 data_iter = iter(dataloader) x, y = next(data_iter) x = x.to(device, non_blocking=True) y = y.to(device, non_blocking=True) outputs = model(input_ids=x, labels=y) # 前向,框架自动算 loss loss = outputs.loss optimizer.zero_grad(set_to_none=True) loss.backward() # 反向 torch.nn.utils.clip_grad_norm_(model.parameters(), train_config.grad_clip) optimizer.step() scheduler.step() step += 1 # ...日志 + 定期存 checkpoint... # 8) 最终保存 final 模型(只存权重 + 配置,不存优化器) torch.save({ "step": step, "model_state_dict": model.state_dict(), "gpt_config": gpt_config.__dict__, "train_config": train_config.__dict__, "loss": loss.item(), }, final_path) if __name__ == "__main__": train()
| 字段 | 中途 checkpoint | final |
|---|---|---|
model_state_dict |
✅ | ✅ |
optimizer_state_dict |
✅ | ❌ |
scheduler_state_dict |
✅ | ❌ |
step / gpt_config / loss |
✅ | ✅ |
final 不存优化器/调度器——因为它只用来推理,不再训练,存了也是浪费空间(优化器状态与模型权重等大)。最终模型是推理层和 Web UI 的默认加载目标。
⚠️ Windows 用户:入口必须放在
if __name__ == "__main__":保护块内(本项目已遵守),否则 spawn 多进程会递归启动。详见《第 1 章 环境准备与首次运行》。
set_seed 固定 4 个随机源;select_device 简单的 auto/显式选择。default=None,配合 is not None 实现「传了才覆盖」。骨架读完后,去《第 5 章 训练循环 · 下:逐行讲解》看余弦退火的数学推导、AdamW 的设计、断点续训的细节、以及训练循环里 9 个容易踩坑的细节。