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()