Download source/tests/test_expanded_data.py from andyshu/opensysone: direct link, hf CLI and curl.
- Browser
- Download file 7.27 kB
-
https://huggingface.co/andyshu/opensysone/resolve/294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_expanded_data.py
- Command line
-
hf download hf://andyshu/opensysone@294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_expanded_data.py
-
curl -L -o test_expanded_data.py https://huggingface.co/andyshu/opensysone/resolve/294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_expanded_data.py
7.27 kB
| 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() | |