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