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