lijn14
创建工程
c1a46f7
Raw
History Blame
6.21 kB
"""
训练器模块 — 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")