Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
+196
@@ -0,0 +1,196 @@
|
||||
"""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.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):
|
||||
loss = model(x, labels=y, ignore_index=IGNORE_INDEX) / 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 = loss.item() * args.grad_acc
|
||||
if raw < best:
|
||||
best = raw
|
||||
if tracker is not None:
|
||||
tracker.log(
|
||||
{"sft/loss": raw, "sft/lr": optim.param_groups[0]["lr"]},
|
||||
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()
|
||||
Reference in New Issue
Block a user