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).
90 lines
2.9 KiB
Python
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"}
|