Skip to content

Commit bd350aa

Browse files
authored
Merge pull request #165 from stantheman0128/fix/164-consolidate-test-split-leak
fix(sleep): keep all-test batches out of consolidate train/val
2 parents 61735e3 + 5f8e85c commit bd350aa

2 files changed

Lines changed: 96 additions & 5 deletions

File tree

‎skillopt_sleep/consolidate.py‎

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -54,12 +54,13 @@ def _norm(s: str) -> str:
5454

5555
train = [t for t in tasks if _norm(t.split) == "train"]
5656
val = [t for t in tasks if _norm(t.split) == "val"]
57-
# be robust if a split is empty: fall back so a night still does something,
58-
# but never silently use test as val.
59-
test = [t for t in tasks if _norm(t.split) == "test"]
57+
# Be robust if a split is empty: fall back so a night still does something,
58+
# but never silently use test as train or val. An all-test batch therefore
59+
# returns empty train/val (caller scores test separately; gate is a no-op).
6060
if not val:
61-
# prefer train as the gate reference over nothing; last resort all-but-test
62-
val = train or [t for t in tasks if _norm(t.split) != "test"] or tasks
61+
# Prefer train as the gate reference; otherwise any non-test tasks.
62+
# Do not fall back to the full task list (that would leak held-out test).
63+
val = train or [t for t in tasks if _norm(t.split) != "test"]
6364
if not train:
6465
train = val
6566
return train, val

‎tests/test_consolidate_split.py‎

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,90 @@
1+
"""Regression tests for consolidate._split hold-out contract.
2+
3+
test-split tasks must never enter train or val. An all-test batch must not
4+
silently fall back to using the held-out set for consolidation.
5+
"""
6+
from __future__ import annotations
7+
8+
import os
9+
import tempfile
10+
import unittest
11+
12+
from skillopt_sleep.consolidate import _split
13+
from skillopt_sleep.tasks_file import load_tasks_file, write_tasks_file
14+
from skillopt_sleep.types import TaskRecord
15+
16+
17+
def _ids(tasks):
18+
return sorted(t.id for t in tasks)
19+
20+
21+
class TestConsolidateSplit(unittest.TestCase):
22+
def test_all_test_batch_does_not_leak_into_train_or_val(self):
23+
tasks = [
24+
TaskRecord(id="t0", project="p", intent="do X", split="test"),
25+
TaskRecord(id="t1", project="p", intent="do Y", split="test"),
26+
TaskRecord(id="t2", project="p", intent="do Z", split="test"),
27+
]
28+
train, val = _split(tasks)
29+
self.assertEqual(train, [])
30+
self.assertEqual(val, [])
31+
32+
def test_all_test_via_tasks_file_path_does_not_leak(self):
33+
tasks = [
34+
TaskRecord(id="t0", project="p", intent="do X", split="test"),
35+
TaskRecord(id="t1", project="p", intent="do Y", split="test"),
36+
TaskRecord(id="t2", project="p", intent="do Z", split="test"),
37+
]
38+
with tempfile.TemporaryDirectory() as tmp:
39+
path = write_tasks_file(
40+
os.path.join(tmp, "tasks.json"),
41+
{"tasks": [t.to_dict() for t in tasks]},
42+
)
43+
loaded, _ = load_tasks_file(path)
44+
train, val = _split(loaded)
45+
self.assertEqual(_ids(train), [])
46+
self.assertEqual(_ids(val), [])
47+
self.assertEqual({t.split for t in loaded}, {"test"})
48+
49+
def test_train_only_falls_back_val_to_train(self):
50+
tasks = [
51+
TaskRecord(id="a", project="p", intent="A", split="train"),
52+
TaskRecord(id="b", project="p", intent="B", split="train"),
53+
]
54+
train, val = _split(tasks)
55+
self.assertEqual(_ids(train), ["a", "b"])
56+
self.assertEqual(_ids(val), ["a", "b"])
57+
58+
def test_train_plus_test_without_val_gates_on_train_not_test(self):
59+
tasks = [
60+
TaskRecord(id="tr", project="p", intent="train", split="train"),
61+
TaskRecord(id="te", project="p", intent="test", split="test"),
62+
]
63+
train, val = _split(tasks)
64+
self.assertEqual(_ids(train), ["tr"])
65+
self.assertEqual(_ids(val), ["tr"])
66+
self.assertNotIn("te", _ids(train) + _ids(val))
67+
68+
def test_explicit_train_val_test_keeps_partitions(self):
69+
tasks = [
70+
TaskRecord(id="tr", project="p", intent="train", split="train"),
71+
TaskRecord(id="va", project="p", intent="val", split="val"),
72+
TaskRecord(id="te", project="p", intent="test", split="test"),
73+
]
74+
train, val = _split(tasks)
75+
self.assertEqual(_ids(train), ["tr"])
76+
self.assertEqual(_ids(val), ["va"])
77+
78+
def test_legacy_holdout_name_maps_to_val(self):
79+
tasks = [
80+
TaskRecord(id="tr", project="p", intent="train", split="replay"),
81+
TaskRecord(id="va", project="p", intent="val", split="holdout"),
82+
TaskRecord(id="te", project="p", intent="test", split="test"),
83+
]
84+
train, val = _split(tasks)
85+
self.assertEqual(_ids(train), ["tr"])
86+
self.assertEqual(_ids(val), ["va"])
87+
88+
89+
if __name__ == "__main__":
90+
unittest.main()

0 commit comments

Comments
 (0)