| """ | |
| 训练器模块 — Person C 负责实现 | |
| 这是训练的核心控制器,负责: | |
| 1. 训练循环 (training loop) | |
| 2. 验证循环 (validation loop) | |
| 3. 检查点保存/加载 | |
| 4. 日志记录 | |
| 5. 分布式训练支持 | |
| 6. 混合精度训练 | |
| 7. 梯度累积 | |
| 8. 早停 | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import os | |
| import time | |
| from pathlib import Path | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| from torch.utils.data import DataLoader | |
| from tqdm import tqdm | |
| logger = logging.getLogger(__name__) | |
| class Trainer: | |
| """ | |
| 翻译模型训练器。 | |
| TODO [Person C]: 实现以下所有方法。 | |
| 使用方法: | |
| trainer = Trainer(model, train_loader, val_loader, config) | |
| trainer.train() | |
| """ | |
| def __init__( | |
| self, | |
| model: nn.Module, | |
| train_loader: DataLoader, | |
| val_loader: DataLoader, | |
| config: dict, | |
| optimizer=None, | |
| scheduler=None, | |
| criterion=None, | |
| evaluator=None, | |
| ): | |
| """ | |
| TODO [Person C]: 初始化训练器: | |
| 1. 保存模型、数据加载器、配置 | |
| 2. 如果 optimizer 为 None,调用 build_optimizer 创建 | |
| 3. 如果 scheduler 为 None,调用 build_scheduler 创建 | |
| 4. 如果 criterion 为 None,创建 LabelSmoothedCrossEntropyLoss | |
| 5. 设置混合精度: torch.amp.GradScaler (如果配置了 fp16/bf16) | |
| 6. 设置分布式训练: 根据配置初始化 DDP/FSDP/DeepSpeed | |
| 7. 初始化日志记录器 (TensorBoard / WandB) | |
| 8. 初始化训练状态 (epoch, step, best_metric) | |
| """ | |
| self.model = model | |
| self.train_loader = train_loader | |
| self.val_loader = val_loader | |
| self.config = config | |
| raise NotImplementedError("TODO: Person C 实现 Trainer.__init__") | |
| def train(self): | |
| """ | |
| 主训练循环。 | |
| TODO [Person C]: 实现以下逻辑: | |
| 1. for epoch in range(start_epoch, num_epochs): | |
| 2. train_loss = self._train_one_epoch(epoch) | |
| 3. val_metrics = self._validate(epoch) | |
| 4. self._log_metrics(epoch, train_loss, val_metrics) | |
| 5. self._save_checkpoint(epoch, val_metrics) | |
| 6. if self._should_early_stop(val_metrics): break | |
| 7. self._log_final_results() | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 train") | |
| def _train_one_epoch(self, epoch: int) -> float: | |
| """ | |
| 训练一个 epoch。 | |
| TODO [Person C]: 实现以下逻辑: | |
| 1. model.train() | |
| 2. 遍历 train_loader: | |
| a. 将 batch 移到设备 | |
| b. 混合精度上下文: with torch.amp.autocast('cuda'): | |
| c. logits = model(src_ids, tgt_input_ids, masks...) | |
| d. loss = criterion(logits, labels) | |
| e. loss = loss / gradient_accumulation_steps | |
| f. scaler.scale(loss).backward() | |
| g. 每 gradient_accumulation_steps 步: | |
| - scaler.unscale_(optimizer) | |
| - torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) | |
| - scaler.step(optimizer) | |
| - scaler.update() | |
| - scheduler.step() | |
| - optimizer.zero_grad() | |
| h. 记录 loss, lr 等指标 | |
| 3. 返回平均训练 loss | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 _train_one_epoch") | |
| def _validate(self, epoch: int) -> dict: | |
| """ | |
| 验证。 | |
| TODO [Person C]: 实现以下逻辑: | |
| 1. model.eval() | |
| 2. with torch.no_grad(): | |
| 3. 遍历 val_loader,计算 loss | |
| 4. 每 N 步调用 evaluator 计算 BLEU 等指标 | |
| 5. 返回 {"val_loss": ..., "bleu": ..., "comet": ...} | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 _validate") | |
| def _save_checkpoint(self, epoch: int, metrics: dict): | |
| """ | |
| 保存检查点。 | |
| TODO [Person C]: 实现以下逻辑: | |
| 1. 构建 checkpoint dict: | |
| { | |
| "epoch": epoch, | |
| "step": self.global_step, | |
| "model_state_dict": model.state_dict(), | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "scheduler_state_dict": scheduler.state_dict(), | |
| "metrics": metrics, | |
| "config": config, | |
| } | |
| 2. 保存到 checkpoint_dir/checkpoint_epoch_{epoch}.pt | |
| 3. 如果是最佳模型,额外保存为 best_model.pt | |
| 4. 清理旧的检查点 (保留最近 max_checkpoints 个) | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 _save_checkpoint") | |
| def _load_checkpoint(self, checkpoint_path: str): | |
| """ | |
| 加载检查点继续训练。 | |
| TODO [Person C]: | |
| 1. torch.load(checkpoint_path) | |
| 2. 恢复 model, optimizer, scheduler 状态 | |
| 3. 恢复 epoch, step 计数器 | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 _load_checkpoint") | |
| def _should_early_stop(self, metrics: dict) -> bool: | |
| """ | |
| 判断是否应该早停。 | |
| TODO [Person C]: | |
| 1. 比较当前指标与最佳指标 | |
| 2. 如果连续 patience 个 epoch 没有改善,返回 True | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 _should_early_stop") | |
| def _log_metrics(self, epoch: int, train_loss: float, val_metrics: dict): | |
| """ | |
| 记录训练指标到 TensorBoard / WandB。 | |
| TODO [Person C]: | |
| 1. 使用 self.logger 记录 train_loss, val_loss, bleu, lr 等 | |
| 2. 打印到控制台 (使用 rich 库的 table 格式) | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 _log_metrics") | |
| def _setup_distributed(self): | |
| """ | |
| 设置分布式训练。 | |
| TODO [Person C]: 根据 config["training"]["distributed"]["strategy"]: | |
| 1. "ddp": 使用 torch.nn.parallel.DistributedDataParallel | |
| 2. "fsdp": 使用 torch.distributed.fsdp.FullyShardedDataParallel | |
| 3. "deepspeed": 使用 deepspeed.initialize() | |
| 也可以使用 HuggingFace Accelerate 统一处理。 | |
| """ | |
| raise NotImplementedError("TODO: Person C 实现 _setup_distributed") | |