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