import math import traceback import random import time import json from pathlib import Path from contextlib import nullcontext from tqdm import tqdm import torch import torch.nn.functional as F from torch.amp import GradScaler, autocast from torch.utils.data import DataLoader from transformers import get_cosine_schedule_with_warmup from safetensors.torch import load_file from lmr.checkpointing import Checkpointing from lmr.ddp import setup_ddp, cleanup_ddp, initialize_model_ddp, unwrap_model, initialize_samplers_ddp from lmr.utils.logger import Logger class Trainer: def __init__(self, config, model, tokenizer, splits, checkpointing, samplers=None): """ 适配 main.py 的参数: config: 这里的 config 对应 main.py 里的 config.training splits: 数据集字典 """ self.training_config = config self.model = model self.tokenizer = tokenizer self.splits = splits self.checkpointing = checkpointing self.samplers = samplers # 默认设备设置(会在 _setup_training 中根据 DDP 更新) self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 识别 Token self.pad_token_id = getattr(self.tokenizer, 'pad_token_id', 0) if self.pad_token_id is None: self.pad_token_id = 0 self.eos_token_id = getattr(self.tokenizer, 'eos_token_id', None) # 精度设置 - 增加默认值以防 config 缺失 precision = getattr(self.training_config, "precision", "bfloat16") self.autocast_dtype = getattr(torch, precision) if hasattr(torch, precision) else torch.bfloat16 self.use_ddp = False self.rank = 0 self.world_size = 1 self.debug_prompts = [ "Question: What is 15 + 32?\nAnswer:", "Question: There are 5 birds on a tree. 2 fly away. How many are left?\nSolution:", ] # ========================================================================= # 权重加载逻辑 (保留你提供的鲁棒版本) # ========================================================================= def _is_lora_param(self, param_name): return any(indicator in param_name for indicator in ['lora_A', 'lora_B', 'lora_dropout']) def load_only_model_weights(self, checkpoint_path, map_location="cpu", strict=False, verbose=True): path_obj = Path(checkpoint_path) if not path_obj.exists(): return model_state = {} if path_obj.is_dir(): index_file = path_obj / "model.safetensors.index.json" if index_file.exists(): with open(index_file, 'r') as f: index_data = json.load(f) weight_map = index_data.get("weight_map", {}) for shard_name in set(weight_map.values()): model_state.update(load_file(str(path_obj / shard_name), device=str(map_location))) else: possible = list(path_obj.glob("*.safetensors")) + list(path_obj.glob("*.pt")) if possible: path_obj = possible[0] if not model_state and path_obj.is_file(): if path_obj.suffix == ".safetensors": model_state = load_file(str(path_obj), device=str(map_location)) else: ckpt = torch.load(path_obj, map_location=map_location) model_state = ckpt.get("model", ckpt.get("state_dict", ckpt)) ckpt_keys_map = {k.replace("module.", "").replace("_orig_mod.", "").replace("model.", ""): k for k in model_state.keys()} load_target = unwrap_model(self.model) target_state = load_target.state_dict() filtered_state = {} for k_target, v_target in target_state.items(): k_clean = k_target.replace("module.", "").replace("_orig_mod.", "").replace("model.", "") if k_clean in ckpt_keys_map: v_ckpt = model_state[ckpt_keys_map[k_clean]] if v_ckpt.shape == v_target.shape: filtered_state[k_target] = v_ckpt load_target.load_state_dict(filtered_state, strict=strict) if verbose and self.rank == 0: print(f"✅ Loaded {len(filtered_state)} parameters from {checkpoint_path}") # ========================================================================= # 训练设置 & 初始化 # ========================================================================= def _setup_training(self): # 1. Dataloaders self.train_dataloader = self._get_dataloader("train") # 适配不同数据集可能的 split 命名 val_key = "val" if "val" in self.splits else "validation" self.validation_dataloader = self._get_dataloader(val_key) if val_key in self.splits else None # 2. 梯度累积 if getattr(self.training_config, "use_grad_accum", False): if self.training_config.grad_accum_steps == "auto": tps = self.training_config.batch_size * 1024 * self.world_size self.grad_accum_steps = max(1, getattr(self.training_config, "tokens_per_step", 32768) // tps) else: self.grad_accum_steps = self.training_config.grad_accum_steps else: self.grad_accum_steps = 1 self.steps_per_epoch = len(self.train_dataloader) // self.grad_accum_steps self.tokens_per_batch = self.training_config.batch_size * 1024 self.tokens_per_step = self.grad_accum_steps * self.tokens_per_batch * self.world_size self.tokens_per_epoch = self.tokens_per_step * self.steps_per_epoch # 3. 基础恢复 try: self.checkpointing.load_model_states("recent") except: pass # 4. 设备 & 编译 self.device = torch.device(f"cuda:{self.rank}") # if getattr(self.training_config, "compile", False): # self.model = torch.compile(self.model, mode=getattr(self.training_config, "compile_mode", "default")) self.model.to(self.device) # 5. DDP 包装 if self.use_ddp: self.model = initialize_model_ddp(self.model, self.rank) # 6. 优化器 & 调度器 self._initialize_optimizer() self._initialize_scheduler() self._initialize_scaler() # 7. 训练状态恢复 try: self.checkpointing.load_training_states("recent") except: pass def _initialize_optimizer(self): # 自动判断 LoRA is_lora = any(getattr(unwrap_model(self.model).config, f"use_lora_{x}", False) for x in ["phi_attention", "icl_attention"]) if is_lora: params = [p for n, p in self.model.named_parameters() if self._is_lora_param(n)] for n, p in self.model.named_parameters(): p.requires_grad = self._is_lora_param(n) else: params = self.model.parameters() self.optimizer = torch.optim.AdamW( params, lr=self.training_config.lr, betas=getattr(self.training_config, "betas", (0.9, 0.95)), weight_decay=self.training_config.weight_decay ) self.checkpointing.optimizer = self.optimizer def _initialize_scheduler(self): warmup = getattr(self.training_config, "warmup_steps", 100) total = self.steps_per_epoch * self.training_config.max_epochs self.scheduler = get_cosine_schedule_with_warmup(self.optimizer, warmup, total) self.checkpointing.scheduler = self.scheduler def _initialize_scaler(self): self.scaler = GradScaler("cuda") if self.autocast_dtype == torch.float16 else None self.checkpointing.scaler = self.scaler def _get_dataloader(self, split_name): return DataLoader( self.splits[split_name], batch_size=self.training_config.batch_size, num_workers=getattr(self.training_config, "num_workers", 4), sampler=self.samplers[split_name] if self.samplers else None, shuffle=(split_name == "train" and self.samplers is None), pin_memory=True, drop_last=True ) # ========================================================================= # 核心循环 # ========================================================================= def _calculate_training_tokens(self, epoch, step): return epoch * self.tokens_per_epoch + step * self.tokens_per_step def _step_loss(self, batch): batch = batch.to(self.device, non_blocking=True) input_tokens = batch[:, :] vocab_limit = self.model.config.vocab_size if (input_tokens >= vocab_limit).any(): print(f"警告:发现越界 Token! 最大 ID: {input_tokens.max()}") input_tokens[input_tokens >= vocab_limit] = 0 # 临时替换为 0 防止崩溃 target_tokens = batch[:, 1:].clone() if self.pad_token_id is not None: target_tokens[target_tokens == self.pad_token_id] = 0 with autocast(device_type="cuda", dtype=self.autocast_dtype): logits = self.model(input_tokens) loss = unwrap_model(self.model).calculate_loss(logits[:, :-1], target_tokens) return loss def _train(self): self._setup_training() start_epoch = self.checkpointing.epoch start_step = self.checkpointing.step tokens_trained = self.checkpointing.tokens_trained if self.rank == 0: Logger.log(f"🚀 Starting training from Epoch {start_epoch}, Step {start_step}") for epoch in range(start_epoch, self.training_config.max_epochs): if self.use_ddp: self.train_dataloader.sampler.set_epoch(epoch) pbar = tqdm(total=self.steps_per_epoch, desc=f"Epoch {epoch}", disable=(self.rank != 0)) self.model.train() accum_loss = 0.0 for micro_step, batch in enumerate(self.train_dataloader): step = micro_step // self.grad_accum_steps is_update_step = ((micro_step + 1) % self.grad_accum_steps == 0) # 跳过已训练步骤 (Resume) if epoch == start_epoch and step < start_step: if is_update_step: pbar.update(1) continue sync_ctx = self.model.no_sync() if (self.use_ddp and not is_update_step) else nullcontext() with sync_ctx: loss = self._step_loss(batch) / self.grad_accum_steps if self.scaler: self.scaler.scale(loss).backward() else: loss.backward() accum_loss += loss.item() if is_update_step: if self.scaler: self.scaler.unscale_(self.optimizer) torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0) self.scaler.step(self.optimizer) self.scaler.update() else: torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0) self.optimizer.step() self.optimizer.zero_grad(set_to_none=True) self.scheduler.step() if pbar: pbar.set_postfix(loss=f"{accum_loss:.4f}") pbar.update(1) accum_loss = 0.0 # Epoch 结束逻辑 if self.rank == 0: val_loss = self._validate() if self.validation_dataloader else 0.0 tokens_trained = self._calculate_training_tokens(epoch + 1, 0) self.checkpointing.save_checkpoint(epoch + 1, 0, 0, val_loss, tokens_trained) Logger.log(f"Epoch {epoch} Done. Val Loss: {val_loss:.4f}") @torch.no_grad() def _validate(self): self.model.eval() total_loss = 0 for batch in tqdm(self.validation_dataloader, desc="Validating", leave=False, disable=(self.rank != 0)): total_loss += self._step_loss(batch).item() return total_loss / len(self.validation_dataloader) def train(self): """外部唯一调用入口""" if getattr(self.training_config, "use_ddp", False): self.use_ddp = True self.rank, self.world_size = setup_ddp() try: self.samplers = initialize_samplers_ddp(self.splits, self.rank, self.world_size) self._train() finally: cleanup_ddp() else: self._train()