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:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+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