Download tests/test_hub_runtime.py from StarLiu714/AF3-NA-plus: direct link, hf CLI and curl.
- Browser
- Download file 6.33 kB
-
https://huggingface.co/StarLiu714/AF3-NA-plus/resolve/main/tests/test_hub_runtime.py
- Command line
-
hf download hf://StarLiu714/AF3-NA-plus/tests/test_hub_runtime.py
-
curl -L -o test_hub_runtime.py https://huggingface.co/StarLiu714/AF3-NA-plus/resolve/main/tests/test_hub_runtime.py
6.33 kB
| from __future__ import annotations | |
| import json | |
| import pathlib | |
| import tempfile | |
| import unittest | |
| import torch | |
| from safetensors.torch import load_file, save_file | |
| RELEASE_ROOT = pathlib.Path(__file__).resolve().parents[1] | |
| import sys | |
| sys.path.insert(0, str(RELEASE_ROOT)) | |
| from evotemplate_na.inference import load_artifact # noqa: E402 | |
| from evotemplate_na.model import EvoTemplateConfig, EvoTemplateNA # noqa: E402 | |
| from hub_loader import HubLoadError, load_release, load_release_spec # noqa: E402 | |
| from hub_predict import normalize_input, predict # noqa: E402 | |
| def tiny_model_config() -> dict: | |
| return { | |
| "config_version": 1, | |
| "model_type": "evotemplate_na", | |
| "student": { | |
| "model_dim": 16, | |
| "num_layers": 1, | |
| "num_heads": 2, | |
| "ffn_dim": 32, | |
| "max_chains": 4, | |
| "dropout": 0.0, | |
| "rotary_base": 10000.0, | |
| "vocab_size": 11, | |
| "polymer_vocab_size": 3, | |
| }, | |
| "pair_decoder": { | |
| "token_dim": 16, | |
| "pair_rank": 4, | |
| "pair_dim": 8, | |
| "num_pair_blocks": 1, | |
| "relative_position_dim": 4, | |
| "max_relative_position": 8, | |
| "base_embedding_dim": 4, | |
| "chain_pair_embedding_dim": 2, | |
| "polymer_pair_embedding_dim": 2, | |
| "slot_embedding_dim": 4, | |
| "num_slots": 4, | |
| "pb_bins": 40, | |
| "auxiliary_distance_bins": 10, | |
| "pairing_classes": 3, | |
| "block_size": 2, | |
| "dropout": 0.0, | |
| "allow_cross_chain": True, | |
| }, | |
| } | |
| def tokenizer_config(max_chains: int) -> dict: | |
| return { | |
| "schema_version": 1, | |
| "max_chains": max_chains, | |
| "vocabulary": [ | |
| "A", "C", "G", "U", "T", "N", "MASK", "BOS", "EOS", "PAD", | |
| "CHAIN_BOUNDARY", | |
| ], | |
| "one_token_per_nucleotide": True, | |
| "explicit_chain_boundary": True, | |
| "reject_t_in_rna": True, | |
| "reject_u_in_dna": True, | |
| "silent_ut_conversion": False, | |
| } | |
| def build_tiny_release(root: pathlib.Path) -> None: | |
| model_config = tiny_model_config() | |
| tokenizer = tokenizer_config(4) | |
| torch.manual_seed(17) | |
| model = EvoTemplateNA(EvoTemplateConfig.from_dict(model_config)) | |
| state = { | |
| name: value.detach().to(dtype=torch.bfloat16).contiguous() | |
| for name, value in model.state_dict().items() | |
| } | |
| save_file(state, str(root / "model.safetensors")) | |
| config = { | |
| "schema_version": 1, | |
| "model_type": "evotemplate_na", | |
| "precision": "bfloat16", | |
| "model_config": model_config, | |
| "checkpoint_file": "model.safetensors", | |
| "quantization": None, | |
| } | |
| (root / "config.json").write_text(json.dumps(config), encoding="utf-8") | |
| (root / "tokenizer.json").write_text( | |
| json.dumps(tokenizer), encoding="utf-8" | |
| ) | |
| class HubRuntimeTest(unittest.TestCase): | |
| def test_real_release_spec_loads_slim_metadata(self) -> None: | |
| spec = load_release_spec(RELEASE_ROOT) | |
| self.assertEqual(spec.precision, "bfloat16") | |
| self.assertEqual(spec.model_config.student.max_chains, 128) | |
| def test_tiny_bf16_safetensors_round_trip_and_k4_artifact(self) -> None: | |
| with tempfile.TemporaryDirectory() as temporary: | |
| root = pathlib.Path(temporary) | |
| release = root / "release" | |
| release.mkdir() | |
| build_tiny_release(release) | |
| loaded = load_release(release, device="cpu", dtype="float32") | |
| self.assertEqual(loaded.dtype, torch.float32) | |
| request = root / "input.json" | |
| request.write_text( | |
| json.dumps( | |
| { | |
| "name": "tiny-rna-dna", | |
| "chains": [ | |
| { | |
| "chain_id": "R", | |
| "polymer_type": "RNA", | |
| "sequence": "AG", | |
| }, | |
| { | |
| "chain_id": "D", | |
| "polymer_type": "DNA", | |
| "sequence": "CT", | |
| }, | |
| ], | |
| } | |
| ), | |
| encoding="utf-8", | |
| ) | |
| output = root / "artifact" | |
| result = predict( | |
| model_dir=release, | |
| input_path=request, | |
| output_directory=output, | |
| device="cpu", | |
| dtype="float32", | |
| block_size=2, | |
| ) | |
| self.assertEqual(result.length, 4) | |
| artifact = load_artifact(output) | |
| self.assertEqual(artifact.arrays["dgram_39"].shape, (4, 4, 4, 39)) | |
| def test_strict_loader_rejects_incomplete_safetensors(self) -> None: | |
| with tempfile.TemporaryDirectory() as temporary: | |
| release = pathlib.Path(temporary) | |
| build_tiny_release(release) | |
| weights = release / "model.safetensors" | |
| state = load_file(str(weights), device="cpu") | |
| state.pop(next(iter(state))) | |
| replacement = release / "replacement.safetensors" | |
| save_file(state, str(replacement)) | |
| replacement.replace(weights) | |
| with self.assertRaisesRegex(HubLoadError, "key set differs"): | |
| load_release(release, device="cpu", dtype="float32") | |
| def test_af3_input_expands_na_chain_ids_and_ignores_other_molecules(self) -> None: | |
| payload = { | |
| "dialect": "alphafold3", | |
| "name": "af3-input", | |
| "sequences": [ | |
| {"protein": {"id": "P", "sequence": "MKT"}}, | |
| {"rna": {"id": ["R", "R.1"], "sequence": "AGCU"}}, | |
| {"dna": {"id": "D", "sequence": "ACGT"}}, | |
| {"ligand": {"id": "L", "ccdCodes": ["MG"]}}, | |
| ], | |
| } | |
| name, chains, pairs, max_length = normalize_input( | |
| payload, fallback_artifact_id="fallback" | |
| ) | |
| self.assertEqual(name, "af3-input") | |
| self.assertEqual([chain["chain_id"] for chain in chains], ["R", "R.1", "D"]) | |
| self.assertEqual(len(pairs), 6) | |
| self.assertEqual(max_length, 12) | |
| if __name__ == "__main__": | |
| unittest.main() | |