Skip to main content

ferritin_plms/esm3/models/
vqvae.rs

1//! ESM3 VQ-VAE structure token encoder and decoder (vqvae.py port).
2//!
3//! `StructureTokenEncoder` maps backbone coordinates to discrete structure tokens via a
4//! 2-layer geometric mini-transformer + nearest-neighbour VQ codebook lookup.
5//!
6//! `StructureTokenDecoder` is stubbed (not needed for ESM3 inference).
7
8use crate::esm3::layers::transformer_stack::TransformerStack;
9use crate::esm3::models::esm3::ESM3Config;
10use crate::esm3::utils::affine3d::Affine3D;
11use candle_core::{Module, Result, Tensor};
12use candle_nn::{self as nn, VarBuilder};
13
14// ── VQ-VAE config ─────────────────────────────────────────────────────────────
15
16#[derive(Debug, Clone)]
17pub struct VqVaeConfig {
18    // Encoder mini-transformer
19    pub enc_d_model: usize,  // 1024
20    pub enc_n_heads: usize,  // 1
21    pub enc_v_heads: usize,  // 128
22    pub enc_n_layers: usize, // 2
23    // Codebook
24    pub d_codebook: usize, // 128 (projected dimension before VQ)
25    pub n_codes: usize,    // 4096
26    // Decoder (kept for future use; not loaded in MVP)
27    pub dec_d_model: usize,  // 1280
28    pub dec_n_heads: usize,  // 20
29    pub dec_n_layers: usize, // 30
30}
31
32impl Default for VqVaeConfig {
33    fn default() -> Self {
34        Self {
35            enc_d_model: 1024,
36            enc_n_heads: 1,
37            enc_v_heads: 128,
38            enc_n_layers: 2,
39            d_codebook: 128,
40            n_codes: 4096,
41            dec_d_model: 1280,
42            dec_n_heads: 20,
43            dec_n_layers: 30,
44        }
45    }
46}
47
48impl VqVaeConfig {
49    /// Build an `ESM3Config` for the encoder TransformerStack.
50    ///
51    /// All layers use geometric attention; vocab sizes are unused and zeroed.
52    fn encoder_esm3_config(&self) -> ESM3Config {
53        ESM3Config {
54            d_model: self.enc_d_model,
55            n_heads: self.enc_n_heads,
56            n_layers: self.enc_n_layers,
57            n_layers_geom: self.enc_n_layers, // every layer uses geometric attention
58            v_head_transformer: self.enc_v_heads,
59            expansion_ratio: 8.0 / 3.0,
60            scale_residue: false,
61            mask_and_zero_frameless: true,
62            qk_layernorm: true,
63            bias: false,
64            // Vocab/embedding sizes unused by TransformerStack
65            d_sequence_vocab: 0,
66            d_structure_vocab: self.n_codes,
67            d_ss8_vocab: 0,
68            d_sasa_vocab: 0,
69            n_function_tracks: 0,
70            d_function_vocab: 0,
71            d_residue_vocab: 0,
72            n_rbf_bins: 0,
73        }
74    }
75}
76
77// ── Inference VQ codebook ─────────────────────────────────────────────────────
78
79/// Inference-only VQ codebook: loads the `(n_codes, d_codebook)` embedding table and
80/// performs nearest-neighbour quantization via L2 distances.
81struct VqCodebook {
82    embeddings: Tensor, // (n_codes, d_codebook)
83}
84
85impl VqCodebook {
86    pub fn load(vb: VarBuilder, n_codes: usize, d_codebook: usize) -> Result<Self> {
87        let embeddings = vb.get((n_codes, d_codebook), "embeddings")?;
88        Ok(Self { embeddings })
89    }
90
91    /// Quantize `z` to the nearest codebook entry.
92    ///
93    /// `z`: `(B, L, d_codebook)`.
94    /// Returns `(B, L)` u32 structure token indices.
95    pub fn quantize(&self, z: &Tensor) -> Result<Tensor> {
96        let (b, l, d) = z.dims3()?;
97        let z_flat = z.reshape((b * l, d))?; // (B*L, d)
98
99        // ||z - e||^2 = ||z||^2 - 2*(z @ e^T) + ||e||^2
100        let z_sq = z_flat.sqr()?.sum_keepdim(1)?; // (B*L, 1)
101        let z_et = z_flat.matmul(&self.embeddings.transpose(0, 1)?)?; // (B*L, n_codes)
102        let e_sq = self.embeddings.sqr()?.sum_keepdim(1)?.transpose(0, 1)?; // (1, n_codes)
103
104        // distances: (B*L, n_codes)
105        let distances = z_sq
106            .broadcast_sub(&z_et.affine(2.0, 0.0)?)?
107            .broadcast_add(&e_sq)?;
108
109        let indices = distances.argmin(1)?; // (B*L,) u32
110        indices.reshape((b, l))
111    }
112}
113
114// ── StructureTokenEncoder ─────────────────────────────────────────────────────
115
116/// Encodes backbone coordinates into discrete structure tokens via a 2-layer geometric
117/// mini-transformer followed by nearest-neighbour VQ codebook lookup.
118///
119/// Weight layout (in encoder checkpoint):
120/// - `encoder.blocks.*`   — transformer
121/// - `pre_vq_proj.weight` — `(d_codebook, enc_d_model)`
122/// - `codebook.embeddings`— `(n_codes, d_codebook)`
123pub struct StructureTokenEncoder {
124    transformer: TransformerStack,
125    pre_vq_proj: nn::Linear,
126    codebook: VqCodebook,
127    config: VqVaeConfig,
128}
129
130impl StructureTokenEncoder {
131    pub fn load(vb: VarBuilder, config: VqVaeConfig) -> Result<Self> {
132        let enc_cfg = config.encoder_esm3_config();
133        let transformer = TransformerStack::load(vb.pp("encoder"), &enc_cfg)?;
134        let pre_vq_proj =
135            nn::linear_no_bias(config.enc_d_model, config.d_codebook, vb.pp("pre_vq_proj"))?;
136        let codebook = VqCodebook::load(vb.pp("codebook"), config.n_codes, config.d_codebook)?;
137        Ok(Self {
138            transformer,
139            pre_vq_proj,
140            codebook,
141            config,
142        })
143    }
144
145    /// Encode backbone coordinates into structure tokens.
146    ///
147    /// - `coords`:      `(B, L, 3, 3)` backbone `(N, CA, C)` positions.
148    /// - `sequence_id`: optional `(B, L)` bin-packing IDs.
149    /// - `chain_id`:    optional `(B, L)` chain IDs.
150    ///
151    /// Returns `(B, L)` u32 structure token indices.
152    pub fn encode(
153        &self,
154        coords: &Tensor,
155        sequence_id: Option<&Tensor>,
156        chain_id: Option<&Tensor>,
157    ) -> Result<Tensor> {
158        let b = coords.dim(0)?;
159        let l = coords.dim(1)?;
160        let device = coords.device();
161        let dtype = coords.dtype();
162
163        let (affine, affine_mask) = Affine3D::build_affine3d_from_coordinates(coords)?;
164
165        // Initial hidden state: zeros (all structural info enters through geometric attention)
166        let x = Tensor::zeros((b, l, self.config.enc_d_model), dtype, device)?;
167
168        let (x, _pre_norm) = self.transformer.forward(
169            &x,
170            sequence_id,
171            Some(&affine),
172            Some(&affine_mask),
173            chain_id,
174        )?;
175
176        let z = self.pre_vq_proj.forward(&x)?;
177        self.codebook.quantize(&z)
178    }
179}
180
181// ── StructureTokenDecoder (stub) ──────────────────────────────────────────────
182
183/// Stub for the structure token decoder. Not needed for ESM3 inference.
184///
185/// The decoder (d_model=1280, n_heads=20, n_layers=30) maps structure tokens back to
186/// coordinates, but ESM3 inference only needs the encoder.
187pub struct StructureTokenDecoder;
188
189impl StructureTokenDecoder {
190    pub fn stub() -> Self {
191        Self
192    }
193}
194
195// ── Tests ─────────────────────────────────────────────────────────────────────
196
197#[cfg(test)]
198mod tests {
199    use super::*;
200    use candle_core::{Device, Tensor};
201
202    #[test]
203    fn test_vq_codebook_quantize_shape() -> Result<()> {
204        let device = &Device::Cpu;
205        let n_codes = 16usize;
206        let d = 8usize;
207        let b = 2usize;
208        let l = 5usize;
209
210        let embeddings = Tensor::randn(0f32, 1f32, (n_codes, d), device)?;
211        let codebook = VqCodebook { embeddings };
212
213        let z = Tensor::randn(0f32, 1f32, (b, l, d), device)?;
214        let tokens = codebook.quantize(&z)?;
215
216        assert_eq!(tokens.shape().dims(), &[b, l]);
217        Ok(())
218    }
219
220    #[test]
221    fn test_vq_codebook_nearest_neighbour() -> Result<()> {
222        let device = &Device::Cpu;
223        // Two codebook entries: [1,0] and [-1,0]
224        let embeddings = Tensor::new(&[[1f32, 0.], [-1., 0.]], device)?;
225        let codebook = VqCodebook { embeddings };
226
227        // z close to code 0 ([1,0]): should map to index 0
228        let z = Tensor::new(&[[[0.9f32, 0.1]]], device)?; // (1,1,2)
229        let tokens = codebook.quantize(&z)?;
230        assert_eq!(tokens.to_vec2::<u32>()?, vec![vec![0u32]]);
231
232        // z close to code 1 ([-1,0]): should map to index 1
233        let z = Tensor::new(&[[[-0.8f32, 0.1]]], device)?;
234        let tokens = codebook.quantize(&z)?;
235        assert_eq!(tokens.to_vec2::<u32>()?, vec![vec![1u32]]);
236        Ok(())
237    }
238
239    #[test]
240    fn test_vqvae_config_encoder_esm3_config() {
241        let cfg = VqVaeConfig::default();
242        let enc = cfg.encoder_esm3_config();
243        assert_eq!(enc.d_model, 1024);
244        assert_eq!(enc.n_heads, 1);
245        assert_eq!(enc.n_layers, 2);
246        assert_eq!(
247            enc.n_layers_geom, 2,
248            "all layers should use geometric attention"
249        );
250        assert_eq!(enc.v_head_transformer, 128);
251        assert!(!enc.scale_residue);
252        assert!(enc.mask_and_zero_frameless);
253    }
254}