Spaces:
Sleeping
Sleeping
File size: 9,372 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 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 | """bioai.models.discrete_vae -- discrete VAE for dsRNA precursor design.
Replaces the WGAN-GP from the old scaffold. WGAN-GP is unstable on discrete
sequences (gradient penalty assumes a continuous data manifold) and was
overkill anyway: we don't need adversarial training to *generate* 200-nt
dsRNA precursors, we need a smooth latent space we can sample from and
optimise against. A discrete VAE (Transformer encoder/decoder, continuous
latent, Gumbel or argmax decoding at inference) is the standard solution
and trains in seconds on a CPU laptop -- a hard requirement for the demo.
Specs (per Task D deliverable #4):
* Encoder: 2-layer Transformer over 200-nt precursor tokens (vocab: A,C,G,T)
* Latent: continuous, dim 64, reparameterisation trick
* Decoder: 2-layer Transformer, outputs logits over 4 vocab at each position
* Loss: standard ELBO = BCE reconstruction + KL(N(mu,sigma) || N(0,1))
* API: ``encode(x) -> (mu, logvar)``, ``decode(z) -> logits``,
``forward(x) -> (recon_logits, mu, logvar)``, ``sample(num_samples) -> tokens``
This is a **stretch goal**: the primary mode of the pipeline is ranking real
candidates via the SiRNACNN + Reynolds features; the VAE is only here to
demonstrate the design-from-scratch capability. The training script in
``bioai/training/train_vae.py`` verifies it converges on a tiny batch.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
# --------------------------------------------------------------------------- #
# Token embedding (sinusoidal positional + learned nucleotide embedding)
# --------------------------------------------------------------------------- #
class _TokenWithPosition(nn.Module):
def __init__(self, vocab_size: int, d_model: int, max_len: int = 512):
super().__init__()
self.tok = nn.Embedding(vocab_size, d_model)
# Standard sinusoidal positional encoding.
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, d_model, 2, dtype=torch.float) * (-torch.log(torch.tensor(10000.0)) / d_model)
)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer("pos_emb", pe.unsqueeze(0)) # (1, max_len, d_model)
def forward(self, tokens: torch.Tensor) -> torch.Tensor:
# tokens: (B, L) long
return self.tok(tokens) + self.pos_emb[:, : tokens.size(1), :]
# --------------------------------------------------------------------------- #
# DiscreteVAE
# --------------------------------------------------------------------------- #
class DiscreteVAE(nn.Module):
"""Transformer-encoder / Transformer-decoder discrete VAE for 200-nt dsRNA precursors."""
def __init__(
self,
seq_len: int = 200,
vocab_size: int = 4,
latent_dim: int = 64,
hidden_dim: int = 128,
num_heads: int = 4,
num_layers: int = 2,
max_len: int = 512,
):
super().__init__()
self.seq_len = seq_len
self.vocab_size = vocab_size
self.latent_dim = latent_dim
self.hidden_dim = hidden_dim
# Encoder: token+pos embedding -> 2-layer TransformerEncoder
self.enc_embed = _TokenWithPosition(vocab_size, hidden_dim, max_len=max_len)
enc_layer = nn.TransformerEncoderLayer(
d_model=hidden_dim,
nhead=num_heads,
dim_feedforward=hidden_dim * 4,
dropout=0.1,
batch_first=True,
activation="gelu",
)
self.encoder = nn.TransformerEncoder(enc_layer, num_layers=num_layers)
# Latent projection: pooled encoder output -> (mu, logvar)
self.to_mu = nn.Linear(hidden_dim, latent_dim)
self.to_logvar = nn.Linear(hidden_dim, latent_dim)
# Decoder: z (broadcast) + token+pos embedding -> 2-layer TransformerDecoder
self.dec_embed = _TokenWithPosition(vocab_size, hidden_dim, max_len=max_len)
# Project latent z back to hidden_dim, then broadcast across positions.
self.z_proj = nn.Linear(latent_dim, hidden_dim)
dec_layer = nn.TransformerDecoderLayer(
d_model=hidden_dim,
nhead=num_heads,
dim_feedforward=hidden_dim * 4,
dropout=0.1,
batch_first=True,
activation="gelu",
)
self.decoder = nn.TransformerDecoder(dec_layer, num_layers=num_layers)
# Output head: per-position logits over vocab
self.lm_head = nn.Linear(hidden_dim, vocab_size)
# ------------------------------------------------------------------ #
def encode(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""``x``: ``(B, L)`` long tokens. Returns ``(mu, logvar)`` each ``(B, latent_dim)``."""
emb = self.enc_embed(x) # (B, L, H)
h = self.encoder(emb) # (B, L, H)
# Mean-pool over sequence length for a fixed-size summary.
pooled = h.mean(dim=1) # (B, H)
mu = self.to_mu(pooled)
logvar = self.to_logvar(pooled)
# Clamp logvar to keep KL finite and reparam stable.
logvar = torch.clamp(logvar, min=-10.0, max=10.0)
return mu, logvar
# ------------------------------------------------------------------ #
def reparameterize(self, mu: torch.Tensor, logvar: torch.Tensor) -> torch.Tensor:
std = torch.exp(0.5 * logvar)
eps = torch.randn_like(std)
return mu + eps * std
# ------------------------------------------------------------------ #
def decode(self, z: torch.Tensor, tgt_tokens: torch.Tensor | None = None) -> torch.Tensor:
"""``z``: ``(B, latent_dim)``. Returns ``logits`` ``(B, L, vocab_size)``.
If ``tgt_tokens`` is given, use it as the decoder input embeddings
(teacher forcing). Otherwise we use a learned-position-only input
(zeros) so the decoder is purely latent-conditioned -- this lets
:meth:`sample` work without a target sequence.
"""
B = z.size(0)
L = self.seq_len
z_h = self.z_proj(z).unsqueeze(1) # (B, 1, H)
if tgt_tokens is not None:
dec_in = self.dec_embed(tgt_tokens) # (B, L, H)
else:
# Zero tokens + positional encoding (still gets pos info).
zeros = torch.zeros(B, L, dtype=torch.long, device=z.device)
dec_in = self.dec_embed(zeros)
# Broadcast the latent across positions and add.
memory = z_h.expand(-1, L, -1) # (B, L, H)
# Use TransformerDecoder with `memory` = the latent broadcast.
# We pass `tgt` = dec_in (positions + optional teacher forcing).
h = self.decoder(tgt=dec_in, memory=memory)
logits = self.lm_head(h) # (B, L, V)
return logits
# ------------------------------------------------------------------ #
def forward(self, x: torch.Tensor):
"""``x``: ``(B, L)`` long tokens. Returns ``(recon_logits, mu, logvar)``.
``recon_logits`` is ``(B, L, vocab_size)``.
"""
mu, logvar = self.encode(x)
z = self.reparameterize(mu, logvar)
# Teacher forcing: feed the ground-truth tokens as decoder input
# (shifted right by inserting a leading zero-token, per LM convention).
B, L = x.shape
shifted = torch.cat(
[torch.zeros(B, 1, dtype=x.dtype, device=x.device), x[:, :-1]], dim=1
)
recon_logits = self.decode(z, tgt_tokens=shifted)
return recon_logits, mu, logvar
# ------------------------------------------------------------------ #
@torch.no_grad()
def sample(self, num_samples: int, device: str | torch.device = "cpu") -> torch.Tensor:
"""Generate ``num_samples`` new precursors. Returns ``(num_samples, seq_len)`` long tokens."""
self.eval()
z = torch.randn(num_samples, self.latent_dim, device=device)
logits = self.decode(z, tgt_tokens=None) # (N, L, V)
tokens = logits.argmax(dim=-1) # (N, L)
return tokens
# ------------------------------------------------------------------ #
@staticmethod
def elbo_loss(
recon_logits: torch.Tensor,
x: torch.Tensor,
mu: torch.Tensor,
logvar: torch.Tensor,
beta: float = 1.0,
) -> tuple[torch.Tensor, dict]:
"""Standard ELBO loss.
Reconstruction term: BCE over per-position vocab logits (treats each
position as an independent 4-way classification, which is standard
for discrete-sequence VAEs).
KL term: closed-form KL(N(mu, sigma) || N(0,1)).
"""
B, L, V = recon_logits.shape
# F.cross_entropy expects (N, C, ...) so we permute.
recon_loss = F.cross_entropy(
recon_logits.reshape(B * L, V),
x.reshape(B * L).long(),
reduction="mean",
)
# KL divergence to N(0,1), per-sample averaged.
kl = -0.5 * torch.mean(
torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1)
)
total = recon_loss + beta * kl
return total, {"recon": recon_loss.item(), "kl": kl.item(), "total": total.item()}
|