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