opensysone / source /tests /test_data_transition.py
andyshu's picture
Back up verified OpenSysOne training snapshot and pinned source
1a0a7fb verified
Raw History Blame
3.96 kB
import copy
import json
from pathlib import Path
import tempfile
import unittest
from data_transition import (SPLITS, PROTECTED_SPLITS, data_signature,
sha256, verify_train_data_transition)
class DataTransitionTests(unittest.TestCase):
def setUp(self):
self.directory = tempfile.TemporaryDirectory()
self.addCleanup(self.directory.cleanup)
root = Path(self.directory.name)
self.old, self.new = root / "old", root / "new"
self.provenance = {"revision": "pinned-model"}
for path in (self.old, self.new):
path.mkdir()
for split in SPLITS:
(path / f"{split}.jsonl").write_text('{"id":"original"}\n')
(self.new / "train.jsonl").write_text('{"id":"original"}\n{"id":"new"}\n')
self.parent = {"split_sha256": self.hashes(self.old)}
(self.old / "manifest.json").write_text(json.dumps(self.parent))
self.current = {"split_sha256": self.hashes(self.new), "base_dataset": {
"path": str(self.old), "manifest_sha256": sha256(self.old / "manifest.json"),
"split_sha256": self.parent["split_sha256"]}}
self.write_current()
self.saved = {"config": {"dataset": str(self.old), "max_tokens": 512},
"model_provenance": self.provenance,
"data_signature": data_signature(self.parent, self.provenance, 512)}
self.config = {"dataset": str(self.new), "max_tokens": 512}
def hashes(self, path):
return {split: sha256(path / f"{split}.jsonl") for split in SPLITS}
def write_current(self):
(self.new / "manifest.json").write_text(json.dumps(self.current))
def verify(self):
return verify_train_data_transition(self.saved, self.config,
data_signature(self.current, self.provenance, 512))
def test_expansion_records_both_signatures_and_protected_hashes(self):
untouched = copy.deepcopy(self.saved)
result = self.verify()
self.assertEqual(self.saved, untouched)
self.assertEqual(result["parent_data_signature"], self.saved["data_signature"])
self.assertNotEqual(result["parent_data_signature"], result["data_signature"])
self.assertEqual(set(result["protected_split_sha256"]), set(PROTECTED_SPLITS))
def test_reserved_changes_rejected_even_with_valid_updated_hash(self):
for split in PROTECTED_SPLITS:
with self.subTest(split=split):
original = (self.new / f"{split}.jsonl").read_bytes()
(self.new / f"{split}.jsonl").write_text('changed')
self.current["split_sha256"] = self.hashes(self.new)
self.write_current()
with self.assertRaisesRegex(ValueError, "reserved"):
self.verify()
(self.new / f"{split}.jsonl").write_bytes(original)
def test_actual_bytes_missing_splits_and_false_lineage_rejected(self):
(self.new / "train.jsonl").write_text('corrupted')
with self.assertRaisesRegex(ValueError, "checksum"):
self.verify()
self.current["split_sha256"] = self.hashes(self.new)
self.current["base_dataset"]["manifest_sha256"] = "wrong"
self.write_current()
with self.assertRaisesRegex(ValueError, "lineage"):
self.verify()
del self.current["split_sha256"]["holdout"]
self.write_current()
with self.assertRaisesRegex(ValueError, "five registered"):
self.verify()
def test_wrong_parent_signature_or_token_policy_rejected(self):
self.saved["data_signature"] = "incorrect"
with self.assertRaisesRegex(ValueError, "Parent data signature"):
self.verify()
self.config["max_tokens"] = 1024
with self.assertRaisesRegex(ValueError, "max_tokens"):
self.verify()
if __name__ == "__main__":
unittest.main()