TD3B / scoring /functions /peptiverse_binding.py
chq1155
Add PeptiVerse affinity backend
ee96220
Raw History Blame Contribute Delete
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)
@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()