File size: 7,270 Bytes
1a0a7fb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import copy
import json
from pathlib import Path
import tempfile
import unittest

from scripts import prepare_expanded_data as expanded


def raw_row(family, index, state=None, group=None):
    state = state or f"{family} distinct training situation {index}"
    if family == "hellaswag":
        return {"ind": index, "ctx": state, "source_id": group or f"video-{index}",
                "endings": [f"ending {i}" for i in range(4)], "label": "2"}
    if family == "piqa":
        return {"goal": state, "sol1": "first solution", "sol2": "second solution", "label": 1}
    return {"id": str(index), "question": state,
            "choices": {"label": list("ABCDE"), "text": [f"answer {i}" for i in range(5)]},
            "answerKey": "D"}


def base_dataset(path):
    path.mkdir()
    for split in expanded.SPLITS:
        row = {"id": "original:" + split, "group": "old-group:" + split,
               "family": "social" if split == "holdout" else "original",
               "source_split": split, "state": "original state " + split,
               "question": "Choose", "choices": ["one", "two"], "target": 0}
        # Deliberately use noncanonical serialization to prove byte preservation.
        (path / f"{split}.jsonl").write_text(json.dumps(row, separators=(",", ":")) + "\n")
    manifest = {"version": "public-decisions-v1", "sources": {"original": {"revision": "old-pin"}},
                "split_sha256": {split: expanded.sha256(path / f"{split}.jsonl") for split in expanded.SPLITS},
                "split_family_counts": {split: {"social" if split == "holdout" else "original": 1}
                                        for split in expanded.SPLITS}}
    (path / "manifest.json").write_text(json.dumps(manifest))


class ExpandedDataTests(unittest.TestCase):
    def setUp(self):
        self.temporary = tempfile.TemporaryDirectory()
        self.addCleanup(self.temporary.cleanup)
        self.root = Path(self.temporary.name)
        self.base = self.root / "base"
        base_dataset(self.base)
        self.train = {family: [raw_row(family, i) for i in range(8)] for family in expanded.SOURCES}
        self.reserved = {family: [] for family in expanded.SOURCES}

    def build(self, name="expanded", **kwargs):
        return expanded.prepare_dataset(self.base, self.root / name, self.train, self.reserved,
                                        copy.deepcopy(expanded.SOURCES), max_new_per_family=3,
                                        diagnostic_per_family=1, **kwargs)

    def test_answer_mapping_survives_deterministic_candidate_permutation(self):
        expected = {"hellaswag": "ending 2", "piqa": "second solution", "commonsenseqa": "answer 3"}
        for family, answer in expected.items():
            row = expanded.convert_row(family, 9, raw_row(family, 9))
            self.assertEqual(row["choices"][row["target"]], answer)
            self.assertEqual(row, expanded.convert_row(family, 9, raw_row(family, 9)))
            self.assertEqual(row["source_split"], "train")
            self.assertEqual(row["source_revision"], expanded.SOURCES[family]["revision"])

    def test_immutable_replay_protected_files_and_separate_diagnostics(self):
        original_bytes = {p.name: p.read_bytes() for p in self.base.iterdir()}
        first = self.build()
        second = self.build("again")
        self.assertEqual(first["split_sha256"], second["split_sha256"])
        self.assertEqual(first["diagnostics"]["sha256"], second["diagnostics"]["sha256"])
        output = self.root / "expanded"
        self.assertEqual(set(first["split_sha256"]), set(expanded.SPLITS))
        for split in expanded.PROTECTED:
            self.assertEqual((output / f"{split}.jsonl").read_bytes(), original_bytes[f"{split}.jsonl"])
        self.assertTrue((output / "train.jsonl").read_bytes().startswith(original_bytes["train.jsonl"]))
        self.assertEqual({p.name: p.read_bytes() for p in self.base.iterdir()}, original_bytes)
        rows = list(expanded.read_rows(output / "train.jsonl"))
        diagnostics = list(expanded.read_rows(output / first["diagnostics"]["path"]))
        self.assertEqual(len(rows), 10)
        self.assertEqual(len(diagnostics), 3)
        self.assertFalse(first["diagnostics"]["selection_eligible"])
        self.assertFalse({r["group"] for r in rows} & {r["group"] for r in diagnostics})
        self.assertFalse(any(r["family"] == "social" for r in rows))

    def test_group_and_normalized_text_exclusions_precede_diagnostics_and_cap(self):
        # Same upstream video, different words: group exclusion must still win.
        self.reserved["hellaswag"] = [raw_row("hellaswag", 100, "different context", "video-0")]
        # Cross-family frozen-evaluation identity is detected despite Unicode/spacing.
        self.train["piqa"].append(raw_row("piqa", 50, "ORIGINAL   state TEST"))
        self.train["piqa"].append(raw_row("piqa", 51, "PIQA DISTINCT TRAINING SITUATION 0"))
        # Official validation/test labels/choices are neither needed nor consulted.
        self.reserved["commonsenseqa"] = [{"question": "commonsenseqa distinct training situation 0"}]
        manifest = self.build()
        self.assertEqual(manifest["audit"]["hellaswag"]["removed"]["official_evaluation_overlap"], 1)
        self.assertEqual(manifest["audit"]["commonsenseqa"]["removed"]["official_evaluation_overlap"], 1)
        removed = manifest["audit"]["piqa"]["removed"]
        self.assertEqual(removed["original_evaluation_overlap"], 1)
        self.assertEqual(removed["duplicate_training_identity_or_state"], 1)
        all_new = list(expanded.read_rows(self.root / "expanded/train.jsonl"))
        all_new += list(expanded.read_rows(self.root / "expanded/diagnostics/new_sources.jsonl"))
        self.assertFalse(any(r["family"] == "hellaswag" and r["source_record_id"] == "0" for r in all_new))

    def test_complete_diagnostic_groups_are_removed_from_training(self):
        for index in range(8, 20):
            self.train["hellaswag"].append(raw_row("hellaswag", index, group=f"video-{index % 8}"))
        manifest = self.build()
        diagnostic = next(r for r in expanded.read_rows(self.root / "expanded/diagnostics/new_sources.jsonl")
                          if r["family"] == "hellaswag")
        training = list(expanded.read_rows(self.root / "expanded/train.jsonl"))
        self.assertFalse(any(r["group"] == diagnostic["group"] for r in training))
        self.assertGreater(manifest["audit"]["hellaswag"]["diagnostic_group_rows_excluded_from_training"], 1)

    def test_rejects_changed_base_and_existing_output(self):
        self.build()
        with self.assertRaises(FileExistsError):
            self.build()
        with (self.base / "validation.jsonl").open("a") as handle:
            handle.write("\n")
        with self.assertRaisesRegex(ValueError, "Base dataset hash mismatch"):
            self.build("changed")

    def test_invalid_candidate_sets_never_reach_training(self):
        bad = raw_row("piqa", 20)
        bad["sol2"] = " FIRST   solution "
        self.train["piqa"].append(bad)
        manifest = self.build()
        self.assertEqual(manifest["audit"]["piqa"]["removed"]["invalid_source_row"], 1)


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