Files
K3/tests/integration/test_sft_data.py
T
dela e7185cbf49 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).
2026-08-25 21:39:58 +08:00

90 lines
2.9 KiB
Python

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
class _Tok:
vocab_size = 32
def encode(self, text: str) -> list[int]:
return [min((ord(c) % 30) + 1, 31) for c in text[:12]] or [1]
def decode(self, ids: list[int]) -> str:
return "x" * len(ids)
def test_instruction_matches_eval_template():
assert instruction_prompt("你好", "en") == "Translate to English:\n你好"
assert instruction_prompt("Hello", "zh") == "Translate to Chinese:\nHello"
def test_prompt_tokens_are_ignored():
tok = _Tok()
src, tgt = "ab", "cd"
ids, labels = encode_sft_row(tok, src, tgt, "en", max_len=64)
prompt_n = len(tok.encode(instruction_prompt(src, "en")))
assert labels[:prompt_n] == [IGNORE_INDEX] * prompt_n
assert all(v != IGNORE_INDEX for v in labels[prompt_n:])
assert ids[prompt_n:] == tok.encode(tgt)
def test_collate_and_jsonl(tmp_path):
path = tmp_path / "tiny.jsonl"
path.write_text(
'{"src": "a", "tgt": "b", "target_lang": "en"}\n'
'{"src": "c", "tgt": "d", "target_lang": "zh"}\n',
encoding="utf-8",
)
rows = load_sft_rows(path)
assert len(rows) == 2
x, y = collate_sft(rows, _Tok(), max_len=32)
assert x.shape == y.shape
assert x.size(0) == 2
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
path = Path(__file__).resolve().parents[2] / "data" / "sft" / "toy.jsonl"
rows = load_sft_rows(path)
assert len(rows) >= 20
langs = {r["target_lang"] for r in rows}
assert langs == {"en", "zh"}