Download tests/test_peptiverse_binding.py from ChatterjeeLab/TD3B: direct link, hf CLI and curl.
- Browser
- Download file 4 kB
-
https://huggingface.co/ChatterjeeLab/TD3B/resolve/main/tests/test_peptiverse_binding.py
- Command line
-
hf download hf://ChatterjeeLab/TD3B/tests/test_peptiverse_binding.py
-
curl -L -o test_peptiverse_binding.py https://huggingface.co/ChatterjeeLab/TD3B/resolve/main/tests/test_peptiverse_binding.py
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() | |