Download scoring/functions/peptiverse_binding.py from ChatterjeeLab/TD3B: direct link, hf CLI and curl.
- Browser
- Download file 10.2 kB
-
https://huggingface.co/ChatterjeeLab/TD3B/resolve/main/scoring/functions/peptiverse_binding.py
- Command line
-
hf download hf://ChatterjeeLab/TD3B/scoring/functions/peptiverse_binding.py
-
curl -L -o peptiverse_binding.py https://huggingface.co/ChatterjeeLab/TD3B/resolve/main/scoring/functions/peptiverse_binding.py
10.2 kB
| """PeptiVerse binding-affinity adapter for TD3B. | |
| This module implements the pooled target-sequence/binder-SMILES model published | |
| in ChatterjeeLab/PeptiVerse without loading PeptiVerse's unrelated predictors. | |
| """ | |
| import logging | |
| from pathlib import Path | |
| from typing import Dict, List, Optional | |
| import torch | |
| import torch.nn as nn | |
| logger = logging.getLogger(__name__) | |
| DEFAULT_REPO_ID = "ChatterjeeLab/PeptiVerse" | |
| DEFAULT_CHECKPOINT_FILE = ( | |
| "training_classifiers/binding_affinity/" | |
| "chemberta_smiles_pooled/best_model.pt" | |
| ) | |
| DEFAULT_ESM_MODEL = "facebook/esm2_t33_650M_UR50D" | |
| DEFAULT_CHEMBERTA_MODEL = "DeepChem/ChemBERTa-77M-MLM" | |
| class PeptiVersePooledAffinityModel(nn.Module): | |
| """PeptiVerse's bidirectional cross-attention affinity head.""" | |
| def __init__( | |
| self, | |
| target_dim: int, | |
| binder_dim: int, | |
| hidden_dim: int, | |
| n_heads: int, | |
| n_layers: int, | |
| dropout: float, | |
| ) -> None: | |
| super().__init__() | |
| self.t_proj = nn.Sequential( | |
| nn.Linear(target_dim, hidden_dim), nn.LayerNorm(hidden_dim) | |
| ) | |
| self.b_proj = nn.Sequential( | |
| nn.Linear(binder_dim, hidden_dim), nn.LayerNorm(hidden_dim) | |
| ) | |
| self.layers = nn.ModuleList() | |
| for _ in range(n_layers): | |
| self.layers.append( | |
| nn.ModuleDict( | |
| { | |
| "attn_tb": nn.MultiheadAttention( | |
| hidden_dim, n_heads, dropout=dropout | |
| ), | |
| "attn_bt": nn.MultiheadAttention( | |
| hidden_dim, n_heads, dropout=dropout | |
| ), | |
| "n1t": nn.LayerNorm(hidden_dim), | |
| "n2t": nn.LayerNorm(hidden_dim), | |
| "n1b": nn.LayerNorm(hidden_dim), | |
| "n2b": nn.LayerNorm(hidden_dim), | |
| "fft": nn.Sequential( | |
| nn.Linear(hidden_dim, 4 * hidden_dim), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(4 * hidden_dim, hidden_dim), | |
| ), | |
| "ffb": nn.Sequential( | |
| nn.Linear(hidden_dim, 4 * hidden_dim), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(4 * hidden_dim, hidden_dim), | |
| ), | |
| } | |
| ) | |
| ) | |
| self.shared = nn.Sequential( | |
| nn.Linear(2 * hidden_dim, hidden_dim), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| ) | |
| self.reg = nn.Linear(hidden_dim, 1) | |
| self.cls = nn.Linear(hidden_dim, 3) | |
| def forward(self, target: torch.Tensor, binder: torch.Tensor): | |
| target = self.t_proj(target).unsqueeze(0) | |
| binder = self.b_proj(binder).unsqueeze(0) | |
| for layer in self.layers: | |
| target_attn, _ = layer["attn_tb"](target, binder, binder) | |
| target = layer["n1t"]((target + target_attn).transpose(0, 1)).transpose(0, 1) | |
| target = layer["n2t"]( | |
| (target + layer["fft"](target)).transpose(0, 1) | |
| ).transpose(0, 1) | |
| binder_attn, _ = layer["attn_bt"](binder, target, target) | |
| binder = layer["n1b"]((binder + binder_attn).transpose(0, 1)).transpose(0, 1) | |
| binder = layer["n2b"]( | |
| (binder + layer["ffb"](binder)).transpose(0, 1) | |
| ).transpose(0, 1) | |
| hidden = self.shared(torch.cat([target[0], binder[0]], dim=-1)) | |
| return self.reg(hidden).squeeze(-1), self.cls(hidden) | |
| class PeptiVerseBindingAffinity: | |
| """Score target/peptide-SMILES pairs with PeptiVerse's pK regressor.""" | |
| backend_name = "peptiverse" | |
| def __init__( | |
| self, | |
| device=None, | |
| checkpoint_path: Optional[str] = None, | |
| repo_id: str = DEFAULT_REPO_ID, | |
| revision: Optional[str] = None, | |
| cache_dir: Optional[str] = None, | |
| local_files_only: bool = False, | |
| esm_name: str = DEFAULT_ESM_MODEL, | |
| chemberta_name: str = DEFAULT_CHEMBERTA_MODEL, | |
| batch_size: int = 32, | |
| max_protein_length: int = 1022, | |
| max_smiles_length: int = 512, | |
| ) -> None: | |
| from transformers import AutoModel, AutoTokenizer, EsmModel, EsmTokenizer | |
| self.device = torch.device( | |
| "cuda" if torch.cuda.is_available() else "cpu" | |
| ) if device is None else torch.device(device) | |
| self.batch_size = max(1, int(batch_size)) | |
| self.max_protein_length = max_protein_length | |
| self.max_smiles_length = max_smiles_length | |
| resolved_checkpoint = self._resolve_checkpoint( | |
| checkpoint_path=checkpoint_path, | |
| repo_id=repo_id, | |
| revision=revision, | |
| cache_dir=cache_dir, | |
| local_files_only=local_files_only, | |
| ) | |
| checkpoint = torch.load( | |
| resolved_checkpoint, map_location=self.device, weights_only=False | |
| ) | |
| if checkpoint.get("mode") != "pooled": | |
| raise ValueError( | |
| f"Expected a pooled PeptiVerse checkpoint, got {checkpoint.get('mode')!r}" | |
| ) | |
| state_dict = checkpoint["state_dict"] | |
| params = checkpoint.get("best_params", {}) | |
| model = PeptiVersePooledAffinityModel( | |
| target_dim=int(state_dict["t_proj.0.weight"].shape[1]), | |
| binder_dim=int(state_dict["b_proj.0.weight"].shape[1]), | |
| hidden_dim=int(params.get("hidden_dim", state_dict["t_proj.0.weight"].shape[0])), | |
| n_heads=int(params.get("n_heads", 4)), | |
| n_layers=int(params.get("n_layers", self._infer_layers(state_dict))), | |
| dropout=float(params.get("dropout", 0.0)), | |
| ) | |
| model.load_state_dict(state_dict, strict=True) | |
| self.model = model.to(self.device).eval() | |
| model_kwargs = { | |
| "cache_dir": cache_dir, | |
| "local_files_only": local_files_only, | |
| } | |
| self.target_tokenizer = EsmTokenizer.from_pretrained(esm_name, **model_kwargs) | |
| self.target_encoder = EsmModel.from_pretrained( | |
| esm_name, add_pooling_layer=False, **model_kwargs | |
| ).to(self.device).eval() | |
| self.binder_tokenizer = AutoTokenizer.from_pretrained( | |
| chemberta_name, **model_kwargs | |
| ) | |
| self.binder_encoder = AutoModel.from_pretrained( | |
| chemberta_name, **model_kwargs | |
| ).to(self.device).eval() | |
| self.target_cache: Dict[str, torch.Tensor] = {} | |
| logger.info("Loaded PeptiVerse affinity checkpoint: %s", resolved_checkpoint) | |
| def _resolve_checkpoint( | |
| checkpoint_path: Optional[str], | |
| repo_id: str, | |
| revision: Optional[str], | |
| cache_dir: Optional[str], | |
| local_files_only: bool, | |
| ) -> str: | |
| if checkpoint_path is not None: | |
| path = Path(checkpoint_path).expanduser() | |
| if not path.is_file(): | |
| raise FileNotFoundError(f"PeptiVerse checkpoint not found: {path}") | |
| return str(path) | |
| from huggingface_hub import hf_hub_download | |
| return hf_hub_download( | |
| repo_id=repo_id, | |
| filename=DEFAULT_CHECKPOINT_FILE, | |
| revision=revision, | |
| cache_dir=cache_dir, | |
| local_files_only=local_files_only, | |
| ) | |
| def _infer_layers(state_dict) -> int: | |
| layer_ids = { | |
| int(key.split(".")[1]) | |
| for key in state_dict | |
| if key.startswith("layers.") | |
| } | |
| return max(layer_ids) + 1 | |
| def _special_token_ids(tokenizer) -> List[int]: | |
| values = [ | |
| getattr(tokenizer, f"{name}_token_id", None) | |
| for name in ("pad", "cls", "sep", "bos", "eos", "mask") | |
| ] | |
| return sorted({int(value) for value in values if value is not None}) | |
| def _pool(self, texts, tokenizer, encoder, max_length: int) -> torch.Tensor: | |
| tokens = tokenizer( | |
| list(texts), | |
| return_tensors="pt", | |
| padding=True, | |
| truncation=True, | |
| max_length=max_length, | |
| ) | |
| tokens = {name: value.to(self.device) for name, value in tokens.items()} | |
| attention_mask = tokens.get( | |
| "attention_mask", torch.ones_like(tokens["input_ids"]) | |
| ).bool() | |
| valid_mask = attention_mask | |
| special_ids = self._special_token_ids(tokenizer) | |
| if special_ids: | |
| special = torch.tensor(special_ids, device=self.device) | |
| valid_mask = valid_mask & ~torch.isin(tokens["input_ids"], special) | |
| hidden = encoder(**tokens).last_hidden_state | |
| weights = valid_mask.unsqueeze(-1).to(hidden.dtype) | |
| return (hidden * weights).sum(dim=1) / weights.sum(dim=1).clamp_min(1.0) | |
| def get_protein_embedding(self, prot_seq: str) -> torch.Tensor: | |
| prot_seq = prot_seq.strip() | |
| if prot_seq not in self.target_cache: | |
| self.target_cache[prot_seq] = self._pool( | |
| [prot_seq], | |
| self.target_tokenizer, | |
| self.target_encoder, | |
| self.max_protein_length, | |
| ) | |
| return self.target_cache[prot_seq] | |
| def forward(self, input_seqs, prot_seq: str): | |
| input_seqs = list(input_seqs) | |
| if not input_seqs: | |
| return [] | |
| target = self.get_protein_embedding(prot_seq) | |
| scores = [] | |
| for start in range(0, len(input_seqs), self.batch_size): | |
| batch = input_seqs[start:start + self.batch_size] | |
| binder = self._pool( | |
| batch, | |
| self.binder_tokenizer, | |
| self.binder_encoder, | |
| self.max_smiles_length, | |
| ) | |
| affinity, _ = self.model(target.expand(len(batch), -1), binder) | |
| scores.extend(affinity.detach().cpu().tolist()) | |
| return scores | |
| def __call__(self, input_seqs, prot_seq: str): | |
| return self.forward(input_seqs, prot_seq) | |
| def clear_cache(self) -> None: | |
| self.target_cache.clear() | |