import json import sys import torch from kda.training.data import ( chunk_ids, fetch_wiki_texts, interleave_balanced, split_heldout, ) def test_interleave_is_one_to_one_monolingual_blocks(): a = list(range(10)) b = list(range(100, 112)) out = interleave_balanced(a, b, block=4) assert out == [0, 1, 2, 3, 100, 101, 102, 103, 4, 5, 6, 7, 104, 105, 106, 107] def test_split_heldout_keeps_at_least_one_train(): chunks = torch.arange(10).view(10, 1, 1) train, held = split_heldout(chunks, frac=0.01, min_heldout=1) assert train.size(0) == 9 assert held.size(0) == 1 empty_train, empty_held = split_heldout(chunks[:1], frac=0.5) assert empty_train.size(0) == 1 assert empty_held.size(0) == 0 def test_chunk_ids_drops_tail(): ids = list(range(10)) chunks = chunk_ids(ids, batch=2, seq_len=4) assert chunks.shape == (1, 2, 4) def test_wiki_cache_roundtrip(tmp_path, monkeypatch): cache = tmp_path / "pretrain" cache.mkdir() path = cache / "wiki-zh-n2-limit3.jsonl" path.write_text( "\n".join(json.dumps({"text": f"article {i}"}) for i in range(3)) + "\n", encoding="utf-8", ) def _boom(*_a, **_k): raise AssertionError("must not hit the network") fake = type(sys)("datasets") fake.load_dataset = _boom monkeypatch.setitem(sys.modules, "datasets", fake) texts = fetch_wiki_texts(3, lang="zh", cache_dir=cache) assert texts == ["article 0", "article 1", "article 2"]