Spaces:
Sleeping
Sleeping
| """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 | |
| # ------------------------------------------------------------------ # | |
| 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 | |
| # ------------------------------------------------------------------ # | |
| 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()} | |