Pull OPUS-100 en-zh for SFT instead of a checked-in jsonl
train_sft --data opus-100 streams Helsinki-NLP/opus-100, writes both directions, and skips frozen eval sentences. Runtime cache stays under data/sft/ (gitignored).
This commit is contained in:
+26
-7
@@ -1,9 +1,9 @@
|
||||
"""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
|
||||
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 \\
|
||||
--limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -18,7 +18,7 @@ from kda.layers.latent_moe import moe_router_losses
|
||||
from kda.training.data import (
|
||||
IGNORE_INDEX,
|
||||
iter_sft_batches,
|
||||
load_sft_rows,
|
||||
resolve_sft_rows,
|
||||
load_tokenizer,
|
||||
)
|
||||
from kda.training.eval_mt import evaluate_pairs
|
||||
@@ -72,7 +72,22 @@ def _read_lines(path: str) -> list[str]:
|
||||
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(
|
||||
"--data",
|
||||
default="opus-100",
|
||||
help="local jsonl/tsv, or 'opus-100' to stream Helsinki-NLP/opus-100 en-zh",
|
||||
)
|
||||
p.add_argument(
|
||||
"--limit",
|
||||
type=int,
|
||||
default=100_000,
|
||||
help="OPUS source pairs to pull (each becomes zh2en + en2zh unless --one-dir)",
|
||||
)
|
||||
p.add_argument(
|
||||
"--one-dir",
|
||||
action="store_true",
|
||||
help="only zh→en rows when pulling OPUS",
|
||||
)
|
||||
p.add_argument("--out", default="ckpts/k3_sft.pt")
|
||||
p.add_argument("--tokenizer", default=None)
|
||||
p.add_argument("--batch", type=int, default=4)
|
||||
@@ -103,9 +118,13 @@ def main() -> None:
|
||||
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)
|
||||
rows = resolve_sft_rows(
|
||||
args.data,
|
||||
limit=args.limit,
|
||||
both_dirs=not args.one_dir,
|
||||
)
|
||||
if not rows:
|
||||
raise SystemExit(f"no SFT rows in {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)
|
||||
|
||||
Reference in New Issue
Block a user