"""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") # ------------------------------------------------------------------ # @staticmethod 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)