TD3B / tests /test_peptiverse_binding.py
chq1155
Add PeptiVerse affinity backend
ee96220
Raw History Blame Contribute Delete
4 kB
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from scoring.functions.peptiverse_binding import (
PeptiVerseBindingAffinity,
PeptiVersePooledAffinityModel,
)
class DummyTokenizer:
pad_token_id = 0
cls_token_id = 1
eos_token_id = 2
sep_token_id = None
bos_token_id = None
mask_token_id = None
def __call__(self, texts, **kwargs):
del texts, kwargs
return {
"input_ids": torch.tensor([[1, 3, 4, 2, 0]]),
"attention_mask": torch.tensor([[1, 1, 1, 1, 0]]),
}
class DummyEncoder(torch.nn.Module):
def forward(self, input_ids, attention_mask):
del attention_mask
hidden = input_ids.float().unsqueeze(-1).repeat(1, 1, 2)
return SimpleNamespace(last_hidden_state=hidden)
class DummyAffinityHead(torch.nn.Module):
def forward(self, target, binder):
del target
return binder.sum(dim=-1), torch.zeros(len(binder), 3)
class PeptiVerseBindingTests(unittest.TestCase):
def test_affinity_head_shapes(self):
model = PeptiVersePooledAffinityModel(
target_dim=8,
binder_dim=6,
hidden_dim=12,
n_heads=3,
n_layers=2,
dropout=0.0,
).eval()
affinity, classes = model(torch.randn(4, 8), torch.randn(4, 6))
self.assertEqual(tuple(affinity.shape), (4,))
self.assertEqual(tuple(classes.shape), (4, 3))
def test_pool_excludes_special_tokens(self):
predictor = object.__new__(PeptiVerseBindingAffinity)
predictor.device = torch.device("cpu")
pooled = predictor._pool(
["unused"], DummyTokenizer(), DummyEncoder(), max_length=8
)
expected = torch.tensor([[3.5, 3.5]])
self.assertTrue(torch.equal(pooled, expected))
def test_factory_keeps_original_as_default(self):
from scoring.functions import binding
original = object()
with patch.object(
binding, "MultiTargetBindingAffinity", return_value=original
) as constructor:
result = binding.create_multi_target_affinity_predictor(
tokenizer=object(), base_path="/tmp", device="cpu"
)
self.assertIs(result, original)
constructor.assert_called_once()
def test_factory_selects_peptiverse(self):
from scoring.functions import binding
peptiverse = object()
with patch.object(
binding, "PeptiVerseBindingAffinity", return_value=peptiverse
) as constructor:
result = binding.create_multi_target_affinity_predictor(
backend="peptiverse",
device="cpu",
peptiverse_checkpoint="model.pt",
)
self.assertIs(result, peptiverse)
constructor.assert_called_once_with(
device="cpu",
checkpoint_path="model.pt",
repo_id="ChatterjeeLab/PeptiVerse",
revision=None,
cache_dir=None,
local_files_only=False,
batch_size=32,
)
def test_forward_batches_binder_smiles(self):
predictor = object.__new__(PeptiVerseBindingAffinity)
predictor.batch_size = 2
predictor.binder_tokenizer = object()
predictor.binder_encoder = object()
predictor.max_smiles_length = 16
predictor.model = DummyAffinityHead()
predictor.get_protein_embedding = lambda _: torch.zeros(1, 3)
batch_sizes = []
def fake_pool(texts, tokenizer, encoder, max_length):
del tokenizer, encoder, max_length
batch_sizes.append(len(texts))
return torch.ones(len(texts), 2)
predictor._pool = fake_pool
scores = predictor.forward(["a", "b", "c", "d", "e"], "TARGET")
self.assertEqual(batch_sizes, [2, 2, 1])
self.assertEqual(scores, [2.0] * 5)
if __name__ == "__main__":
unittest.main()