Compare commits

..
2 Commits
Author SHA1 Message Date
dela 9652a9a7eb 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.
2026-08-26 10:00:11 +08:00
dela 071dfaf42c Sanitize SwanLab env before login so 0.9 nested project does not crash
OpenBayes sets SWANLAB_PROJECT as a string; swanlab>=0.9 parses that as
ProjectSettings and raises QuoteAwareEnvSettingsSource. Drop it, keep
SWANLAB_PROJ_NAME, and share run-id extraction with train_k3.
2026-08-26 10:00:06 +08:00
5 changed files with 351 additions and 106 deletions
+41
View File
@@ -0,0 +1,41 @@
"""Sanitize SwanLab env before import/init.
swanlab>=0.9 ``Settings.project`` is a nested model. A string
``SWANLAB_PROJECT`` (OpenBayes and older docs) makes pydantic raise
``error parsing value for field "project" from source
_QuoteAwareEnvSettingsSource``. Project name belongs in
``SWANLAB_PROJ_NAME`` / ``init(project=...)``.
"""
from __future__ import annotations
import os
def prepare_swanlab_env(default_project: str = "kda") -> str:
"""Drop nested ``SWANLAB_PROJECT``, keep a plain project name.
Must run before ``import swanlab`` / ``swanlab.login`` / ``init``.
"""
raw = os.environ.pop("SWANLAB_PROJECT", None)
name = os.environ.get("SWANLAB_PROJ_NAME") or raw or default_project
name = str(name).strip().strip("\"'")
if not name or name[0] in "{[":
name = default_project
os.environ.pop("SWANLAB_PROJECT", None)
os.environ["SWANLAB_PROJ_NAME"] = name
return name
def swanlab_run_id(run) -> str | None:
for attr in ("id", "run_id"):
val = getattr(run, attr, None)
if isinstance(val, str) and val:
return val
public = getattr(run, "public", None)
if public is not None:
for attr in ("cloud_run_id", "run_id", "id"):
val = getattr(public, attr, None)
if isinstance(val, str) and val:
return val
return None
+13
View File
@@ -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)
+39
View File
@@ -0,0 +1,39 @@
import os
from kda.training.swanlab_env import prepare_swanlab_env, swanlab_run_id
def test_prepare_drops_string_project(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", "kda")
monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False)
assert prepare_swanlab_env() == "kda"
assert "SWANLAB_PROJECT" not in os.environ
assert os.environ["SWANLAB_PROJ_NAME"] == "kda"
def test_prepare_strips_quotes(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", '"kda"')
monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False)
assert prepare_swanlab_env() == "kda"
assert "SWANLAB_PROJECT" not in os.environ
def test_prepare_prefers_proj_name(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", "ignored")
monkeypatch.setenv("SWANLAB_PROJ_NAME", "mine")
assert prepare_swanlab_env() == "mine"
assert os.environ["SWANLAB_PROJ_NAME"] == "mine"
def test_prepare_rejects_json_blob(monkeypatch):
monkeypatch.setenv("SWANLAB_PROJECT", '{"name": "x"}')
monkeypatch.delenv("SWANLAB_PROJ_NAME", raising=False)
assert prepare_swanlab_env() == "kda"
def test_swanlab_run_id():
class _Run:
id = "ilgne5ro"
assert swanlab_run_id(_Run()) == "ilgne5ro"
assert swanlab_run_id(object()) is None
+9 -21
View File
@@ -22,6 +22,7 @@ from kda.models.causal_lm import CausalLM
from kda.models.k3_config import K3Config from kda.models.k3_config import K3Config
from kda.training.data import iter_indexed, load_pretrain_chunks, load_tokenizer from kda.training.data import iter_indexed, load_pretrain_chunks, load_tokenizer
from kda.training.schedule import lr_scale, tokens_per_micro, total_opt_steps from kda.training.schedule import lr_scale, tokens_per_micro, 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
_TOY_TRAIN = { _TOY_TRAIN = {
@@ -66,25 +67,14 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
group["lr"] = lr group["lr"] = lr
def _swanlab_run_id(run) -> str | None: def _init_swanlab(
for attr in ("id", "run_id"): cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None
val = getattr(run, attr, None) ):
if isinstance(val, str) and val:
return val
public = getattr(run, "public", None)
if public is not None:
for attr in ("cloud_run_id", "run_id", "id"):
val = getattr(public, attr, None)
if isinstance(val, str) and val:
return val
return None
def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None):
"""Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op.""" """Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op."""
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:
@@ -92,8 +82,6 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None
return None return None
try: try:
swanlab.login(api_key=key, save=False) swanlab.login(api_key=key, save=False)
# swanlab 0.9 Settings.project is nested; a string SWANLAB_PROJECT env crashes init.
project = os.environ.pop("SWANLAB_PROJECT", None) or "kda"
run_id = resume_id or os.environ.get("SWANLAB_RUN_ID") run_id = resume_id or os.environ.get("SWANLAB_RUN_ID")
init_kw = dict( init_kw = dict(
project=project, project=project,
@@ -125,7 +113,7 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None
init_kw["resume"] = True init_kw["resume"] = True
print(f"swanlab resume id={run_id}") print(f"swanlab resume id={run_id}")
run = swanlab.init(**init_kw) run = swanlab.init(**init_kw)
got = _swanlab_run_id(run) got = swanlab_run_id(run)
if got: if got:
args.swanlab_id = got args.swanlab_id = got
print(f"swanlab run id {got}") print(f"swanlab run id {got}")
@@ -523,9 +511,9 @@ def main() -> None:
if elapsed > 0: if elapsed > 0:
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
ended = ( ended = (args.max_tokens is not None and tokens >= args.max_tokens) or (
args.max_tokens is not None and tokens >= args.max_tokens args.max_tokens is None and micro_step >= args.steps
) or (args.max_tokens is None and micro_step >= args.steps) )
log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended
eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended
ckpt_now = micro_step % args.ckpt_every == 0 or ended ckpt_now = micro_step % args.ckpt_every == 0 or ended
+249 -85
View File
@@ -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__":