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()