Files
K3/train_sft.py
T
dela 7a12f61de1 LatentMoE: K3 sigmoid routing and Switch aux/z-loss
Route with σ(W_r x), Top-k(s+b), then L1-normalize over the selected set.
Add Switch/GShard aux and router z-loss into train_k3 and train_sft.
Wiki parquet URLs honor HF_ENDPOINT for mirrored downloads.
2026-08-25 19:50:07 +08:00

207 lines
6.9 KiB
Python

"""Instruction SFT for zh↔en translation. Prompt template matches eval_mt.
用法:
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/train.jsonl
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data data/sft/opus.jsonl \\
--seq-len 512 --batch 4 --lr 5e-5 --epochs 2
"""
from __future__ import annotations
import argparse
import os
from dataclasses import asdict
import torch
from kda.layers.latent_moe import moe_router_losses
from kda.training.data import (
IGNORE_INDEX,
iter_sft_batches,
load_sft_rows,
load_tokenizer,
)
from kda.training.eval_mt import evaluate_pairs
from kda.training.schedule import lr_scale, total_opt_steps
from kda.training.toy import load_ckpt
def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
for group in optim.param_groups:
group["lr"] = lr
def _init_swanlab(args: argparse.Namespace):
key = os.environ.get("SWANLAB_API_KEY")
if not key:
return None
try:
import swanlab
except ImportError:
return None
try:
swanlab.login(api_key=key, save=False)
project = os.environ.pop("SWANLAB_PROJECT", None) or "kda"
return swanlab.init(
project=project,
name=f"sft-{os.path.basename(args.ckpt)}",
config={
"ckpt": args.ckpt,
"data": args.data,
"lr": args.lr,
"batch": args.batch,
"seq_len": args.seq_len,
"epochs": args.epochs,
},
)
except Exception as exc:
print(f"swanlab init failed ({exc}); continuing without cloud monitor")
return None
def _read_lines(path: str) -> list[str]:
from pathlib import Path
return [
ln.strip()
for ln in Path(path).read_text(encoding="utf-8").splitlines()
if ln.strip()
]
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--ckpt", required=True)
p.add_argument("--data", required=True, help="jsonl {src,tgt,target_lang} or TSV")
p.add_argument("--out", default="ckpts/k3_sft.pt")
p.add_argument("--tokenizer", default=None)
p.add_argument("--batch", type=int, default=4)
p.add_argument("--seq-len", type=int, default=256)
p.add_argument("--lr", type=float, default=1e-3)
p.add_argument("--warmup", type=int, default=20)
p.add_argument("--epochs", type=int, default=2)
p.add_argument("--max-steps", type=int, default=None)
p.add_argument("--grad-acc", type=int, default=1)
p.add_argument("--eval-every", type=int, default=50)
p.add_argument("--src", default=None, help="frozen eval src (not used as train)")
p.add_argument("--ref", default=None)
p.add_argument("--target-lang", default="en", choices=["en", "zh"])
p.add_argument("--device", default="auto")
args = p.parse_args()
device = args.device
if device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
use_bf16 = device == "cuda" and torch.cuda.is_bf16_supported()
if device == "cuda" and not use_bf16:
raise SystemExit("KDA training needs bf16")
model, cfg = load_ckpt(args.ckpt)
model.to(device)
payload = torch.load(args.ckpt, map_location="cpu", weights_only=False)
tok_src = args.tokenizer or payload.get("tokenizer")
if not tok_src:
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
tok = load_tokenizer(tok_src)
rows = load_sft_rows(args.data)
if not rows:
raise SystemExit(f"no SFT rows in {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)
max_micro = args.max_steps
if max_micro is None:
max_micro = steps_per_epoch * args.epochs
horizon = total_opt_steps(
max_tokens=None,
max_micro=max_micro,
batch=args.batch,
seq_len=args.seq_len,
grad_acc=args.grad_acc,
)
optim = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01)
tracker = _init_swanlab(args)
model.train()
step = 0
opt_step = 0
best = float("inf")
for _, x, y in iter_sft_batches(rows, tok, args.batch, args.seq_len):
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:
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}")
if tracker is not None:
tracker.finish()
if __name__ == "__main__":
main()