Biopesticide-AI / bioai /models /caduceus_adapter.py
flvcko's picture
Biopesticide-AI: AMD Hackathon Unicorn Track submission
914512c
Raw
History Blame Contribute Delete
9.7 kB
"""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)