opensysone / source /tests /test_expanded_data.py
andyshu's picture
Back up verified OpenSysOne training snapshot and pinned source
1a0a7fb verified
Raw History Blame
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()