Save SFT last/best checkpoints during training and on interrupt
train_sft used to torch.save only after the full epoch budget, so Ctrl+C dropped all translation weights. Write _last every --ckpt-every steps and on KeyboardInterrupt; write _best when frozen eval (success, chrF) improves; --resume continues from _last.
This commit is contained in:
@@ -0,0 +1,13 @@
|
|||||||
|
from train_sft import _is_better_eval, _sibling
|
||||||
|
|
||||||
|
|
||||||
|
def test_sibling_last_best():
|
||||||
|
assert _sibling("ckpts/k3_sft.pt", "_last") == "ckpts/k3_sft_last.pt"
|
||||||
|
assert _sibling("ckpts/k3_sft.pt", "_best") == "ckpts/k3_sft_best.pt"
|
||||||
|
|
||||||
|
|
||||||
|
def test_best_prefers_success_then_chrf():
|
||||||
|
assert _is_better_eval(1.0, 70.0, 0.95, 90.0)
|
||||||
|
assert not _is_better_eval(0.95, 99.0, 1.0, 70.0)
|
||||||
|
assert _is_better_eval(1.0, 91.0, 1.0, 81.0)
|
||||||
|
assert not _is_better_eval(1.0, 70.0, 1.0, 81.0)
|
||||||
+249
-85
@@ -4,6 +4,7 @@
|
|||||||
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/toy.jsonl
|
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/toy.jsonl
|
||||||
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\
|
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\
|
||||||
--limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2
|
--limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2
|
||||||
|
uv run python train_sft.py --resume --out ckpts/k3_sft.pt
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -23,6 +24,7 @@ from kda.training.data import (
|
|||||||
)
|
)
|
||||||
from kda.training.eval_mt import evaluate_pairs
|
from kda.training.eval_mt import evaluate_pairs
|
||||||
from kda.training.schedule import lr_scale, total_opt_steps
|
from kda.training.schedule import lr_scale, total_opt_steps
|
||||||
|
from kda.training.swanlab_env import prepare_swanlab_env, swanlab_run_id
|
||||||
from kda.training.toy import load_ckpt
|
from kda.training.toy import load_ckpt
|
||||||
|
|
||||||
|
|
||||||
@@ -31,20 +33,43 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
|
|||||||
group["lr"] = lr
|
group["lr"] = lr
|
||||||
|
|
||||||
|
|
||||||
def _init_swanlab(args: argparse.Namespace):
|
def _sibling(path: str, suffix: str) -> str:
|
||||||
|
root, ext = os.path.splitext(path)
|
||||||
|
return f"{root}{suffix}{ext}"
|
||||||
|
|
||||||
|
|
||||||
|
def _save(path: str, payload: dict) -> None:
|
||||||
|
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||||
|
tmp = path + ".tmp"
|
||||||
|
torch.save(payload, tmp)
|
||||||
|
os.replace(tmp, path)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_better_eval(
|
||||||
|
success: float, chrf: float, best_success: float, best_chrf: float
|
||||||
|
) -> bool:
|
||||||
|
if success > best_success:
|
||||||
|
return True
|
||||||
|
if success == best_success and chrf > best_chrf:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _init_swanlab(args: argparse.Namespace, resume_id: str | None = None):
|
||||||
key = os.environ.get("SWANLAB_API_KEY")
|
key = os.environ.get("SWANLAB_API_KEY")
|
||||||
if not key:
|
if not key:
|
||||||
return None
|
return None
|
||||||
|
project = prepare_swanlab_env()
|
||||||
try:
|
try:
|
||||||
import swanlab
|
import swanlab
|
||||||
except ImportError:
|
except ImportError:
|
||||||
|
print("SWANLAB_API_KEY set but swanlab is not installed")
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
swanlab.login(api_key=key, save=False)
|
swanlab.login(api_key=key, save=False)
|
||||||
project = os.environ.pop("SWANLAB_PROJECT", None) or "kda"
|
init_kw = dict(
|
||||||
return swanlab.init(
|
|
||||||
project=project,
|
project=project,
|
||||||
name=f"sft-{os.path.basename(args.ckpt)}",
|
name=f"sft-{os.path.basename(args.ckpt or args.out)}",
|
||||||
config={
|
config={
|
||||||
"ckpt": args.ckpt,
|
"ckpt": args.ckpt,
|
||||||
"data": args.data,
|
"data": args.data,
|
||||||
@@ -52,8 +77,20 @@ def _init_swanlab(args: argparse.Namespace):
|
|||||||
"batch": args.batch,
|
"batch": args.batch,
|
||||||
"seq_len": args.seq_len,
|
"seq_len": args.seq_len,
|
||||||
"epochs": args.epochs,
|
"epochs": args.epochs,
|
||||||
|
"grad_acc": args.grad_acc,
|
||||||
|
"limit": args.limit,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
if resume_id:
|
||||||
|
init_kw["id"] = resume_id
|
||||||
|
init_kw["resume"] = True
|
||||||
|
print(f"swanlab resume id={resume_id}")
|
||||||
|
run = swanlab.init(**init_kw)
|
||||||
|
got = swanlab_run_id(run)
|
||||||
|
if got:
|
||||||
|
args.swanlab_id = got
|
||||||
|
print(f"swanlab run id {got}")
|
||||||
|
return run
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
print(f"swanlab init failed ({exc}); continuing without cloud monitor")
|
print(f"swanlab init failed ({exc}); continuing without cloud monitor")
|
||||||
return None
|
return None
|
||||||
@@ -69,9 +106,50 @@ def _read_lines(path: str) -> list[str]:
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _payload(
|
||||||
|
cfg,
|
||||||
|
model,
|
||||||
|
optim: torch.optim.Optimizer,
|
||||||
|
args: argparse.Namespace,
|
||||||
|
*,
|
||||||
|
tok_src: str,
|
||||||
|
step: int,
|
||||||
|
opt_step: int,
|
||||||
|
row_index: int,
|
||||||
|
best_success: float,
|
||||||
|
best_chrf: float,
|
||||||
|
best_train: float,
|
||||||
|
):
|
||||||
|
return {
|
||||||
|
"config": asdict(cfg),
|
||||||
|
"model_state": model.state_dict(),
|
||||||
|
"optimizer_state": optim.state_dict(),
|
||||||
|
"tokenizer": tok_src,
|
||||||
|
"sft_data": args.data,
|
||||||
|
"pretrained_ckpt": args.ckpt,
|
||||||
|
"sft_step": step,
|
||||||
|
"opt_step": opt_step,
|
||||||
|
"row_index": row_index,
|
||||||
|
"best_success": best_success,
|
||||||
|
"best_chrf": best_chrf,
|
||||||
|
"best_train": best_train,
|
||||||
|
"batch": args.batch,
|
||||||
|
"seq_len": args.seq_len,
|
||||||
|
"grad_acc": args.grad_acc,
|
||||||
|
"swanlab_id": getattr(args, "swanlab_id", None),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
def main() -> None:
|
||||||
p = argparse.ArgumentParser(description=__doc__)
|
p = argparse.ArgumentParser(description=__doc__)
|
||||||
p.add_argument("--ckpt", required=True)
|
p.add_argument("--ckpt", default=None, help="pretrained (or SFT) checkpoint to start from")
|
||||||
|
p.add_argument(
|
||||||
|
"--resume",
|
||||||
|
nargs="?",
|
||||||
|
const="__last__",
|
||||||
|
default=None,
|
||||||
|
help="resume SFT; default path is <out>_last",
|
||||||
|
)
|
||||||
p.add_argument(
|
p.add_argument(
|
||||||
"--data",
|
"--data",
|
||||||
default="opus-100",
|
default="opus-100",
|
||||||
@@ -98,12 +176,26 @@ def main() -> None:
|
|||||||
p.add_argument("--max-steps", type=int, default=None)
|
p.add_argument("--max-steps", type=int, default=None)
|
||||||
p.add_argument("--grad-acc", type=int, default=1)
|
p.add_argument("--grad-acc", type=int, default=1)
|
||||||
p.add_argument("--eval-every", type=int, default=50)
|
p.add_argument("--eval-every", type=int, default=50)
|
||||||
|
p.add_argument(
|
||||||
|
"--ckpt-every",
|
||||||
|
type=int,
|
||||||
|
default=500,
|
||||||
|
help="write _last this many steps (0 = only interrupt + end)",
|
||||||
|
)
|
||||||
p.add_argument("--src", default=None, help="frozen eval src (not used as train)")
|
p.add_argument("--src", default=None, help="frozen eval src (not used as train)")
|
||||||
p.add_argument("--ref", default=None)
|
p.add_argument("--ref", default=None)
|
||||||
p.add_argument("--target-lang", default="en", choices=["en", "zh"])
|
p.add_argument("--target-lang", default="en", choices=["en", "zh"])
|
||||||
p.add_argument("--device", default="auto")
|
p.add_argument("--device", default="auto")
|
||||||
args = p.parse_args()
|
args = p.parse_args()
|
||||||
|
|
||||||
|
resume_path = args.resume
|
||||||
|
if resume_path == "__last__":
|
||||||
|
resume_path = _sibling(args.out, "_last")
|
||||||
|
if resume_path is None and not args.ckpt:
|
||||||
|
raise SystemExit("need --ckpt or --resume")
|
||||||
|
if resume_path is not None and not os.path.isfile(resume_path):
|
||||||
|
raise SystemExit(f"resume checkpoint not found: {resume_path}")
|
||||||
|
|
||||||
device = args.device
|
device = args.device
|
||||||
if device == "auto":
|
if device == "auto":
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
@@ -111,10 +203,13 @@ def main() -> None:
|
|||||||
if device == "cuda" and not use_bf16:
|
if device == "cuda" and not use_bf16:
|
||||||
raise SystemExit("KDA training needs bf16")
|
raise SystemExit("KDA training needs bf16")
|
||||||
|
|
||||||
model, cfg = load_ckpt(args.ckpt)
|
start_path = resume_path or args.ckpt
|
||||||
|
model, cfg = load_ckpt(start_path)
|
||||||
model.to(device)
|
model.to(device)
|
||||||
payload = torch.load(args.ckpt, map_location="cpu", weights_only=False)
|
loaded = torch.load(start_path, map_location="cpu", weights_only=False)
|
||||||
tok_src = args.tokenizer or payload.get("tokenizer")
|
if args.ckpt is None:
|
||||||
|
args.ckpt = loaded.get("pretrained_ckpt")
|
||||||
|
tok_src = args.tokenizer or loaded.get("tokenizer")
|
||||||
if not tok_src:
|
if not tok_src:
|
||||||
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
|
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
|
||||||
tok = load_tokenizer(tok_src)
|
tok = load_tokenizer(tok_src)
|
||||||
@@ -125,7 +220,6 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
if not rows:
|
if not rows:
|
||||||
raise SystemExit(f"no SFT rows from {args.data}")
|
raise SystemExit(f"no SFT rows from {args.data}")
|
||||||
print(f"SFT {len(rows)} rows from {args.data}; model {cfg.__class__.__name__}")
|
|
||||||
|
|
||||||
steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1)
|
steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1)
|
||||||
max_micro = args.max_steps
|
max_micro = args.max_steps
|
||||||
@@ -138,87 +232,157 @@ def main() -> None:
|
|||||||
seq_len=args.seq_len,
|
seq_len=args.seq_len,
|
||||||
grad_acc=args.grad_acc,
|
grad_acc=args.grad_acc,
|
||||||
)
|
)
|
||||||
|
print(
|
||||||
|
f"SFT {len(rows)} rows from {args.data}; model {cfg.__class__.__name__} "
|
||||||
|
f"{max_micro} steps ({args.epochs} epoch, batch {args.batch}) "
|
||||||
|
f"eval/{args.eval_every} ckpt/{args.ckpt_every}"
|
||||||
|
)
|
||||||
|
|
||||||
optim = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01)
|
optim = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01)
|
||||||
tracker = _init_swanlab(args)
|
|
||||||
model.train()
|
|
||||||
step = 0
|
step = 0
|
||||||
opt_step = 0
|
opt_step = 0
|
||||||
best = float("inf")
|
row_start = 0
|
||||||
|
best_train = float("inf")
|
||||||
|
best_success = -1.0
|
||||||
|
best_chrf = -1.0
|
||||||
|
if resume_path is not None:
|
||||||
|
if loaded.get("optimizer_state"):
|
||||||
|
optim.load_state_dict(loaded["optimizer_state"])
|
||||||
|
step = int(loaded.get("sft_step", 0))
|
||||||
|
opt_step = int(loaded.get("opt_step", 0))
|
||||||
|
row_start = int(loaded.get("row_index", 0))
|
||||||
|
best_train = float(loaded.get("best_train", best_train))
|
||||||
|
best_success = float(loaded.get("best_success", best_success))
|
||||||
|
best_chrf = float(loaded.get("best_chrf", best_chrf))
|
||||||
|
print(
|
||||||
|
f"resume {resume_path} step {step} opt {opt_step} "
|
||||||
|
f"row {row_start} best success {best_success:.2f} chrf {best_chrf:.2f}"
|
||||||
|
)
|
||||||
|
|
||||||
for _, x, y in iter_sft_batches(rows, tok, args.batch, args.seq_len):
|
tracker = _init_swanlab(
|
||||||
if step >= max_micro:
|
args,
|
||||||
break
|
resume_id=loaded.get("swanlab_id") if resume_path is not None else None,
|
||||||
x, y = x.to(device), y.to(device)
|
|
||||||
_set_lr(optim, args.lr * lr_scale(opt_step, args.warmup, horizon))
|
|
||||||
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16):
|
|
||||||
task = model(x, labels=y, ignore_index=IGNORE_INDEX)
|
|
||||||
aux, z_loss = moe_router_losses(model)
|
|
||||||
loss = (task + aux + z_loss) / args.grad_acc
|
|
||||||
loss.backward()
|
|
||||||
if (step + 1) % args.grad_acc == 0:
|
|
||||||
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
|
||||||
optim.step()
|
|
||||||
optim.zero_grad(set_to_none=True)
|
|
||||||
opt_step += 1
|
|
||||||
raw = float(task.detach())
|
|
||||||
if raw < best:
|
|
||||||
best = raw
|
|
||||||
if tracker is not None:
|
|
||||||
tracker.log(
|
|
||||||
{
|
|
||||||
"sft/loss": raw,
|
|
||||||
"sft/lr": optim.param_groups[0]["lr"],
|
|
||||||
"moe/aux": float(aux.detach()),
|
|
||||||
"moe/z": float(z_loss.detach()),
|
|
||||||
},
|
|
||||||
step=step,
|
|
||||||
)
|
|
||||||
if step % args.eval_every == 0 or step == max_micro - 1:
|
|
||||||
print(
|
|
||||||
f"step {step:4d} sft loss {raw:.4f} lr {optim.param_groups[0]['lr']:.2e}"
|
|
||||||
)
|
|
||||||
if args.src and args.ref:
|
|
||||||
model.eval()
|
|
||||||
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
|
|
||||||
out = evaluate_pairs(
|
|
||||||
model,
|
|
||||||
tok,
|
|
||||||
srcs,
|
|
||||||
refs,
|
|
||||||
target_lang=args.target_lang,
|
|
||||||
device=device,
|
|
||||||
max_new=64,
|
|
||||||
limit=None,
|
|
||||||
)
|
|
||||||
printable = {k: v for k, v in out.items() if k != "hyps"}
|
|
||||||
print(printable)
|
|
||||||
if tracker is not None:
|
|
||||||
tracker.log(
|
|
||||||
{
|
|
||||||
"eval/success_rate": printable["success_rate"],
|
|
||||||
"eval/chrf": printable["chrf"],
|
|
||||||
"eval/copy_rate": printable["copy_rate"],
|
|
||||||
},
|
|
||||||
step=step,
|
|
||||||
)
|
|
||||||
model.train()
|
|
||||||
step += 1
|
|
||||||
|
|
||||||
os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
|
|
||||||
torch.save(
|
|
||||||
{
|
|
||||||
"config": asdict(cfg),
|
|
||||||
"model_state": model.state_dict(),
|
|
||||||
"optimizer_state": optim.state_dict(),
|
|
||||||
"tokenizer": tok_src,
|
|
||||||
"sft_data": args.data,
|
|
||||||
"pretrained_ckpt": args.ckpt,
|
|
||||||
},
|
|
||||||
args.out,
|
|
||||||
)
|
)
|
||||||
print(f"best sft loss {best:.4f}; checkpoint -> {args.out}")
|
last_path = _sibling(args.out, "_last")
|
||||||
if tracker is not None:
|
best_path = _sibling(args.out, "_best")
|
||||||
tracker.finish()
|
model.train()
|
||||||
|
next_row = row_start
|
||||||
|
|
||||||
|
def dump() -> dict:
|
||||||
|
return _payload(
|
||||||
|
cfg,
|
||||||
|
model,
|
||||||
|
optim,
|
||||||
|
args,
|
||||||
|
tok_src=tok_src,
|
||||||
|
step=step,
|
||||||
|
opt_step=opt_step,
|
||||||
|
row_index=next_row,
|
||||||
|
best_success=best_success,
|
||||||
|
best_chrf=best_chrf,
|
||||||
|
best_train=best_train,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
for row_index, x, y in iter_sft_batches(
|
||||||
|
rows, tok, args.batch, args.seq_len, start=row_start
|
||||||
|
):
|
||||||
|
if step >= max_micro:
|
||||||
|
break
|
||||||
|
x, y = x.to(device), y.to(device)
|
||||||
|
_set_lr(optim, args.lr * lr_scale(opt_step, args.warmup, horizon))
|
||||||
|
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16):
|
||||||
|
task = model(x, labels=y, ignore_index=IGNORE_INDEX)
|
||||||
|
aux, z_loss = moe_router_losses(model)
|
||||||
|
loss = (task + aux + z_loss) / args.grad_acc
|
||||||
|
loss.backward()
|
||||||
|
if (step + 1) % args.grad_acc == 0:
|
||||||
|
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
||||||
|
optim.step()
|
||||||
|
optim.zero_grad(set_to_none=True)
|
||||||
|
opt_step += 1
|
||||||
|
raw = float(task.detach())
|
||||||
|
if raw < best_train:
|
||||||
|
best_train = raw
|
||||||
|
next_row = row_index + args.batch
|
||||||
|
if tracker is not None:
|
||||||
|
tracker.log(
|
||||||
|
{
|
||||||
|
"sft/loss": raw,
|
||||||
|
"sft/lr": optim.param_groups[0]["lr"],
|
||||||
|
"moe/aux": float(aux.detach()),
|
||||||
|
"moe/z": float(z_loss.detach()),
|
||||||
|
},
|
||||||
|
step=step,
|
||||||
|
)
|
||||||
|
eval_now = step % args.eval_every == 0 or step == max_micro - 1
|
||||||
|
ckpt_now = args.ckpt_every > 0 and step > 0 and (
|
||||||
|
step % args.ckpt_every == 0 or step == max_micro - 1
|
||||||
|
)
|
||||||
|
if eval_now:
|
||||||
|
print(
|
||||||
|
f"step {step:4d} sft loss {raw:.4f} "
|
||||||
|
f"lr {optim.param_groups[0]['lr']:.2e}"
|
||||||
|
)
|
||||||
|
if args.src and args.ref:
|
||||||
|
model.eval()
|
||||||
|
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
|
||||||
|
out = evaluate_pairs(
|
||||||
|
model,
|
||||||
|
tok,
|
||||||
|
srcs,
|
||||||
|
refs,
|
||||||
|
target_lang=args.target_lang,
|
||||||
|
device=device,
|
||||||
|
max_new=64,
|
||||||
|
limit=None,
|
||||||
|
)
|
||||||
|
printable = {k: v for k, v in out.items() if k != "hyps"}
|
||||||
|
print(printable)
|
||||||
|
if tracker is not None:
|
||||||
|
tracker.log(
|
||||||
|
{
|
||||||
|
"eval/success_rate": printable["success_rate"],
|
||||||
|
"eval/chrf": printable["chrf"],
|
||||||
|
"eval/copy_rate": printable["copy_rate"],
|
||||||
|
},
|
||||||
|
step=step,
|
||||||
|
)
|
||||||
|
if _is_better_eval(
|
||||||
|
printable["success_rate"],
|
||||||
|
printable["chrf"],
|
||||||
|
best_success,
|
||||||
|
best_chrf,
|
||||||
|
):
|
||||||
|
best_success = float(printable["success_rate"])
|
||||||
|
best_chrf = float(printable["chrf"])
|
||||||
|
_save(best_path, dump())
|
||||||
|
print(
|
||||||
|
f" best success {best_success:.2f} "
|
||||||
|
f"chrf {best_chrf:.2f} -> {best_path}"
|
||||||
|
)
|
||||||
|
model.train()
|
||||||
|
if ckpt_now:
|
||||||
|
_save(last_path, dump())
|
||||||
|
print(f" last -> {last_path}")
|
||||||
|
step += 1
|
||||||
|
payload = dump()
|
||||||
|
_save(last_path, payload)
|
||||||
|
_save(args.out, payload)
|
||||||
|
print(
|
||||||
|
f"best train {best_train:.4f} best success {best_success:.2f} "
|
||||||
|
f"chrf {best_chrf:.2f}; checkpoint -> {args.out}"
|
||||||
|
)
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
print("interrupt; writing last checkpoint")
|
||||||
|
_save(last_path, dump())
|
||||||
|
print(f" last -> {last_path}")
|
||||||
|
if os.path.isfile(best_path):
|
||||||
|
print(f" best remains {best_path}")
|
||||||
|
raise SystemExit(130) from None
|
||||||
|
finally:
|
||||||
|
if tracker is not None:
|
||||||
|
tracker.finish()
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user