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