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")
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)