blizzarman/polyglot-tutor
0
1import pytest2 3from tutor.ml.cefr.preprocessing import passages_from_record4from tutor.ml.cefr.splitting import SPLITS, assign_splits5 6 7def _strata(n_docs: int, stratum: str, prefix: str = "doc") -> dict[str, str]:8 return {f"{prefix}:{i}": stratum for i in range(n_docs)}9 10 11def test_ratios_on_a_large_stratum() -> None:12 assignment = assign_splits(_strata(100, "corpus|B1"), ratios=(0.8, 0.1, 0.1), seed=13)13 counts = {split: sum(1 for value in assignment.values() if value == split) for split in SPLITS}14 assert counts == {"train": 80, "val": 10, "test": 10}15 16 17def test_small_strata_go_to_train_only() -> None:18 assignment = assign_splits(_strata(5, "corpus|A1"), seed=13)19 assert set(assignment.values()) == {"train"}20 21 22def test_deterministic_and_insertion_order_independent() -> None:23 strata = {**_strata(50, "x|B1", "a"), **_strata(50, "y|B2", "b")}24 reversed_strata = dict(reversed(list(strata.items())))25 assert assign_splits(strata, seed=13) == assign_splits(reversed_strata, seed=13)26 assert assign_splits(strata, seed=13) != assign_splits(strata, seed=14)27 28 29def test_arm_comparability_en_assignment_unaffected_by_other_corpora() -> None:30 """The guarantee behind ADR 0003 arms 1 vs 2: adding multilingual corpora31 must not move a single English document between splits."""32 en_only = _strata(40, "cambridge_exams_en|B2", "cambridge_exams_en")33 multilingual = {34 **en_only,35 **_strata(300, "elg_cefr_nl|B1", "elg_cefr_nl"),36 **_strata(200, "readme_fr|A2", "readme_fr"),37 }38 assignment_en = assign_splits(en_only, seed=13)39 assignment_multi = assign_splits(multilingual, seed=13)40 for doc_id, split in assignment_en.items():41 assert assignment_multi[doc_id] == split42 43 44def test_bad_ratios_raise() -> None:45 with pytest.raises(ValueError, match="sum to 1"):46 assign_splits(_strata(10, "s"), ratios=(0.5, 0.2, 0.2))47 48 49def test_no_chunk_leakage_by_construction() -> None:50 """Chunks inherit their document's split: a doc_id can never straddle splits."""51 long_text = ". ".join(" ".join(f"w{i}" for i in range(11)) + " end" for _ in range(60)) + "."52 passages = []53 for doc_index in range(30):54 passages += passages_from_record(55 text=long_text,56 level_raw="B2",57 lang="en",58 corpus="cambridge_exams_en",59 doc_id=f"cambridge_exams_en:{doc_index}",60 source_format="document-level",61 )62 assert len(passages) > 30 # documents really did produce multiple chunks63 64 doc_strata = {p.doc_id: f"{p.corpus}|{p.level}" for p in passages}65 assignment = assign_splits(doc_strata, seed=13)66 parts = {67 split: {p.doc_id for p in passages if assignment[p.doc_id] == split} for split in SPLITS68 }69 assert parts["train"] & parts["val"] == set()70 assert parts["train"] & parts["test"] == set()71 assert parts["val"] & parts["test"] == set()72 