FST_code / src /lmr /checkpointing /checkpointing copy.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
12.5 kB
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
# State fields
self.epoch = 0
self.step = 0
self.train_loss = float("inf")
self.val_loss = float("inf")
self.tokens_trained = 0
# DDP flags
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()
# ---------- ddp ----------
def _barrier(self):
if self.use_ddp:
torch.distributed.barrier()
# ---------- logging + util ----------
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 loss
# 文件夹格式通常是: best_epoch_010_step_5000_val=0.4500
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)
# ---------- state saving ----------
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)
# 1. 保存模型权重 (使用 HF 原生分片逻辑)
model = unwrap_model(self.model)
if hasattr(model, "save_pretrained"):
# 这是生成 model.safetensors.index.json 和 model-000xx.safetensors 的关键
# max_shard_size 可以控制切分大小,默认通常是 5GB 或 10GB
model.save_pretrained(save_path, safe_serialization=False, max_shard_size="5GB")
# 同时也保存 tokenizer (如果有的话,建议在 Trainer 里传入 tokenizer 并挂载到 self.model 上,或者手动调 tokenizer.save_pretrained)
# if hasattr(self, 'tokenizer') and self.tokenizer:
# self.tokenizer.save_pretrained(save_path)
else:
# 如果不是 HF 模型,回退到存单文件
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"))
# 2. 保存自定义 Metadata
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)
# 3. 保存 Optimizer/Scheduler (放在同一个文件夹里)
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):
# 移除旧的 best 文件夹
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) # Best 通常只存权重,不存 optimizer 以节省空间
self._update_log("best", dirname)
def _save_recent(self):
# 移除旧的 recent 文件夹
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) # Recent 必须存 optimizer 用于恢复
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)
# ---------- checkpoint saving ----------
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)
# 保存 recent (用于断点续训)
self._save_recent()
# 保存 epoch 存档 (可选)
if step is None:
self.step = 0
self._save_epoch()
# 只有当你想每隔多少步存一个永久档时才打开这个
# else:
# self._save_step()
# 保存 best (用于推理)
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()
# ---------- loading (适配文件夹结构) ----------
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}")
# 这里的加载逻辑其实在 Trainer.load_only_model_weights 里有更详细的实现
# 这里主要是为了让 Checkpointing 类也能独立运作
model = unwrap_model(self.model)
# 优先尝试 HF 原生加载 (支持分片)
if hasattr(model, "from_pretrained"):
# 注意:from_pretrained 是类方法,但我们已经有实例了。
# 对于已有实例,通常没有直接的 "load_pretrained" 方法来处理分片。
# 所以我们还是依赖 Trainer 里的那个能够处理 index.json 的 load_only_model_weights 函数。
# 这里我们只负责加载 metadata。
pass
# 读取 Metadata
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 # 返回路径给 Trainer 用
def load_training_states(self, checkpoint_type="recent"):
ckpt_dir = self._get_checkpoint_path(checkpoint_type)
if not ckpt_dir: return
# Optimizer 在文件夹里面
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()