Spaces:
Sleeping
Sleeping
| """bioai.models.caduceus_adapter -- best-effort Caduceus adapter with CNN fallback. | |
| Caduceus (Schiff et al. 2024, https://arxiv.org/abs/2403.09235) is a long-context | |
| Mamba-2 model for DNA that respects reverse-complement symmetry. The | |
| ``kuleshov-group/caduceus-ph-1`` checkpoint on HuggingFace expects integer | |
| **token** input (BPE tokenizer over A/C/G/T), not one-hot. | |
| This adapter tries to load Caduceus from HuggingFace; if anything fails | |
| (network error, OOM, missing transformers, missing trust_remote_code deps, | |
| etc.) it transparently falls back to :class:`bioai.models.sirna_cnn.SiRNACNN`. | |
| The demo video prints which backend is active so judges can see at a glance. | |
| The forward signature matches ``SiRNACNN.forward`` -- ``(B, 4, 21)`` one-hot | |
| in, ``(efficacy_pred, safety_pred)`` out -- so callers can swap the two | |
| without touching their code. | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| from typing import Tuple | |
| import torch | |
| import torch.nn as nn | |
| from .sirna_cnn import SiRNACNN, resolve_device | |
| logger = logging.getLogger(__name__) | |
| # HuggingFace id for Caduceus-Ph-1 (the 1.3M-param pretraining checkpoint). | |
| CADUCEUS_MODEL_ID = "kuleshov-group/caduceus-ph-1" | |
| class CaduceusAdapter(nn.Module): | |
| """Caduceus-backed siRNA predictor with automatic CNN fallback. | |
| Parameters | |
| ---------- | |
| seq_len: | |
| Expected siRNA length (default 21). | |
| num_safety_species: | |
| Number of safety-panel species (default 6). | |
| device: | |
| ``'auto' | 'cpu' | 'cuda'``. | |
| force_backend: | |
| ``None`` (try Caduceus then fall back) or ``'cnn'`` (skip the | |
| Caduceus attempt, useful in CI / low-RAM environments). | |
| """ | |
| def __init__( | |
| self, | |
| seq_len: int = 21, | |
| num_safety_species: int = 6, | |
| device: str = "auto", | |
| force_backend: str | None = None, | |
| ): | |
| super().__init__() | |
| self.seq_len = seq_len | |
| self.num_safety_species = num_safety_species | |
| self.device = resolve_device(device) | |
| self.backend: str = "cnn" # set during _try_load_caduceus | |
| # Small heads that map the Caduceus pooled embedding -> predictions. | |
| # Initialised here so they live on the right device even if Caduceus | |
| # later fails (we still need them for the CNN fallback). | |
| self.caduceus_embed_dim: int = 256 | |
| self.efficacy_head = nn.Linear(self.caduceus_embed_dim, 1).to(self.device) | |
| self.safety_head = nn.Linear(self.caduceus_embed_dim, num_safety_species).to(self.device) | |
| self._caduceus = None | |
| self._cnn_fallback: SiRNACNN | None = None | |
| if force_backend == "cnn": | |
| self.backend = "cnn" | |
| print("[CaduceusAdapter] force_backend='cnn' -> using SiRNACNN") | |
| else: | |
| self._try_load_caduceus() | |
| if self.backend == "cnn": | |
| self._init_cnn_fallback() | |
| # ------------------------------------------------------------------ # | |
| def _try_load_caduceus(self) -> None: | |
| """Attempt to download + load Caduceus from HuggingFace. | |
| Any exception (ImportError, network, OOM, model-config issue) drops | |
| us back to ``self.backend = 'cnn'``. | |
| """ | |
| try: | |
| from transformers import AutoModel, AutoTokenizer # type: ignore | |
| except Exception as exc: # pragma: no cover - depends on env | |
| print(f"[CaduceusAdapter] transformers unavailable ({exc!r}); falling back to SiRNACNN") | |
| self.backend = "cnn" | |
| return | |
| try: | |
| print(f"[CaduceusAdapter] Attempting to load {CADUCEUS_MODEL_ID} from HuggingFace...") | |
| tokenizer = AutoTokenizer.from_pretrained(CADUCEUS_MODEL_ID, trust_remote_code=True) | |
| model = AutoModel.from_pretrained( | |
| CADUCEUS_MODEL_ID, trust_remote_code=True, add_pooling_layer=False | |
| ) | |
| model.eval() | |
| for p in model.parameters(): | |
| p.requires_grad_(False) | |
| # Discover the real embedding dim from the config so our heads | |
| # match (Caduceus-Ph-1 is 256, but later checkpoints may differ). | |
| cfg_dim = getattr(getattr(model, "config", None), "d_model", None) | |
| if cfg_dim is not None and cfg_dim != self.caduceus_embed_dim: | |
| self.caduceus_embed_dim = int(cfg_dim) | |
| self.efficacy_head = nn.Linear(self.caduceus_embed_dim, 1).to(self.device) | |
| self.safety_head = nn.Linear(self.caduceus_embed_dim, self.num_safety_species).to(self.device) | |
| model.to(self.device) | |
| self._caduceus = model | |
| self._tokenizer = tokenizer | |
| self.backend = "caduceus" | |
| print(f"[CaduceusAdapter] Caduceus loaded (embed_dim={self.caduceus_embed_dim}). Backend=caduceus") | |
| except Exception as exc: | |
| print(f"[CaduceusAdapter] Caduceus load failed ({exc!r}); falling back to SiRNACNN") | |
| self.backend = "cnn" | |
| # ------------------------------------------------------------------ # | |
| def _init_cnn_fallback(self) -> None: | |
| self._cnn_fallback = SiRNACNN( | |
| seq_len=self.seq_len, | |
| num_safety_species=self.num_safety_species, | |
| ).to(self.device) | |
| print("[CaduceusAdapter] SiRNACNN fallback initialised. Backend=cnn") | |
| # ------------------------------------------------------------------ # | |
| def _onehot_to_tokens(x: torch.Tensor) -> torch.Tensor: | |
| """``(B, 4, L)`` one-hot -> ``(B, L)`` integer tokens (argmax).""" | |
| return x.argmax(dim=1).long() | |
| def _tokens_to_caduceus_ids(self, tokens: torch.Tensor) -> torch.Tensor: | |
| """Convert (B, L) int tokens to Caduceus input_ids. | |
| Caduceus-Ph-1 uses a simple BPE where A/C/G/T map to specific ids. | |
| Most checkpoints add ``[CLS]``/``[SEP]`` automatically via the | |
| tokenizer; we let the tokenizer handle it. If the tokenizer is not | |
| callable in the expected way, we fall back to A=5,C=6,G=7,T=8 (the | |
| raw Caduceus BPE ids observed in the published config) and prepend | |
| CLS=1 / append SEP=2. | |
| """ | |
| import torch as _t | |
| # Convert tokens back to a DNA string, then re-tokenise. | |
| idx_to_base = {0: "A", 1: "C", 2: "G", 3: "T"} | |
| seqs = [ | |
| "".join(idx_to_base.get(int(t), "A") for t in row) | |
| for row in tokens.cpu().numpy() | |
| ] | |
| try: | |
| enc = self._tokenizer( | |
| seqs, return_tensors="pt", padding=True, truncation=True, | |
| ) | |
| return enc["input_ids"].to(self.device) | |
| except Exception as exc: | |
| print(f"[CaduceusAdapter] tokenizer call failed ({exc!r}); using raw BPE ids") | |
| # Raw Caduceus BPE ids: A=5, C=6, G=7, T=8 (per the published config). | |
| raw_map = {0: 5, 1: 6, 2: 7, 3: 8} | |
| ids = tokens.clone().cpu() | |
| for k, v in raw_map.items(): | |
| ids[ids == k] = v | |
| cls_col = _t.full((ids.size(0), 1), 1, dtype=ids.dtype) | |
| sep_col = _t.full((ids.size(0), 1), 2, dtype=ids.dtype) | |
| return _t.cat([cls_col, ids, sep_col], dim=1).to(self.device) | |
| # ------------------------------------------------------------------ # | |
| def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: | |
| """Same signature as :class:`SiRNACNN`: ``(B, 4, seq_len)`` one-hot in, | |
| ``(efficacy_pred, safety_pred)`` out. | |
| """ | |
| x = x.to(self.device) | |
| if self.backend == "cnn" or self._cnn_fallback is not None and self._caduceus is None: | |
| assert self._cnn_fallback is not None | |
| return self._cnn_fallback.forward(x) | |
| # Caduceus path | |
| try: | |
| tokens = self._onehot_to_tokens(x) | |
| input_ids = self._tokens_to_caduceus_ids(tokens) | |
| with torch.no_grad(): | |
| outputs = self._caduceus(input_ids) | |
| # Caduceus returns last_hidden_state (B, L, D). Mean-pool over L. | |
| hidden = outputs.last_hidden_state if hasattr(outputs, "last_hidden_state") else outputs[0] | |
| mask = (input_ids != 0).float().unsqueeze(-1) # treat pad=0 | |
| pooled = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1.0) | |
| eff = torch.sigmoid(self.efficacy_head(pooled)) | |
| safe = torch.sigmoid(self.safety_head(pooled)) | |
| return eff, safe | |
| except Exception as exc: | |
| # Never crash the pipeline -- fall through to CNN. | |
| print(f"[CaduceusAdapter] forward failed ({exc!r}); using CNN fallback for this batch") | |
| if self._cnn_fallback is None: | |
| self._init_cnn_fallback() | |
| return self._cnn_fallback.forward(x) | |
| # ------------------------------------------------------------------ # | |
| # Passthroughs so callers can treat this like the underlying model | |
| # ------------------------------------------------------------------ # | |
| def parameters(self, recurse: bool = True): | |
| if self.backend == "cnn" and self._cnn_fallback is not None: | |
| return self._cnn_fallback.parameters(recurse=recurse) | |
| # When using Caduceus (frozen), only the heads are trainable. | |
| return list(self.efficacy_head.parameters()) + list(self.safety_head.parameters()) | |
| def train(self, mode: bool = True): | |
| if self.backend == "cnn" and self._cnn_fallback is not None: | |
| self._cnn_fallback.train(mode) | |
| else: | |
| super().train(mode) | |
| # Caduceus itself stays in eval mode (frozen backbone). | |
| if self._caduceus is not None: | |
| self._caduceus.eval() | |
| return self | |
| def eval(self): | |
| return self.train(False) | |