AF3-NA-plus / tests /test_hub_runtime.py
StarLiu714's picture
Initial release AF3-NA+ v0.0.1 preview package
a371c7e
Raw History Blame Contribute Delete
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()