| 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 |
| |
| |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
| |
| 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) |
|
|
| |
| 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): |
| |
| self.train_dataloader = self._get_dataloader("train") |
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| try: self.checkpointing.load_model_states("recent") |
| except: pass |
|
|
| |
| self.device = torch.device(f"cuda:{self.rank}") |
| |
| |
| self.model.to(self.device) |
|
|
| |
| if self.use_ddp: |
| self.model = initialize_model_ddp(self.model, self.rank) |
|
|
| |
| self._initialize_optimizer() |
| self._initialize_scheduler() |
| self._initialize_scaler() |
|
|
| |
| try: self.checkpointing.load_training_states("recent") |
| except: pass |
|
|
| def _initialize_optimizer(self): |
| |
| 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 |
|
|
| 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) |
|
|
| |
| 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 |
|
|
| |
| 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() |