feat(train): crash-resilient resume + EMA loss logging (stage1 & stage2)

Addresses both points from review:

1) Resume (so a killed run converges across multiple restarts, like gen_supervisor):
   - checkpoints now carry optimizer state + EMA value in a 'train_state' field
   - --resume auto-picks the latest step ckpt and continues; --resume_from for explicit
   - stage2 resume rebuilds the LoRA structure via enable_lora (same path as fresh)
     then overlays saved adapter weights via set_peft_model_state_dict. Using
     PeftModel.from_pretrained instead reset requires_grad and shrank the trainable
     set 129->33 tensors, corrupting the optimizer state — verified and fixed.
   - Smoke-verified: stage1 step4->10, stage2 step4->8, both restore optimizer +
     EMA + LoRA with no state mismatch.

2) EMA loss logging (honest answer to 'is the loss swing a real problem?'):
   - the old log printed the raw per-micro-batch loss every N steps; on
     heterogeneous 6-task data that swings 2-3x purely from sampling (anomaly
     JSON is low-loss, forecast numbers high-loss), so it looked 'unstable'
     while the model may well have been converging.
   - EMA(alpha=0.05) now logged alongside raw + as a tensorboard scalar
     (total_loss_ema / sft_loss_ema). Read the EMA, not the raw, to judge
     convergence. lr left unchanged (1e-4) pending EMA evidence.

