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