bio2nl-models / test_inference.py
marisming's picture
Publish 9 pretraining and 27 complete source-selected models
9fdd079 verified
Raw History Blame Contribute Delete
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):
@classmethod
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()