Download test_inference.py from dnagpt/bio2nl-models: direct link, hf CLI and curl.
- Browser
- Download file 7.35 kB
-
https://huggingface.co/dnagpt/bio2nl-models/resolve/main/test_inference.py
- Command line
-
hf download hf://dnagpt/bio2nl-models/test_inference.py
-
curl -L -o test_inference.py https://huggingface.co/dnagpt/bio2nl-models/resolve/main/test_inference.py
7.35 kB
| """Integrity/adversarial tests; no large models, GPU or training required.""" | |
| import copy | |
| import json | |
| from pathlib import Path | |
| import tempfile | |
| import unittest | |
| import numpy as np | |
| import torch | |
| from transformers import GPT2Config, GPT2Model | |
| import inference as api | |
| import portable_core as core | |
| HERE = Path(__file__).resolve().parent | |
| class IntegrityTests(unittest.TestCase): | |
| def setUpClass(cls): | |
| torch.set_num_threads(2) | |
| def setUp(self): | |
| self.temp = tempfile.TemporaryDirectory() | |
| self.root = Path(self.temp.name) | |
| cfg = GPT2Config(vocab_size=37, n_positions=32, n_ctx=32, n_embd=16, n_layer=1, n_head=2, | |
| n_inner=32, eos_token_id=1, pad_token_id=0) | |
| cfg._attn_implementation = "sdpa" | |
| model = core.TransferClassifier(GPT2Model(cfg), 0) | |
| self.payload = {"schema_version": 1, "kind": "source_classifier", "config": cfg.to_dict(), | |
| "metadata": {"job": {"condition": "EP", "pt_seed": 0, "ft_seed": 0, "smoke": False}, | |
| "protocol_sha256": "a" * 64, "prepared_manifest_sha256": "b" * 64, "epoch": 2}, | |
| "state_dict": core.snapshot_state(model)} | |
| self.path = self.root / "toy.pt" | |
| self.save() | |
| def tearDown(self): | |
| self.temp.cleanup() | |
| def save(self): | |
| self.payload["state_sha256"] = core.state_digest(self.payload["state_dict"]) | |
| torch.save(self.payload, self.path) | |
| self.kw = dict(expected_sha256=api.file_sha(self.path), expected_kind="source_classifier", | |
| expected_job=self.payload["metadata"]["job"], expected_state_sha256=self.payload["state_sha256"], | |
| expected_protocol_sha256="a" * 64, expected_manifest_sha256="b" * 64, expected_epoch=2, | |
| expected_config=self.payload["config"]) | |
| def test_intact_reloads(self): | |
| model, payload = core.load_checkpoint(self.path, **self.kw) | |
| self.assertEqual(core.state_digest(model.state_dict()), payload["state_sha256"]) | |
| self.assertFalse(model.training) | |
| self.assertTrue(all(not x.requires_grad for x in model.parameters())) | |
| def test_corrupt_file_rejected(self): | |
| with self.path.open("ab") as stream: | |
| stream.write(b"tampered") | |
| with self.assertRaisesRegex(ValueError, "file SHA"): | |
| core.load_checkpoint(self.path, **self.kw) | |
| def test_missing_head_rejected_even_with_new_digest(self): | |
| del self.payload["state_dict"]["score.weight"] | |
| self.save() | |
| with self.assertRaises(RuntimeError): | |
| core.load_checkpoint(self.path, **self.kw) | |
| def test_wrong_head_shape_rejected_even_with_new_digest(self): | |
| self.payload["state_dict"]["score.weight"] = torch.zeros(3, 16) | |
| self.save() | |
| with self.assertRaises(RuntimeError): | |
| core.load_checkpoint(self.path, **self.kw) | |
| def test_provenance_rejected(self): | |
| for key in ("expected_protocol_sha256", "expected_manifest_sha256", "expected_state_sha256"): | |
| changed = dict(self.kw, **{key: "c" * 64}) | |
| with self.assertRaises(ValueError): | |
| core.load_checkpoint(self.path, **changed) | |
| def test_wrong_job_rejected(self): | |
| changed = dict(self.kw, expected_job={**self.kw["expected_job"], "ft_seed": 1}) | |
| with self.assertRaisesRegex(ValueError, "job differs"): | |
| core.load_checkpoint(self.path, **changed) | |
| def test_wrong_epoch_rejected(self): | |
| with self.assertRaisesRegex(ValueError, "epoch differs"): | |
| core.load_checkpoint(self.path, **{**self.kw, "expected_epoch": 1}) | |
| def test_wrong_config_rejected(self): | |
| with self.assertRaisesRegex(ValueError, "config differs"): | |
| core.load_checkpoint(self.path, **{**self.kw, "expected_config": {"n_embd": 32}}) | |
| def test_json_config_label_keys_roundtrip(self): | |
| public_config = json.loads(json.dumps(self.payload["config"])) | |
| digest = api.canonical_sha(public_config) | |
| expected = api.config_for_core(public_config) | |
| model, payload = core.load_checkpoint(self.path, **{**self.kw, "expected_config": expected}) | |
| self.assertEqual(api.original_config_sha_after_load(payload["config"]), digest) | |
| self.assertEqual(api.canonical_sha(public_config), digest) | |
| def test_unexpected_config_runtime_value_rejected(self): | |
| changed = {**self.payload["config"], "attn_implementation": "eager"} | |
| with self.assertRaisesRegex(ValueError, "injected attention"): | |
| api.original_config_sha_after_load(changed) | |
| def test_nonfinite_rejected(self): | |
| self.payload["state_dict"]["score.weight"][0, 0] = float("nan") | |
| self.save() | |
| with self.assertRaisesRegex(ValueError, "nonfinite"): | |
| core.load_checkpoint(self.path, **self.kw) | |
| def test_path_escape_and_ambiguity_rejected(self): | |
| for relative in ("../secret", "/tmp/file", "a/../b", "a//b", "./a", "a\\b"): | |
| with self.assertRaises(ValueError): | |
| api.safe_path(self.root, relative) | |
| def test_catalog_file_tamper(self): | |
| path = HERE / "model_catalog.json" | |
| original = path.read_bytes() | |
| bad = self.root / "catalog.json" | |
| bad.write_bytes(original + b" ") | |
| with self.assertRaisesRegex(ValueError, "Catalog SHA"): | |
| api.read_catalog(bad, api.file_sha(path)) | |
| def test_wrong_parent_rejected_with_rehashed_catalog(self): | |
| data = json.loads((HERE / "model_catalog.json").read_text()) | |
| entry = next(x for x in data["models"] if x["kind"] == "source_classifier") | |
| entry["pretraining_parent"]["id"] = "ES-pt0" | |
| path = self.root / "catalog.json" | |
| path.write_text(json.dumps(data)) | |
| with self.assertRaisesRegex(ValueError, "parent provenance"): | |
| api.read_catalog(path, api.file_sha(path)) | |
| def test_tokenizer_hash_tamper(self): | |
| catalog = json.loads((HERE / "model_catalog.json").read_text()) | |
| raw = (HERE / "tokenizer.json").read_bytes() | |
| (self.root / "tokenizer.json").write_bytes(raw.replace(b"BPE", b"XYZ", 1)) | |
| with self.assertRaisesRegex(ValueError, "SHA"): | |
| api.load_tokenizer(catalog, self.root) | |
| def test_encoding_caps_and_roundtrip(self): | |
| catalog = json.loads((HERE / "model_catalog.json").read_text()) | |
| tokenizer = api.load_tokenizer(catalog, HERE) | |
| pair = {"sentence1": "hello world " * 400, "sentence2": "protein " * 400} | |
| arrays, info = api.encode_pairs(tokenizer, [pair]) | |
| a = tokenizer.encode(pair["sentence1"], add_special_tokens=False).ids[:255] | |
| b = tokenizer.encode(pair["sentence2"], add_special_tokens=False).ids[:255] | |
| self.assertEqual(arrays["input_ids"][0].tolist(), a + [2] + b + [1]) | |
| self.assertTrue(arrays["attention_mask"].all()) | |
| self.assertEqual(info[0]["truncated_endpoints"], [True, True]) | |
| def test_labels_or_empty_endpoints_rejected(self): | |
| tokenizer = api.load_tokenizer(json.loads((HERE / "model_catalog.json").read_text()), HERE) | |
| for pair in ({"sentence1": "a", "sentence2": "b", "label": 1}, {"sentence1": "", "sentence2": "b"}): | |
| with self.assertRaises(ValueError): | |
| api.encode_pairs(tokenizer, [pair]) | |
| if __name__ == "__main__": | |
| unittest.main() | |