File size: 6,496 Bytes
35d483e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
from __future__ import annotations

import random
import sys
import unittest
from copy import deepcopy
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))

from turn_detection.data import (  # noqa: E402
    SplitValidationError,
    assert_no_split_leakage,
    assign_leave_one_out,
    assign_splits,
    build_split_report,
    find_split_leakage,
    iter_leave_one_out_folds,
    parse_split_ratios,
)


def sample_rows(groups: int = 100) -> list[dict]:
    rows = []
    for index in range(groups):
        group_id = f"group-{index:03d}"
        group_keys = [f"record:{index}", f"audio:{index}"]
        # Every tenth component has two linked rows.
        copies = 2 if index % 10 == 0 else 1
        for copy in range(copies):
            rows.append(
                {
                    "record_id": f"row-{index}-{copy}",
                    "group_id": group_id,
                    "group_keys": group_keys,
                    "audio_sha256": f"hash-{index}",
                    "endpoint": bool(index % 2),
                    "language": "hi" if index % 3 == 0 else "en",
                    "dataset": "human" if index % 4 == 0 else "synthetic",
                }
            )
    return rows


class SplitAssignmentTests(unittest.TestCase):
    def test_deterministic_irrespective_of_input_order(self) -> None:
        rows = sample_rows()
        shuffled = deepcopy(rows)
        random.Random(7).shuffle(shuffled)
        first = assign_splits(rows, ratios={"train": 0.8, "validation": 0.2}, seed=11)
        second = assign_splits(shuffled, ratios={"train": 0.8, "validation": 0.2}, seed=11)
        first_map = {row["record_id"]: row["split"] for row in first}
        second_map = {row["record_id"]: row["split"] for row in second}
        self.assertEqual(first_map, second_map)
        assert_no_split_leakage(first)

    def test_ratios_are_close_without_splitting_groups(self) -> None:
        result = assign_splits(sample_rows(), ratios={"train": 0.8, "validation": 0.2})
        report = build_split_report(result)
        train_fraction = report["split_counts"]["train"] / len(result)
        self.assertAlmostEqual(train_fraction, 0.8, delta=0.04)
        self.assertTrue(report["leakage"]["is_valid"])

    def test_balances_fifty_fifty_endpoint_with_correlated_strata(self) -> None:
        rows = []
        for index in range(200):
            endpoint = index % 2 == 0
            rows.append(
                {
                    "record_id": f"balanced-{index}",
                    "group_id": f"balanced-group-{index}",
                    "group_keys": [f"audio:balanced-{index}"],
                    "audio_sha256": f"balanced-{index}",
                    "endpoint": endpoint,
                    # Deliberately correlate language/source with endpoint.
                    "language": "hi" if endpoint else ("en" if index % 4 else "mr"),
                    "dataset": "human" if endpoint else "synthetic",
                }
            )
        result = assign_splits(rows, ratios={"train": 0.8, "validation": 0.2}, seed=19)
        for split, expected_per_label in (("train", 80), ("validation", 20)):
            counts = {
                label: sum(row["split"] == split and row["endpoint"] is label for row in result)
                for label in (False, True)
            }
            self.assertLessEqual(abs(counts[False] - expected_per_label), 2)
            self.assertLessEqual(abs(counts[True] - expected_per_label), 2)

    def test_holdout_field_never_crosses(self) -> None:
        result = assign_splits(
            sample_rows(),
            ratios={"train": 0.5, "validation": 0.5},
            holdout_fields=("dataset",),
        )
        splits_by_dataset: dict[str, set[str]] = {}
        for row in result:
            splits_by_dataset.setdefault(row["dataset"], set()).add(row["split"])
        self.assertTrue(all(len(splits) == 1 for splits in splits_by_dataset.values()))
        report = build_split_report(result, holdout_fields=("dataset",))
        reported = {
            value
            for values in report["holdouts"]["dataset"]["by_split"].values()
            for value in values
        }
        self.assertEqual(reported, set(splits_by_dataset))
        self.assertEqual(report["holdouts"]["dataset"]["crossing_values"], {})
        self.assertEqual(report["domain_shift"]["kind"], "domain_shift_stress_test")
        self.assertIn("confounded", report["domain_shift"]["caution"])
        self.assertIn("synthetic", report["domain_shift"]["slice_counts"])

    def test_leave_one_source_out_helper_names_exact_source(self) -> None:
        rows = sample_rows(20)
        result = assign_leave_one_out(rows, field="dataset", held_out_value="human")
        self.assertTrue(
            all(row["split"] == "validation" for row in result if row["dataset"] == "human")
        )
        self.assertTrue(
            all(row["split"] == "train" for row in result if row["dataset"] == "synthetic")
        )
        report = build_split_report(result, holdout_fields=("dataset",))
        self.assertIn("human", report["holdouts"]["dataset"]["by_split"]["validation"])
        folds = list(iter_leave_one_out_folds(rows, field="dataset", values=("synthetic",)))
        self.assertEqual(folds[0][0], "synthetic")
        self.assertTrue(
            all(
                row["split"] == "validation" for row in folds[0][1] if row["dataset"] == "synthetic"
            )
        )

    def test_leakage_detector_checks_hash_even_if_group_ids_are_wrong(self) -> None:
        rows = [
            {"group_id": "a", "group_keys": ["audio:x"], "audio_sha256": "x", "split": "train"},
            {
                "group_id": "b",
                "group_keys": ["audio:x"],
                "audio_sha256": "x",
                "split": "validation",
            },
        ]
        report = find_split_leakage(rows)
        self.assertFalse(report["is_valid"])
        self.assertIn("x", report["audio_hash_crossings"])
        with self.assertRaises(SplitValidationError):
            assert_no_split_leakage(rows)

    def test_ratio_parser_normalizes_weights(self) -> None:
        ratios = parse_split_ratios(["train=8", "validation=2"])
        self.assertEqual(ratios, {"train": 0.8, "validation": 0.2})
        with self.assertRaises(ValueError):
            parse_split_ratios(["train=1", "train=1"])


if __name__ == "__main__":
    unittest.main()