"""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()}