From e7185cbf49e51b11fc625576d723cc1a0ab54218 Mon Sep 17 00:00:00 2001 From: dela Date: Tue, 25 Aug 2026 21:39:58 +0800 Subject: [PATCH] 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). --- kda/training/data.py | 108 +++++++++++++++++++++++++++++ tests/integration/test_sft_data.py | 39 ++++++++++- train_sft.py | 33 +++++++-- 3 files changed, 172 insertions(+), 8 deletions(-) diff --git a/kda/training/data.py b/kda/training/data.py index ac20af3..3d8f1fe 100644 --- a/kda/training/data.py +++ b/kda/training/data.py @@ -89,6 +89,17 @@ def pretrain_dir() -> Path: return Path("data/pretrain") +def sft_dir() -> Path: + for candidate in ( + os.environ.get("KDA_SFT_DIR"), + "/data/sft", + "data/sft", + ): + if candidate and Path(candidate).is_dir(): + return Path(candidate) + return Path("data/sft") + + def _wiki_files(lang: str, n_shards: int) -> list[str]: if lang not in WIKI_SHARD_TOTAL: raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en") @@ -287,6 +298,103 @@ def encode_sft_row( return ids, labels +def _eval_blocklist(eval_dir: str | Path | None = None) -> set[str]: + """Frozen eval sentences must not appear in SFT bitext.""" + blocked: set[str] = set() + folders = [] + if eval_dir is not None: + folders.append(Path(eval_dir)) + folders.extend( + [ + Path(os.environ["KDA_EVAL_DIR"]) if os.environ.get("KDA_EVAL_DIR") else None, + Path("/data/eval"), + Path("data/eval"), + ] + ) + for folder in folders: + if folder is None or not folder.is_dir(): + continue + for path in folder.glob("*.txt"): + for line in path.read_text(encoding="utf-8").splitlines(): + text = line.strip() + if text: + blocked.add(text) + return blocked + + +def fetch_opus100_enzh( + limit: int, + *, + both_dirs: bool = True, + cache_dir: str | Path | None = None, + eval_dir: str | Path | None = None, +) -> list[dict]: + """Stream Helsinki-NLP/opus-100 ``en-zh`` train. ``limit`` is source pairs.""" + if limit < 1: + raise ValueError(f"limit must be >= 1, got {limit}") + cache = Path(cache_dir) if cache_dir is not None else sft_dir() + cache.mkdir(parents=True, exist_ok=True) + tag = "both" if both_dirs else "enzh" + path = cache / f"opus100-en-zh-{tag}-limit{limit}.jsonl" + if path.exists(): + rows = load_sft_rows(path) + if rows: + return rows + + from datasets import load_dataset + + ds = load_dataset("Helsinki-NLP/opus-100", "en-zh", split="train", streaming=True) + blocked = _eval_blocklist(eval_dir) + rows: list[dict] = [] + n_src = 0 + for row in ds: + trans = row.get("translation") if isinstance(row, dict) else None + blob = trans if isinstance(trans, dict) else row + en = str(blob.get("en") or "").strip() + zh = str(blob.get("zh") or "").strip() + if not en or not zh or en == zh: + continue + if en in blocked or zh in blocked: + continue + if min(len(en), len(zh)) < 2: + continue + n_src += 1 + rows.append({"src": zh, "tgt": en, "target_lang": "en"}) + if both_dirs: + rows.append({"src": en, "tgt": zh, "target_lang": "zh"}) + if n_src >= limit: + break + tmp = path.with_suffix(path.suffix + ".tmp") + with tmp.open("w", encoding="utf-8") as fh: + for row in rows: + fh.write(json.dumps(row, ensure_ascii=False) + "\n") + tmp.replace(path) + return rows + + +def resolve_sft_rows( + source: str, + *, + limit: int = 100_000, + both_dirs: bool = True, + cache_dir: str | Path | None = None, + eval_dir: str | Path | None = None, +) -> list[dict]: + """Local jsonl/tsv, or ``opus-100`` / ``opus`` to pull OPUS-100 en-zh from HF.""" + path = Path(source) + if path.is_file(): + return load_sft_rows(path) + key = source.strip().lower().replace("_", "-") + if key in {"opus", "opus-100", "opus100", "helsinki-nlp/opus-100"}: + print(f"fetching OPUS-100 en-zh (limit {limit} pairs, both_dirs={both_dirs})") + return fetch_opus100_enzh( + limit, both_dirs=both_dirs, cache_dir=cache_dir, eval_dir=eval_dir + ) + raise FileNotFoundError( + f"SFT source {source!r} is not a file; use a jsonl path or 'opus-100'" + ) + + def load_sft_rows(path: str | Path) -> list[dict]: """jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``.""" p = Path(path) diff --git a/tests/integration/test_sft_data.py b/tests/integration/test_sft_data.py index c7b1897..2748fae 100644 --- a/tests/integration/test_sft_data.py +++ b/tests/integration/test_sft_data.py @@ -1,4 +1,11 @@ -from kda.training.data import IGNORE_INDEX, collate_sft, encode_sft_row, load_sft_rows +from kda.training.data import ( + IGNORE_INDEX, + collate_sft, + encode_sft_row, + fetch_opus100_enzh, + load_sft_rows, + resolve_sft_rows, +) from kda.training.prompts import instruction_prompt @@ -42,6 +49,36 @@ def test_collate_and_jsonl(tmp_path): assert (y == IGNORE_INDEX).any() +def test_resolve_sft_rows_reads_local_jsonl(tmp_path): + path = tmp_path / "bitext.jsonl" + path.write_text( + '{"src": "你好", "tgt": "Hello", "target_lang": "en"}\n', + encoding="utf-8", + ) + rows = resolve_sft_rows(str(path)) + assert rows == [{"src": "你好", "tgt": "Hello", "target_lang": "en"}] + + +def test_fetch_opus_skips_eval_sentences(tmp_path, monkeypatch): + eval_dir = tmp_path / "eval" + eval_dir.mkdir() + (eval_dir / "zh2en.src.txt").write_text("禁止句\n", encoding="utf-8") + + class _DS: + def __iter__(self): + yield {"translation": {"en": "Hello", "zh": "你好"}} + yield {"translation": {"en": "skip", "zh": "禁止句"}} + yield {"translation": {"en": "Thanks", "zh": "谢谢"}} + + fake = type("datasets", (), {"load_dataset": staticmethod(lambda *a, **k: _DS())}) + monkeypatch.setitem(__import__("sys").modules, "datasets", fake) + rows = fetch_opus100_enzh(10, cache_dir=tmp_path / "sft", eval_dir=eval_dir) + srcs = {r["src"] for r in rows} + assert "禁止句" not in srcs + assert "你好" in srcs and "Hello" in srcs + assert "谢谢" in srcs and "Thanks" in srcs + + def test_toy_sft_file_parses(): from pathlib import Path diff --git a/train_sft.py b/train_sft.py index 5c80d33..41e7e37 100644 --- a/train_sft.py +++ b/train_sft.py @@ -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)