Adds scripts/train_full_supervisor.sh: stage1->stage2 resume loop, same
crash-convergence guarantee as gen_full_supervisor.
This commit is contained in:
张宗平
2026-07-01 01:06:15 +00:00
parent 1a096ff989
commit 5e3030a20e
5 changed files with 262 additions and 26 deletions
+57
View File
@@ -0,0 +1,57 @@
#!/usr/bin/env bash
# Crash-resilient full-scale training supervisor (stage1 → stage2).
#
# Both tsmm.train.stage1 and stage2 now support --resume: each saved checkpoint
# carries encoder/projector (+LoRA for stage2) + optimizer state + EMA loss, so
# a kill mid-run keeps all progress and the next --resume continues seamlessly.
#
# This supervisor loops "run --resume -> check step" until max_steps is reached.
# Each surviving iteration advances the step counter by >=1 (ckpt_every), so the
# run converges across multiple kills.
#
# Usage: scripts/train_full_supervisor.sh [stage1_max_steps] [stage2_max_steps]
set -uo pipefail # no -e: a killed train sub-call must NOT abort the loop
cd "$(dirname "$0")/.."
export PYTHONPATH=src PYTHONUNBUFFERED=1
S1_MAX="${1:-30000}" # 50万 align, bs=8 ga=4 => 1 epoch ≈ 15625 steps; 2 ep = 31250
S2_MAX="${2:-10000}" # 10万 sft, bs=4 ga=8 => 1 epoch ≈ 3125 steps; 3 ep = 9375
run_until() { # name max ckpt_dir module extra_args...
local name=$1 max=$2 ckpt_dir=$3 module=$4; shift 4
local extra=("$@")
while :; do
local cur=0
# read current step from the latest step ckpt's filename
local latest
latest=$(ls -1 "$ckpt_dir"/${name}_step*.pt 2>/dev/null | sort -V | tail -1)
if [ -n "$latest" ]; then
cur=$(echo "$latest" | grep -oE '[0-9]+' | tail -1)
elif [ -f "$ckpt_dir/${name}_final.pt" ]; then
# final carries step inside the file; assume done if present
cur=$max
fi
if [ "$cur" -ge "$max" ]; then
echo "[sup] $name done: step $cur/$max"; return 0
fi
echo "[sup] $name at step $cur/$max — launching --resume $(date '+%H:%M:%S')"
.venv/bin/python -m "$module" \
--ckpt_dir "$ckpt_dir" --max_steps "$max" --resume "${extra[@]}"
# loop: re-check (train may have been killed or completed)
done
}
echo "[sup] === TRAIN START $(date '+%H:%M:%S') ==="
echo "[sup] stage1 target=$S1_MAX stage2 target=$S2_MAX"
run_until stage1 "$S1_MAX" checkpoints/stage1 tsmm.train.stage1 \
--data data/align.jsonl --bs 8 --grad_accum 4 --ctx 512 --lr 1e-4 \
--lambda_contrast 0.1 --perturb_std 0.1 --log_every 50 --ckpt_every 2000 --tensorboard
run_until stage2 "$S2_MAX" checkpoints/stage2 tsmm.train.stage2 \
--data data/sft.jsonl --stage1_ckpt checkpoints/stage1/stage1_final.pt \
--bs 4 --grad_accum 8 --ctx 1024 --lr 1e-4 --lora_r 16 \
--log_every 50 --ckpt_every 2000 --tensorboard
echo "[sup] === TRAIN ALL DONE $(date '+%H:%M:%S') ==="
ls -la checkpoints/stage1/stage1_final.pt checkpoints/stage2/stage2_final.pt 2>/dev/null
+2
View File
@@ -22,3 +22,5 @@ export PYTHONPATH=src PYTHONUNBUFFERED=1
--log_every 20 \
--ckpt_every 2000 \
--tensorboard
# Resume support: pass --resume to continue from the latest step ckpt.
# Usage: scripts/train_stage1.sh data/align.jsonl 30000 checkpoints/stage1 --resume
+1
View File
@@ -22,3 +22,4 @@ export PYTHONPATH=src PYTHONUNBUFFERED=1
--log_every 20 \
--ckpt_every 2000 \
--tensorboard
# Resume support: pass --resume to continue from the latest step ckpt + lora_adapter.
+94 -17
View File
@@ -14,10 +14,11 @@ Design ref: §4.1/§4.2.
from __future__ import annotations
import argparse
import glob
import json
import os
import random
from typing import Iterator, List
from typing import Iterator, List, Optional
import torch
@@ -32,6 +33,40 @@ LLM_PATH = os.environ.get(
)
class EMA:
"""Exponential moving average of a scalar — for noise-free loss curves.
The per-micro-batch loss on heterogeneous 6-task data swings 2-3× between
steps purely from sampling (anomaly JSON is low-loss, forecast numbers are
high-loss). Logging the raw value every 20 steps is meaningless; the EMA is
what tells you whether the model is actually converging.
"""
def __init__(self, alpha: float = 0.05) -> None:
self.alpha = alpha
self.value: Optional[float] = None
def update(self, x: float) -> float:
if self.value is None:
self.value = float(x)
else:
self.value = (1 - self.alpha) * self.value + self.alpha * float(x)
return self.value
def _latest_step(ckpt_dir: str, prefix: str) -> Optional[str]:
"""Return the highest-numbered step checkpoint path in ``ckpt_dir``, or None."""
steps = sorted(glob.glob(os.path.join(ckpt_dir, f"{prefix}step*.pt")))
return steps[-1] if steps else None
def _step_number(path: str) -> int:
"""``stage1_step1500.pt`` -> 1500."""
base = os.path.basename(path)
num = "".join(ch for ch in base if ch.isdigit())
return int(num) if num else 0
def stream_jsonl(path: str) -> Iterator[dict]:
with open(path, "r", encoding="utf-8") as f:
for line in f:
@@ -89,11 +124,40 @@ def train(args: argparse.Namespace) -> None:
writer = SummaryWriter(os.path.join(args.ckpt_dir, "tb"))
# --- crash-resilient resume ---
# Priority: explicit --resume_from > --resume (auto-pick latest step ckpt).
step = 0
opt.zero_grad()
ema = EMA(alpha=0.05)
resume_path = args.resume_from
if not resume_path and args.resume:
resume_path = _latest_step(args.ckpt_dir, "stage1_")
# treat stage1_final.pt as a step ckpt too if it's all we have
final = os.path.join(args.ckpt_dir, "stage1_final.pt")
if os.path.exists(final):
if resume_path is None or _step_number(final) > _step_number(resume_path):
resume_path = final
if resume_path and os.path.exists(resume_path):
st = torch.load(resume_path, map_location="cpu")
model.encoder.load_state_dict(st["encoder"])
model.projector.load_state_dict(st["projector"])
step = int(st.get("step", 0))
ts = st.get("train_state", {})
if "optimizer" in ts:
try:
opt.load_state_dict(ts["optimizer"])
for g in opt.param_groups:
g["lr"] = args.lr # allow lr override on resume
print(f"[stage1] resumed optimizer from {resume_path}")
except (ValueError, KeyError) as e:
print(f"[stage1] optimizer state mismatch ({e}), fresh AdamW")
if "ema" in ts:
ema.value = ts["ema"]
print(f"[stage1] resumed weights from {resume_path} at step {step}")
accum_loss = 0.0
print(f"[stage1] training: max_steps={args.max_steps} bs={args.bs} "
f"grad_accum={args.grad_accum} ctx={args.ctx} lambda={args.lambda_contrast}")
f"grad_accum={args.grad_accum} ctx={args.ctx} lambda={args.lambda_contrast} "
f"start_step={step}")
while step < args.max_steps:
for batch_samples in batched(stream_jsonl(args.data), args.bs):
if step >= args.max_steps:
@@ -122,36 +186,45 @@ def train(args: argparse.Namespace) -> None:
opt.step()
opt.zero_grad()
ema.update(float(loss))
if writer is not None:
writer.add_scalar("lm_loss", float(lm), step)
writer.add_scalar("contrastive_loss", float(contra), step)
writer.add_scalar("total_loss", float(loss), step)
writer.add_scalar("total_loss_ema", ema.value, step)
if step % args.log_every == 0:
peak = (torch.cuda.max_memory_allocated() / 1e9) if torch.cuda.is_available() else 0.0
print(f" step {step:5d} lm={float(lm):.4f} contra={float(contra):.4f} "
f"total={float(loss):.4f} peak_GB={peak:.2f}")
f"total={float(loss):.4f} ema={ema.value:.4f} peak_GB={peak:.2f}")
if (step + 1) % args.ckpt_every == 0:
ckpt = os.path.join(args.ckpt_dir, f"stage1_step{step+1}.pt")
torch.save({
"encoder": model.encoder.state_dict(),
"projector": model.projector.state_dict(),
"step": step + 1,
}, ckpt)
print(f" saved {ckpt}")
_save_stage1(model, opt, ema, args.ckpt_dir, step + 1)
step += 1
# final ckpt
ckpt = os.path.join(args.ckpt_dir, "stage1_final.pt")
torch.save({
_save_stage1(model, opt, ema, args.ckpt_dir, step, final=True)
print(f"[stage1] done. saved to {args.ckpt_dir} avg_loss={accum_loss/max(step,1):.4f} ema={ema.value:.4f}")
if writer is not None:
writer.close()
def _save_stage1(model, opt, ema, ckpt_dir: str, step: int, final: bool = False) -> None:
"""Save weights + optimizer + EMA so a kill mid-run can resume seamlessly."""
train_state = {
"optimizer": opt.state_dict(),
"ema": ema.value,
}
payload = {
"encoder": model.encoder.state_dict(),
"projector": model.projector.state_dict(),
"step": step,
}, ckpt)
print(f"[stage1] done. saved {ckpt} avg_loss={accum_loss/max(step,1):.4f}")
if writer is not None:
writer.close()
"train_state": train_state,
}
name = "stage1_final.pt" if final else f"stage1_step{step}.pt"
path = os.path.join(ckpt_dir, name)
torch.save(payload, path)
print(f" saved {name} ema={ema.value:.4f}")
def main() -> None:
@@ -169,6 +242,10 @@ def main() -> None:
p.add_argument("--ckpt_every", type=int, default=2000)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--tensorboard", action="store_true")
p.add_argument("--resume", action="store_true",
help="auto-pick the latest step ckpt in --ckpt_dir and continue")
p.add_argument("--resume_from", type=str, default=None,
help="explicit ckpt path to resume from (overrides --resume)")
args = p.parse_args()
train(args)
+108 -9
View File
@@ -27,15 +27,85 @@ def load_stage1_ckpt(model, ckpt_path: str) -> None:
print(f"[stage2] loaded stage1 ckpt: {ckpt_path} (step {state.get('step','?')})")
class EMA:
"""Exponential moving average of a scalar — see stage1.EMA for rationale."""
def __init__(self, alpha: float = 0.05) -> None:
self.alpha = alpha
self.value = None
def update(self, x: float) -> float:
if self.value is None:
self.value = float(x)
else:
self.value = (1 - self.alpha) * self.value + self.alpha * float(x)
return self.value
def _latest_step(ckpt_dir: str, prefix: str):
import glob
steps = sorted(glob.glob(os.path.join(ckpt_dir, f"{prefix}step*.pt")))
return steps[-1] if steps else None
def _step_number(path: str) -> int:
base = os.path.basename(path)
num = "".join(ch for ch in base if ch.isdigit())
return int(num) if num else 0
def train(args: argparse.Namespace) -> None:
device = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(args.seed)
# Decide resume BEFORE enabling LoRA: a stage2 resume ckpt carries its own
# encoder/projector (so we skip the stage1 load); a fresh run loads stage1.
resume_path = args.resume_from
if not resume_path and args.resume:
resume_path = _latest_step(args.ckpt_dir, "stage2_")
final = os.path.join(args.ckpt_dir, "stage2_final.pt")
if os.path.exists(final):
if resume_path is None or _step_number(final) > _step_number(resume_path):
resume_path = final
model = build_model().to(device)
if args.stage1_ckpt and os.path.exists(args.stage1_ckpt):
# When resuming, encoder/projector come from the stage2 ckpt (which already
# absorbed stage1 + further SFT); only a fresh run bootstraps from stage1.
if resume_path and os.path.exists(resume_path):
st = torch.load(resume_path, map_location="cpu")
model.encoder.load_state_dict(st["encoder"])
model.projector.load_state_dict(st["projector"])
resume_step = int(st.get("step", 0))
print(f"[stage2] resumed weights from {resume_path} at step {resume_step}")
elif args.stage1_ckpt and os.path.exists(args.stage1_ckpt):
load_stage1_ckpt(model, args.stage1_ckpt)
resume_step = 0
else:
resume_step = 0
model.freeze_llm()
# Build LoRA structure the SAME way for fresh and resume (so the trainable
# param set matches the saved optimizer state). For resume we then overlay
# the saved adapter weights via load_adapter (which does NOT reset
# requires_grad, unlike PeftModel.from_pretrained).
model.enable_lora(r=args.lora_r)
resume_lora_loaded = False
if resume_path and os.path.exists(resume_path):
lora_dir = os.path.join(args.ckpt_dir, "lora_adapter")
if os.path.exists(lora_dir):
try:
from peft import set_peft_model_state_dict
import json as _json
# load adapter weights into the already-built LoRA model
sd_path = os.path.join(lora_dir, "adapter_model.safetensors")
if os.path.exists(sd_path):
from safetensors.torch import load_file
sd = load_file(sd_path)
set_peft_model_state_dict(model.llm, sd)
print(f"[stage2] resumed LoRA weights from {lora_dir}")
resume_lora_loaded = True
except Exception as e:
print(f"[stage2] LoRA weight load failed ({e}), using fresh LoRA")
model.train()
if hasattr(model.llm, "gradient_checkpointing_enable"):
model.llm.gradient_checkpointing_enable()
@@ -47,17 +117,34 @@ def train(args: argparse.Namespace) -> None:
params = [p for p in model.parameters() if p.requires_grad]
opt = torch.optim.AdamW(params, lr=args.lr)
# restore optimizer + EMA when resuming (LoRA already rebuilt above to match)
ema = EMA(alpha=0.05)
if resume_path and os.path.exists(resume_path):
st = torch.load(resume_path, map_location="cpu")
ts = st.get("train_state", {})
if "optimizer" in ts:
try:
opt.load_state_dict(ts["optimizer"])
for g in opt.param_groups:
g["lr"] = args.lr
print(f"[stage2] resumed optimizer from {resume_path}")
except (ValueError, KeyError) as e:
print(f"[stage2] optimizer state mismatch ({e}), fresh AdamW")
if "ema" in ts:
ema.value = ts["ema"]
os.makedirs(args.ckpt_dir, exist_ok=True)
writer = None
if args.tensorboard:
from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter(os.path.join(args.ckpt_dir, "tb"))
step = 0
step = resume_step
opt.zero_grad()
accum = 0.0
print(f"[stage2] training: max_steps={args.max_steps} bs={args.bs} "
f"grad_accum={args.grad_accum} ctx={args.ctx} lora_r={args.lora_r}")
f"grad_accum={args.grad_accum} ctx={args.ctx} lora_r={args.lora_r} "
f"start_step={step}")
while step < args.max_steps:
for batch_samples in batched(stream_jsonl(args.data), args.bs):
if step >= args.max_steps:
@@ -73,28 +160,36 @@ def train(args: argparse.Namespace) -> None:
opt.step()
opt.zero_grad()
ema.update(float(loss))
if writer is not None:
writer.add_scalar("sft_loss", float(loss), step)
writer.add_scalar("sft_loss_ema", ema.value, step)
if step % args.log_every == 0:
peak = (torch.cuda.max_memory_allocated() / 1e9) if torch.cuda.is_available() else 0.0
print(f" step {step:5d} sft_loss={float(loss):.4f} peak_GB={peak:.2f}")
print(f" step {step:5d} sft_loss={float(loss):.4f} ema={ema.value:.4f} peak_GB={peak:.2f}")
if (step + 1) % args.ckpt_every == 0:
save_stage2_ckpt(model, args.ckpt_dir, step + 1)
save_stage2_ckpt(model, opt, ema, args.ckpt_dir, step + 1)
step += 1
save_stage2_ckpt(model, args.ckpt_dir, step, final=True)
print(f"[stage2] done. saved to {args.ckpt_dir} avg_loss={accum/max(step,1):.4f}")
save_stage2_ckpt(model, opt, ema, args.ckpt_dir, step, final=True)
print(f"[stage2] done. saved to {args.ckpt_dir} avg_loss={accum/max(step,1):.4f} ema={ema.value:.4f}")
if writer is not None:
writer.close()
def save_stage2_ckpt(model, ckpt_dir: str, step: int, final: bool = False) -> None:
def save_stage2_ckpt(model, opt, ema, ckpt_dir: str, step: int, final: bool = False) -> None:
"""Save encoder/projector + LoRA adapter + optimizer + EMA for resume."""
train_state = {
"optimizer": opt.state_dict(),
"ema": ema.value,
}
payload = {
"encoder": model.encoder.state_dict(),
"projector": model.projector.state_dict(),
"step": step,
"train_state": train_state,
}
name = "stage2_final.pt" if final else f"stage2_step{step}.pt"
torch.save(payload, os.path.join(ckpt_dir, name))
@@ -103,7 +198,7 @@ def save_stage2_ckpt(model, ckpt_dir: str, step: int, final: bool = False) -> No
model.llm.save_pretrained(os.path.join(ckpt_dir, "lora_adapter"))
except Exception as e: # pragma: no cover - best-effort adapter dump
print(f" (lora adapter save skipped: {e})")
print(f" saved {name}")
print(f" saved {name} ema={ema.value:.4f}")
def main() -> None:
@@ -121,6 +216,10 @@ def main() -> None:
p.add_argument("--ckpt_every", type=int, default=2000)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--tensorboard", action="store_true")
p.add_argument("--resume", action="store_true",
help="auto-pick latest step ckpt + lora_adapter and continue")
p.add_argument("--resume_from", type=str, default=None,
help="explicit ckpt path to resume from (overrides --resume)")
args = p.parse_args()
train(args)