| """ |
| 训练器模块 — 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") |
|
|