| """ |
| 训练器模块 — Person C 负责实现 |
| |
| 这是训练的核心控制器,负责: |
| 1. 训练循环 (training loop) |
| 2. 验证循环 (validation loop) |
| 3. 检查点保存/加载 |
| 4. 日志记录 |
| 5. 分布式训练支持 |
| 6. 混合精度训练 |
| 7. 梯度累积 |
| 8. 早停 |
| """ |
|
|
| from __future__ import annotations |
|
|
| import json |
| import logging |
| import os |
| import shutil |
| 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 |
|
|
| from easytranslate.training.loss import LabelSmoothedCrossEntropyLoss |
| from easytranslate.training.optimizer import build_optimizer, build_scheduler |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| class Trainer: |
| """ |
| 翻译模型训练器。 |
| |
| 使用方法: |
| 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, |
| ): |
| self.model = model |
| self.train_loader = train_loader |
| self.val_loader = val_loader |
| self.config = config |
|
|
| train_cfg = config.get("training", {}) |
|
|
| self.device = self._resolve_device(train_cfg) |
| self.model = self.model.to(self.device) |
|
|
| self.fp16 = train_cfg.get("fp16", False) |
| self.bf16 = train_cfg.get("bf16", False) |
| self.gradient_accumulation_steps = int(train_cfg.get("gradient_accumulation_steps", 1)) |
| self.max_grad_norm = float(train_cfg.get("max_grad_norm", 1.0)) |
| self.gradient_checkpointing = train_cfg.get("gradient_checkpointing", False) |
|
|
| if self.gradient_checkpointing and hasattr(self.model, "gradient_checkpointing_enable"): |
| self.model.gradient_checkpointing_enable() |
|
|
| self.num_epochs = int(train_cfg.get("epochs", 30)) |
| self.max_steps = int(train_cfg.get("max_steps", -1)) |
|
|
| self.scaler = None |
| self.amp_dtype = None |
| if self.fp16: |
| self.scaler = torch.amp.GradScaler(self.device.type) |
| self.amp_dtype = torch.float16 |
| elif self.bf16: |
| self.amp_dtype = torch.bfloat16 |
|
|
| if optimizer is None: |
| optimizer = build_optimizer(model, train_cfg) |
| self.optimizer = optimizer |
|
|
| total_steps = self.num_epochs * len(train_loader) // self.gradient_accumulation_steps |
| if scheduler is None: |
| scheduler = build_scheduler(optimizer, train_cfg, num_training_steps=total_steps) |
| self.scheduler = scheduler |
|
|
| reg_cfg = train_cfg.get("regularization", {}) |
| if criterion is None: |
| criterion = LabelSmoothedCrossEntropyLoss( |
| smoothing=float(reg_cfg.get("label_smoothing", 0.1)), |
| pad_id=-100, |
| ) |
| self.criterion = criterion |
|
|
| self.evaluator = evaluator |
|
|
| ckpt_cfg = train_cfg.get("checkpoint", {}) |
| self.checkpoint_dir = Path(ckpt_cfg.get("save_dir", "checkpoints/")) |
| self.checkpoint_dir.mkdir(parents=True, exist_ok=True) |
| self.save_every_n_steps = int(ckpt_cfg.get("save_every_n_steps", 5000)) |
| self.save_best = ckpt_cfg.get("save_best", True) |
| self.metric_for_best = ckpt_cfg.get("metric_for_best", "bleu") |
| self.max_checkpoints = int(ckpt_cfg.get("max_checkpoints", 5)) |
|
|
| es_cfg = train_cfg.get("early_stopping", {}) |
| self.early_stopping_enabled = es_cfg.get("enabled", True) |
| self.patience = int(es_cfg.get("patience", 5)) |
| self.min_delta = float(es_cfg.get("min_delta", 0.1)) |
|
|
| log_cfg = config.get("logging", {}) |
| self.log_dir = Path(log_cfg.get("log_dir", "logs/")) |
| self.log_dir.mkdir(parents=True, exist_ok=True) |
| self.log_every_n_steps = int(log_cfg.get("log_every_n_steps", 100)) |
| self.log_backend = log_cfg.get("backend", "tensorboard") |
|
|
| self._init_logger() |
|
|
| self.global_step = 0 |
| self.current_epoch = 0 |
| self.best_metric_value = float("-inf") |
| self.best_epoch = 0 |
| self.epochs_without_improvement = 0 |
| self.train_loss_history: list[float] = [] |
| self.val_metrics_history: list[dict] = [] |
|
|
| self._setup_distributed() |
|
|
| def _resolve_device(self, train_cfg: dict) -> torch.device: |
| device_str = train_cfg.get("device", "auto") |
| if device_str == "auto": |
| if torch.cuda.is_available(): |
| return torch.device("cuda") |
| elif torch.backends.mps.is_available(): |
| return torch.device("mps") |
| else: |
| return torch.device("cpu") |
| return torch.device(device_str) |
|
|
| def _init_logger(self): |
| self.writer = None |
| if self.log_backend in ("tensorboard", "both"): |
| try: |
| from torch.utils.tensorboard import SummaryWriter |
| self.writer = SummaryWriter(log_dir=str(self.log_dir)) |
| except ImportError: |
| logger.warning("TensorBoard not available, skipping") |
|
|
| self.wandb_run = None |
| if self.log_backend in ("wandb", "both"): |
| try: |
| import wandb |
| log_cfg = self.config.get("logging", {}) |
| self.wandb_run = wandb.init( |
| project=log_cfg.get("project_name", "EasyTranslate"), |
| config=self.config, |
| dir=str(self.log_dir), |
| ) |
| except ImportError: |
| logger.warning("WandB not available, skipping") |
|
|
| def _setup_distributed(self): |
| dist_cfg = self.config.get("training", {}).get("distributed", {}) |
| strategy = dist_cfg.get("strategy", "ddp") |
|
|
| if strategy == "ddp" and torch.distributed.is_available() and torch.distributed.is_initialized(): |
| self.model = nn.parallel.DistributedDataParallel(self.model) |
| self.is_distributed = True |
| elif strategy == "deepspeed": |
| try: |
| import deepspeed |
| ds_config = dist_cfg.get("deepspeed_config", "configs/deepspeed_config.json") |
| self.model, self.optimizer, _, _ = deepspeed.initialize( |
| model=self.model, |
| optimizer=self.optimizer, |
| config_params=ds_config, |
| ) |
| self.is_distributed = True |
| except ImportError: |
| logger.warning("DeepSpeed not available, falling back to single GPU") |
| self.is_distributed = False |
| else: |
| self.is_distributed = False |
|
|
| def train(self): |
| logger.info("Starting training for %d epochs", self.num_epochs) |
| logger.info("Device: %s, FP16: %s, BF16: %s", self.device, self.fp16, self.bf16) |
| logger.info("Gradient accumulation steps: %d", self.gradient_accumulation_steps) |
|
|
| for epoch in range(self.current_epoch, self.num_epochs): |
| self.current_epoch = epoch |
| logger.info("=" * 50) |
| logger.info("Epoch %d/%d", epoch + 1, self.num_epochs) |
|
|
| train_loss = self._train_one_epoch(epoch) |
| self.train_loss_history.append(train_loss) |
|
|
| val_metrics = self._validate(epoch) |
| self.val_metrics_history.append(val_metrics) |
|
|
| self._log_metrics(epoch, train_loss, val_metrics) |
| self._save_checkpoint(epoch, val_metrics) |
|
|
| if self._should_early_stop(val_metrics): |
| logger.info("Early stopping triggered at epoch %d", epoch + 1) |
| break |
|
|
| if self.max_steps > 0 and self.global_step >= self.max_steps: |
| logger.info("Reached max steps %d, stopping", self.max_steps) |
| break |
|
|
| self._log_final_results() |
| self._cleanup() |
|
|
| def _train_one_epoch(self, epoch: int) -> float: |
| self.model.train() |
| total_loss = 0.0 |
| num_batches = 0 |
| self.optimizer.zero_grad() |
|
|
| pbar = tqdm(self.train_loader, desc=f"Train Epoch {epoch + 1}", leave=False) |
| for batch_idx, batch in enumerate(pbar): |
| batch = self._move_batch_to_device(batch) |
|
|
| with torch.amp.autocast(self.device.type, enabled=self.amp_dtype is not None, dtype=self.amp_dtype): |
| logits = self.model( |
| batch["src_ids"], |
| batch["tgt_input_ids"], |
| batch.get("src_padding_mask"), |
| batch.get("tgt_padding_mask"), |
| ) |
| loss = self.criterion(logits, batch["labels"]) |
|
|
| loss = loss / self.gradient_accumulation_steps |
|
|
| if self.scaler is not None: |
| self.scaler.scale(loss).backward() |
| else: |
| loss.backward() |
|
|
| total_loss += loss.item() * self.gradient_accumulation_steps |
| num_batches += 1 |
|
|
| if (batch_idx + 1) % self.gradient_accumulation_steps == 0: |
| if self.scaler is not None: |
| self.scaler.unscale_(self.optimizer) |
| nn.utils.clip_grad_norm_(self.model.parameters(), self.max_grad_norm) |
| self.scaler.step(self.optimizer) |
| self.scaler.update() |
| else: |
| nn.utils.clip_grad_norm_(self.model.parameters(), self.max_grad_norm) |
| self.optimizer.step() |
|
|
| self.scheduler.step() |
| self.optimizer.zero_grad() |
| self.global_step += 1 |
|
|
| current_lr = self.scheduler.get_last_lr()[0] |
| pbar.set_postfix({ |
| "loss": f"{loss.item() * self.gradient_accumulation_steps:.4f}", |
| "lr": f"{current_lr:.2e}", |
| "step": self.global_step, |
| }) |
|
|
| if self.global_step % self.log_every_n_steps == 0: |
| self._log_step_metrics(loss.item() * self.gradient_accumulation_steps, current_lr) |
|
|
| if self.save_every_n_steps > 0 and self.global_step % self.save_every_n_steps == 0: |
| self._save_checkpoint(epoch, {"step": self.global_step}, prefix=f"step_{self.global_step}") |
|
|
| avg_loss = total_loss / max(num_batches, 1) |
| logger.info("Epoch %d - Train Loss: %.4f", epoch + 1, avg_loss) |
| return avg_loss |
|
|
| @torch.no_grad() |
| def _validate(self, epoch: int) -> dict: |
| self.model.eval() |
| total_loss = 0.0 |
| num_batches = 0 |
|
|
| pbar = tqdm(self.val_loader, desc=f"Val Epoch {epoch + 1}", leave=False) |
| for batch in pbar: |
| batch = self._move_batch_to_device(batch) |
|
|
| with torch.amp.autocast(self.device.type, enabled=self.amp_dtype is not None, dtype=self.amp_dtype): |
| logits = self.model( |
| batch["src_ids"], |
| batch["tgt_input_ids"], |
| batch.get("src_padding_mask"), |
| batch.get("tgt_padding_mask"), |
| ) |
| loss = self.criterion(logits, batch["labels"]) |
| total_loss += loss.item() |
| num_batches += 1 |
|
|
| val_loss = total_loss / max(num_batches, 1) |
| metrics = {"val_loss": round(val_loss, 4)} |
|
|
| if self.evaluator is not None and self.config.get("evaluation", {}).get("eval_on_epoch_end", True): |
| try: |
| eval_results = self.evaluator.evaluate(self.val_loader) |
| metrics.update(eval_results) |
| except Exception as e: |
| logger.warning("Evaluation failed during validation: %s", e) |
|
|
| logger.info("Epoch %d - Val Loss: %.4f", epoch + 1, val_loss) |
| for k, v in metrics.items(): |
| if k != "val_loss" and not isinstance(v, list): |
| logger.info(" %s: %.4f", k, v) |
|
|
| return metrics |
|
|
| def _save_checkpoint(self, epoch: int, metrics: dict, prefix: str = ""): |
| model_to_save = self.model.module if hasattr(self.model, "module") else self.model |
|
|
| checkpoint = { |
| "epoch": epoch, |
| "step": self.global_step, |
| "model_state_dict": model_to_save.state_dict(), |
| "optimizer_state_dict": self.optimizer.state_dict(), |
| "scheduler_state_dict": self.scheduler.state_dict(), |
| "metrics": metrics, |
| "config": self.config, |
| "train_loss_history": self.train_loss_history, |
| "val_metrics_history": self.val_metrics_history, |
| } |
|
|
| if prefix: |
| ckpt_path = self.checkpoint_dir / f"checkpoint_{prefix}.pt" |
| else: |
| ckpt_path = self.checkpoint_dir / f"checkpoint_epoch_{epoch + 1}.pt" |
|
|
| torch.save(checkpoint, ckpt_path) |
| logger.info("Checkpoint saved: %s", ckpt_path) |
|
|
| current_metric = metrics.get(self.metric_for_best, metrics.get("val_loss", float("inf"))) |
| if self.metric_for_best == "val_loss": |
| current_metric = -current_metric |
|
|
| if self.save_best and current_metric > self.best_metric_value: |
| self.best_metric_value = current_metric |
| self.best_epoch = epoch |
| best_path = self.checkpoint_dir / "best_model.pt" |
| torch.save(checkpoint, best_path) |
| logger.info("New best model saved: %s (metric: %.4f)", best_path, current_metric) |
|
|
| self._cleanup_old_checkpoints() |
|
|
| def _cleanup_old_checkpoints(self): |
| ckpt_files = sorted( |
| self.checkpoint_dir.glob("checkpoint_epoch_*.pt"), |
| key=os.path.getmtime, |
| ) |
| while len(ckpt_files) > self.max_checkpoints: |
| oldest = ckpt_files.pop(0) |
| oldest.unlink() |
| logger.debug("Removed old checkpoint: %s", oldest) |
|
|
| def _load_checkpoint(self, checkpoint_path: str): |
| logger.info("Loading checkpoint from %s", checkpoint_path) |
| checkpoint = torch.load(checkpoint_path, map_location=self.device) |
|
|
| model_to_load = self.model.module if hasattr(self.model, "module") else self.model |
| model_to_load.load_state_dict(checkpoint["model_state_dict"]) |
|
|
| self.optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) |
| self.scheduler.load_state_dict(checkpoint["scheduler_state_dict"]) |
| self.current_epoch = checkpoint["epoch"] + 1 |
| self.global_step = checkpoint["step"] |
| self.best_metric_value = checkpoint.get("metrics", {}).get( |
| self.metric_for_best, float("-inf") |
| ) |
| self.train_loss_history = checkpoint.get("train_loss_history", []) |
| self.val_metrics_history = checkpoint.get("val_metrics_history", []) |
|
|
| logger.info("Resumed from epoch %d, step %d", self.current_epoch, self.global_step) |
|
|
| def _should_early_stop(self, metrics: dict) -> bool: |
| if not self.early_stopping_enabled: |
| return False |
|
|
| current_metric = metrics.get(self.metric_for_best, metrics.get("val_loss", float("inf"))) |
| if self.metric_for_best == "val_loss": |
| current_metric = -current_metric |
|
|
| if current_metric > self.best_metric_value + self.min_delta: |
| self.epochs_without_improvement = 0 |
| return False |
|
|
| self.epochs_without_improvement += 1 |
| logger.info( |
| "No improvement for %d epochs (best: %.4f, current: %.4f)", |
| self.epochs_without_improvement, self.best_metric_value, current_metric, |
| ) |
| return self.epochs_without_improvement >= self.patience |
|
|
| def _log_metrics(self, epoch: int, train_loss: float, val_metrics: dict): |
| if self.writer is not None: |
| self.writer.add_scalar("Loss/train", train_loss, epoch) |
| for k, v in val_metrics.items(): |
| if not isinstance(v, list): |
| self.writer.add_scalar(f"Metrics/{k}", v, epoch) |
| self.writer.add_scalar("LR", self.scheduler.get_last_lr()[0], epoch) |
|
|
| if self.wandb_run is not None: |
| import wandb |
| log_dict = {"epoch": epoch, "train_loss": train_loss} |
| for k, v in val_metrics.items(): |
| if not isinstance(v, list): |
| log_dict[f"val/{k}"] = v |
| log_dict["lr"] = self.scheduler.get_last_lr()[0] |
| wandb.log(log_dict, step=self.global_step) |
|
|
| def _log_step_metrics(self, loss: float, lr: float): |
| if self.writer is not None: |
| self.writer.add_scalar("Loss/train_step", loss, self.global_step) |
| self.writer.add_scalar("LR/step", lr, self.global_step) |
|
|
| if self.wandb_run is not None: |
| import wandb |
| wandb.log({"train/loss_step": loss, "lr": lr}, step=self.global_step) |
|
|
| def _log_final_results(self): |
| logger.info("=" * 50) |
| logger.info("Training completed!") |
| logger.info("Best epoch: %d", self.best_epoch + 1) |
| logger.info("Best %s: %.4f", self.metric_for_best, self.best_metric_value) |
|
|
| summary = { |
| "best_epoch": self.best_epoch, |
| "best_metric": self.best_metric_value, |
| "metric_name": self.metric_for_best, |
| "total_steps": self.global_step, |
| "train_loss_history": self.train_loss_history, |
| "val_metrics_history": self.val_metrics_history, |
| } |
| summary_path = self.checkpoint_dir / "training_summary.json" |
| with open(summary_path, "w", encoding="utf-8") as f: |
| json.dump(summary, f, indent=2, ensure_ascii=False, default=str) |
| logger.info("Training summary saved to %s", summary_path) |
|
|
| def _cleanup(self): |
| if self.writer is not None: |
| self.writer.close() |
| if self.wandb_run is not None: |
| self.wandb_run.finish() |
|
|
| def _move_batch_to_device(self, batch: dict) -> dict: |
| return { |
| k: v.to(self.device, non_blocking=True) if isinstance(v, torch.Tensor) else v |
| for k, v in batch.items() |
| } |
|
|