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