Spaces:
Sleeping
Sleeping
File size: 6,534 Bytes
914512c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 | """bioai.models.sirna_cnn -- dilated 1D-CNN for siRNA (21-nt) efficacy prediction.
HyenaDNA-inspired: instead of a single deep stack we use 4 conv blocks with
**exponentially-growing dilations** (1, 2, 4, 8) and a kernel of 7, which gives
a receptive field of ``1 + (kernel-1) * sum(dilations) = 1 + 6 * 15 = 91`` nt.
That covers a 21-nt siRNA more than 4x over -- every output position sees the
entire input. This is the same "long-range via dilation" trick HyenaDNA uses,
just on a CNN backbone (so it trains in seconds on a CPU laptop, which the
Hackathon demo requires).
Inputs are ``(batch, 4, 21)`` one-hot encodings (A=0, C=1, G=2, T/U=3).
Outputs are ``(efficacy_pred, safety_pred)`` where:
* ``efficacy_pred`` is ``(batch, 1)`` sigmoid -> [0, 1] knockdown fraction
* ``safety_pred`` is ``(batch, num_safety_species)`` sigmoid -> per-species
off-target risk in [0, 1]
Hardware: works on CPU, CUDA, and ROCm (``device='auto'`` -> ``cuda`` if
``torch.cuda.is_available()`` returns True; on ROCm wheels that flag is set).
"""
from __future__ import annotations
from typing import Tuple
import torch
import torch.nn as nn
# --------------------------------------------------------------------------- #
# Device resolution helper (shared by all modules)
# --------------------------------------------------------------------------- #
def resolve_device(device: str = "auto") -> torch.device:
"""Map ``'auto' | 'cpu' | 'cuda'`` to a ``torch.device``.
On AMD ROCm PyTorch wheels, ``torch.cuda.is_available()`` returns ``True``
(ROCm masquerades as CUDA in the PyTorch API), so the ``'auto'`` default
correctly picks the GPU on both NVIDIA and AMD machines.
"""
if device == "auto":
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
if device == "cuda" and not torch.cuda.is_available():
print("CUDA/ROCm requested but unavailable; falling back to CPU.")
return torch.device("cpu")
return torch.device(device)
# --------------------------------------------------------------------------- #
# Dilated conv block
# --------------------------------------------------------------------------- #
class _DilatedConvBlock(nn.Module):
"""Conv1d -> BatchNorm1d -> GELU -> Dropout, with same-length dilated padding."""
def __init__(self, in_channels: int, out_channels: int,
kernel_size: int = 7, dilation: int = 1, dropout: float = 0.1):
super().__init__()
# 'same'-style dilated padding: pad both sides by (kernel-1)/2 * dilation
# kernel=7 -> (7-1)/2 = 3, so total padding = 3 * dilation per side.
pad = ((kernel_size - 1) // 2) * dilation
self.conv = nn.Conv1d(
in_channels, out_channels,
kernel_size=kernel_size,
padding=pad,
dilation=dilation,
)
self.bn = nn.BatchNorm1d(out_channels)
self.act = nn.GELU()
self.drop = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.drop(self.act(self.bn(self.conv(x))))
# --------------------------------------------------------------------------- #
# SiRNACNN
# --------------------------------------------------------------------------- #
class SiRNACNN(nn.Module):
"""Dilated multi-task CNN for siRNA efficacy + safety prediction.
Architecture (HyenaDNA-inspired receptive field via dilations):
input (B, 4, 21)
-> DilatedConvBlock(4, 64, dilation=1) # RF = 7
-> DilatedConvBlock(64, 64, dilation=2) # RF = 21
-> DilatedConvBlock(64, 64, dilation=4) # RF = 49
-> DilatedConvBlock(64, 64, dilation=8) # RF = 105 (>4x the siRNA)
-> AdaptiveMaxPool1d(1) -> flatten # (B, 64)
-> Linear(64, 64) -> GELU # shared trunk
heads:
-> Linear(64, 1) -> sigmoid # efficacy_pred (knockdown)
-> Linear(64, num_safety_species) -> sigmoid # safety_pred
"""
def __init__(self, seq_len: int = 21, num_safety_species: int = 6,
channels: int = 64, kernel_size: int = 7, dropout: float = 0.1):
super().__init__()
self.seq_len = seq_len
self.num_safety_species = num_safety_species
self.channels = channels
dilations = (1, 2, 4, 8)
blocks = []
in_ch = 4
for d in dilations:
blocks.append(_DilatedConvBlock(in_ch, channels, kernel_size, d, dropout))
in_ch = channels
self.trunk_conv = nn.Sequential(*blocks)
self.pool = nn.AdaptiveMaxPool1d(1)
self.trunk_fc = nn.Sequential(
nn.Linear(channels, channels),
nn.GELU(),
)
self.efficacy_head = nn.Linear(channels, 1)
self.safety_head = nn.Linear(channels, num_safety_species)
# ------------------------------------------------------------------ #
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""``x``: ``(batch, 4, seq_len)`` one-hot. Returns ``(efficacy, safety)``."""
if x.dim() == 2:
# accept flattened (batch, 4*seq_len) for backwards-compat
x = x.view(-1, 4, self.seq_len)
elif x.dim() == 4 and x.size(1) == 1:
# accept (batch, 1, 4, seq_len) just in case
x = x.squeeze(1)
if x.size(1) != 4:
raise ValueError(
f"SiRNACNN expects (batch, 4, seq_len); got {tuple(x.shape)}"
)
if x.size(2) != self.seq_len:
# gracefully resize rather than crash for slightly off inputs
x = nn.functional.interpolate(
x, size=self.seq_len, mode="nearest"
)
h = self.trunk_conv(x) # (B, C, L)
h = self.pool(h).squeeze(-1) # (B, C)
h = self.trunk_fc(h) # (B, C)
efficacy_pred = torch.sigmoid(self.efficacy_head(h)) # (B, 1)
safety_pred = torch.sigmoid(self.safety_head(h)) # (B, num_safety)
return efficacy_pred, safety_pred
# ------------------------------------------------------------------ #
@torch.no_grad()
def predict(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Inference helper (no grad, eval mode if caller forgot)."""
was_training = self.training
self.eval()
try:
return self.forward(x)
finally:
if was_training:
self.train()
|