Skip to main content

ferritin_plms/esmfold2/layers/
atom_encoder.rs

1//! Atom encoder: per-atom feature encoder for ESMFold2.
2//!
3//! Weight layout:
4//! ```text
5//! inputs.atom_encoder.blocks.{0..2}.*   — 3 atom transformer blocks
6//! ```
7
8use candle_core::{Result, Tensor};
9use candle_nn::VarBuilder;
10
11/// 3-block atom-level transformer encoder.
12///
13/// Maps per-atom features (`[B, N_atom, d_atom=128]`) to per-token atom
14/// representations (`[B, N_tok, d_token=768]`) via windowed attention and
15/// token-level pooling.
16pub struct AtomEncoder {
17    // TODO: blocks: Vec<AtomTransformerBlock> (3 blocks, SWA window=128)
18    // TODO: proj:   Linear (d_atom → d_token)
19    d_atom: usize,
20    d_token: usize,
21    device: candle_core::Device,
22}
23
24impl AtomEncoder {
25    /// Load the atom encoder from a `VarBuilder` rooted at `inputs.atom_encoder.*`.
26    ///
27    /// # Arguments
28    /// * `vb`       — builder rooted at `inputs.atom_encoder`
29    /// * `d_atom`   — per-atom feature dimension (128)
30    /// * `d_token`  — per-token output dimension (768)
31    /// * `n_blocks` — number of atom transformer blocks (3)
32    pub fn load(vb: VarBuilder, d_atom: usize, d_token: usize, n_blocks: usize) -> Result<Self> {
33        // TODO: load n_blocks atom transformer blocks with SWA window=128
34        let _ = n_blocks; // suppress unused warning until blocks are implemented
35        let device = vb.device().clone();
36        Ok(Self {
37            d_atom,
38            d_token,
39            device,
40        })
41    }
42
43    /// Encodes per-atom features to per-token atom representations.
44    ///
45    /// - Input:  `[B, N_atom, d_atom=128]`
46    /// - Output: `[B, N_tok, d_token=768]`
47    ///
48    /// The number of output tokens depends on the residue-level grouping of atoms.
49    pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
50        // TODO: run atom transformer blocks + token-level mean-pooling + projection
51        let _ = self.d_atom; // used only in weight loading
52        let (b, _n_atom, _) = x.dims3()?;
53        // Placeholder: single token per batch (real impl emits N_tok tokens)
54        Tensor::zeros((b, 1_usize, self.d_token), x.dtype(), &self.device)
55    }
56}