Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
197 lines
6.6 KiB
Python
197 lines
6.6 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.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()
|