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