| import time |
| import torch |
| import re |
| import json |
| import shutil |
| from pathlib import Path |
| from glob import glob |
| from lmr.utils.logger import Logger |
| from lmr.utils.parsing import int_to_formatted_string |
| from lmr.ddp import unwrap_model |
|
|
| class Checkpointing: |
| |
| def __init__(self, model, checkpoint_dir, optimizer=None, scheduler=None, scaler=None, map_device="cpu"): |
| |
| self.model = model |
| self.optimizer = optimizer |
| self.scheduler = scheduler |
| self.scaler = scaler |
| self.map_device = map_device |
| |
| |
| self.epoch = 0 |
| self.step = 0 |
| self.train_loss = float("inf") |
| self.val_loss = float("inf") |
| self.tokens_trained = 0 |
|
|
| |
| self.use_ddp = torch.distributed.is_available() and torch.distributed.is_initialized() |
| self.is_main_process = not self.use_ddp or torch.distributed.get_rank() == 0 |
|
|
| self.checkpoint_dir = Path(checkpoint_dir) |
| |
| if not self.is_main_process: |
| return |
|
|
| self.checkpoint_dir.mkdir(parents=True, exist_ok=True) |
| self.log_path = self.checkpoint_dir / "_checkpoint_log.tsv" |
| self._create_log() |
|
|
| self.best_val_loss = self._get_best_val_loss() |
|
|
| |
|
|
| def _barrier(self): |
| if self.use_ddp: |
| torch.distributed.barrier() |
|
|
| |
| |
| def _remove_old(self, pattern): |
| """ |
| 删除旧的 Checkpoint 文件夹。 |
| pattern 例如: "recent_epoch*_step*" |
| """ |
| if not self.is_main_process: |
| return |
| |
| |
| candidates = sorted(glob(str(self.checkpoint_dir / pattern))) |
| for path_str in candidates: |
| path = Path(path_str) |
| try: |
| if path.is_dir(): |
| shutil.rmtree(path) |
| print(f"🗑️ Removed old checkpoint dir: {path.name}") |
| else: |
| path.unlink() |
| except Exception as e: |
| print(f"⚠️ Failed to remove {path}: {e}") |
|
|
| def _get_best_val_loss(self): |
| |
| |
| best_dirs = glob(str(self.checkpoint_dir / "best_epoch*_val=*")) |
| best_val = float("inf") |
| for dir_path in best_dirs: |
| match = re.search(r"val=([0-9.]+)", dir_path) |
| if match: |
| try: |
| best_val = min(best_val, float(match.group(1))) |
| except ValueError: |
| pass |
| return best_val |
|
|
| def _checkpoint_dirname(self, epoch, step=None, val_loss=None, tokens_trained=None, prefix=None): |
| """生成文件夹名称""" |
| name = f"epoch_{epoch:03d}" |
| |
| if prefix is not None: |
| name = f"{prefix}_{name}" |
| if step is not None: |
| name += f"_step_{step:09d}" |
| if tokens_trained is not None: |
| name += f"_tokens_{int_to_formatted_string(tokens_trained)}" |
| if val_loss is not None: |
| name += f"_val={val_loss:.4f}" |
| |
| return name |
|
|
| def _create_log(self): |
| if self.log_path.exists(): |
| return |
| header = "Time\tCheckpoint_Type\tEpoch\tStep\tTrain_Loss\tVal_Loss\tTokens_Trained\n" |
| with open(self.log_path, "w", encoding="utf-8") as f: |
| f.write(header) |
|
|
| def _update_log(self, kind, dirname): |
| ts = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) |
| epoch_str = f"{self.epoch:03d}" |
| step_str = f"{self.step:09d}" |
| |
| line = ( |
| f"{ts}\t{kind}\t{epoch_str}\t{step_str}\t" |
| f"{'' if self.train_loss is None else self.train_loss}\t" |
| f"{'' if self.val_loss is None else self.val_loss}\t" |
| f"{dirname}\n" |
| ) |
| |
| with open(self.log_path, "a", encoding="utf-8") as f: |
| f.write(line) |
|
|
| |
|
|
| def _save_state(self, dirname, include_training_states=False): |
| """ |
| 创建文件夹并使用 save_pretrained 保存分片模型。 |
| """ |
| save_path = self.checkpoint_dir / dirname |
| save_path.mkdir(parents=True, exist_ok=True) |
| |
| |
| model = unwrap_model(self.model) |
| |
| if hasattr(model, "save_pretrained"): |
| |
| |
| model.save_pretrained(save_path, safe_serialization=False, max_shard_size="5GB") |
| |
| |
| |
| |
| |
| else: |
| |
| Logger.log("⚠️ Model does not have save_pretrained method. Saving single file.") |
| from safetensors.torch import save_file |
| state_dict = {k: v.cpu().contiguous() for k, v in model.state_dict().items()} |
| save_file(state_dict, str(save_path / "model.safetensors")) |
|
|
| |
| metadata = { |
| "epoch": self.epoch, |
| "step": self.step, |
| "train_loss": self.train_loss, |
| "val_loss": self.val_loss, |
| "tokens_trained": self.tokens_trained, |
| } |
| with open(save_path / "trainer_state.json", 'w') as f: |
| json.dump(metadata, f, indent=4) |
|
|
| |
| if include_training_states: |
| train_state = {} |
| if self.optimizer is not None: |
| train_state["optimizer"] = self.optimizer.state_dict() |
| if self.scheduler is not None and hasattr(self.scheduler, "state_dict"): |
| train_state["scheduler"] = self.scheduler.state_dict() |
| if self.scaler is not None and hasattr(self.scaler, "state_dict"): |
| train_state["scaler"] = self.scaler.state_dict() |
| |
| torch.save(train_state, save_path / "optimizer.pt") |
| |
| Logger.log(f"💾 Checkpoint saved to: {save_path}") |
|
|
| def _update_state(self, epoch, step=0, train_loss=None, val_loss=None, tokens_trained=None): |
| self.epoch = int(epoch) |
| self.step = int(step) if step is not None else 0 |
| if train_loss is not None: |
| self.train_loss = train_loss |
| if val_loss is not None: |
| self.val_loss = val_loss |
| if tokens_trained is not None: |
| self.tokens_trained = tokens_trained |
|
|
| def _save_best(self): |
| |
| self._remove_old("best_epoch*") |
| |
| dirname = self._checkpoint_dirname(self.epoch, self.step, self.val_loss, self.tokens_trained, prefix="best") |
| self._save_state(dirname, include_training_states=False) |
| self._update_log("best", dirname) |
| |
| def _save_recent(self): |
| |
| self._remove_old("recent_epoch*") |
| |
| dirname = self._checkpoint_dirname(self.epoch, self.step, self.val_loss, self.tokens_trained, prefix="recent") |
| self._save_state(dirname, include_training_states=True) |
| self._update_log("recent", dirname) |
|
|
| def _save_epoch(self): |
| dirname = self._checkpoint_dirname(self.epoch, None, self.val_loss, self.tokens_trained) |
| self._save_state(dirname, include_training_states=False) |
| self._update_log("epoch", dirname) |
| |
| def _save_step(self): |
| dirname = self._checkpoint_dirname(self.epoch, self.step, self.val_loss, self.tokens_trained) |
| self._save_state(dirname, include_training_states=False) |
| self._update_log("step", dirname) |
|
|
| |
|
|
| def save_checkpoint(self, epoch, step=None, train_loss=None, val_loss=None, tokens_trained=None): |
| |
| if self.is_main_process: |
| self._update_state(epoch, step=step, train_loss=train_loss, val_loss=val_loss, tokens_trained=tokens_trained) |
| |
| |
| self._save_recent() |
|
|
| |
| if step is None: |
| self.step = 0 |
| self._save_epoch() |
| |
| |
| |
|
|
| |
| if (self.val_loss is not None) and (self.val_loss < self.best_val_loss): |
| self.best_val_loss = self.val_loss |
| self._save_best() |
| |
| self._barrier() |
|
|
| |
|
|
| def _get_checkpoint_path(self, checkpoint_type): |
| """寻找对应的文件夹""" |
| if checkpoint_type == "best": |
| pattern = "best_epoch*" |
| elif checkpoint_type == "recent": |
| pattern = "recent_epoch*" |
| elif checkpoint_type.startswith("epoch_"): |
| pattern = f"{checkpoint_type}*" |
| else: |
| |
| path = self.checkpoint_dir / checkpoint_type |
| if path.exists(): return path |
| return None |
|
|
| candidates = sorted(glob(str(self.checkpoint_dir / pattern))) |
| if not candidates: |
| Logger.log(f"No checkpoint found for {checkpoint_type}") |
| return None |
| |
| return Path(candidates[-1]) |
|
|
| def load_model_states(self, checkpoint_type="recent"): |
| ckpt_dir = self._get_checkpoint_path(checkpoint_type) |
| if not ckpt_dir: return |
|
|
| Logger.log(f"📂 Loading model from dir: {ckpt_dir}") |
| |
| |
| |
| model = unwrap_model(self.model) |
| |
| |
| if hasattr(model, "from_pretrained"): |
| |
| |
| |
| |
| pass |
| |
| |
| meta_path = ckpt_dir / "trainer_state.json" |
| if meta_path.exists(): |
| with open(meta_path, 'r') as f: |
| state = json.load(f) |
| self.epoch = int(state.get("epoch", 0)) |
| self.step = int(state.get("step", 0)) |
| self.train_loss = state.get("train_loss", None) |
| self.val_loss = state.get("val_loss", None) |
| self.tokens_trained = state.get("tokens_trained", 0) |
| |
| self._barrier() |
| return ckpt_dir |
|
|
| def load_training_states(self, checkpoint_type="recent"): |
| ckpt_dir = self._get_checkpoint_path(checkpoint_type) |
| if not ckpt_dir: return |
|
|
| |
| opt_path = ckpt_dir / "optimizer.pt" |
| |
| if not opt_path.exists(): |
| Logger.log(f"⚠️ Optimizer state not found in {ckpt_dir}") |
| return |
|
|
| state = torch.load(opt_path, map_location=self.map_device) |
|
|
| if self.optimizer is not None and "optimizer" in state: |
| self.optimizer.load_state_dict(state["optimizer"]) |
| if self.scheduler is not None and "scheduler" in state: |
| self.scheduler.load_state_dict(state["scheduler"]) |
| if self.scaler is not None and "scaler" in state: |
| self.scaler.load_state_dict(state["scaler"]) |
| |
| Logger.log(f"✅ Loaded training states from {opt_path}") |
| self._barrier() |