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:
dela
2026-08-25 21:39:58 +08:00
parent 53d0f4b17a
commit e7185cbf49
3 changed files with 172 additions and 8 deletions
+108
View File
@@ -89,6 +89,17 @@ def pretrain_dir() -> Path:
return Path("data/pretrain") 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]: def _wiki_files(lang: str, n_shards: int) -> list[str]:
if lang not in WIKI_SHARD_TOTAL: if lang not in WIKI_SHARD_TOTAL:
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en") raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
@@ -287,6 +298,103 @@ def encode_sft_row(
return ids, labels 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]: def load_sft_rows(path: str | Path) -> list[dict]:
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``.""" """jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
p = Path(path) p = Path(path)
+38 -1
View File
@@ -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 from kda.training.prompts import instruction_prompt
@@ -42,6 +49,36 @@ def test_collate_and_jsonl(tmp_path):
assert (y == IGNORE_INDEX).any() 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(): def test_toy_sft_file_parses():
from pathlib import Path from pathlib import Path
+26 -7
View File
@@ -1,9 +1,9 @@
"""Instruction SFT for zh↔en translation. Prompt template matches eval_mt. """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_wiki.pt --data data/sft/toy.jsonl
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data data/sft/opus.jsonl \\ uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\
--seq-len 512 --batch 4 --lr 5e-5 --epochs 2 --limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2
""" """
from __future__ import annotations from __future__ import annotations
@@ -18,7 +18,7 @@ from kda.layers.latent_moe import moe_router_losses
from kda.training.data import ( from kda.training.data import (
IGNORE_INDEX, IGNORE_INDEX,
iter_sft_batches, iter_sft_batches,
load_sft_rows, resolve_sft_rows,
load_tokenizer, load_tokenizer,
) )
from kda.training.eval_mt import evaluate_pairs from kda.training.eval_mt import evaluate_pairs
@@ -72,7 +72,22 @@ def _read_lines(path: str) -> list[str]:
def main() -> None: def main() -> None:
p = argparse.ArgumentParser(description=__doc__) p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--ckpt", required=True) 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("--out", default="ckpts/k3_sft.pt")
p.add_argument("--tokenizer", default=None) p.add_argument("--tokenizer", default=None)
p.add_argument("--batch", type=int, default=4) p.add_argument("--batch", type=int, default=4)
@@ -103,9 +118,13 @@ def main() -> None:
if not tok_src: if not tok_src:
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint") raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
tok = load_tokenizer(tok_src) 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: 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__}") 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) steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1)