"""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) @staticmethod 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, ) @staticmethod 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 @staticmethod 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}) @torch.no_grad() 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] @torch.no_grad() 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()