Skip to main content

ferritin_plms/esmfold2/layers/
folding_trunk.rs

1//! Folding trunk: 24-layer Pairformer transformer.
2//!
3//! Weight layout:
4//! ```text
5//! folding_trunk.pair_init.*            — pair representation initializer
6//! folding_trunk.blocks.{0..23}.*       — Pairformer layers
7//! folding_trunk.norm.*                 — final single-repr LayerNorm (TODO)
8//! ```
9
10use super::pair_init::PairInit;
11use super::pairformer::PairformerBlock;
12use candle_core::{Result, Tensor};
13use candle_nn::{self as nn, LayerNorm, LayerNormConfig, Module, VarBuilder};
14
15/// 24-layer Pairformer trunk that jointly refines single and pair representations.
16///
17/// - Single repr: `[B, N, d_single=384]`
18/// - Pair repr:   `[B, N, N, d_pair=256]`
19pub struct FoldingTrunk {
20    pair_init: PairInit,
21    blocks: Vec<PairformerBlock>,
22    single_norm: LayerNorm,
23}
24
25impl FoldingTrunk {
26    /// Load the folding trunk from a `VarBuilder` rooted at `folding_trunk.*`.
27    ///
28    /// `n_heads` comes from `ESMFold2Config::trunk_n_heads` (8 for ESMFold2-Fast).
29    pub fn load(
30        vb: VarBuilder,
31        n_layers: usize,
32        d_single: usize,
33        d_pair: usize,
34        n_heads: usize,
35    ) -> Result<Self> {
36        let pair_init = PairInit::load(
37            vb.pp("pair_init"),
38            d_single,
39            d_pair,
40            32, // n_relpos_bins — from ESMFold2Config::n_relative_residx_bins
41            32, // d_outer — inner dim for outer-product projection
42        )?;
43        let blocks = (0..n_layers)
44            .map(|i| PairformerBlock::load(vb.pp(format!("blocks.{i}")), d_pair, n_heads))
45            .collect::<Result<Vec<_>>>()?;
46        let single_norm = nn::layer_norm(d_single, LayerNormConfig::from(1e-5), vb.pp("norm"))?;
47        Ok(Self {
48            pair_init,
49            blocks,
50            single_norm,
51        })
52    }
53
54    /// Initialise the pair representation from sequence and chain metadata.
55    ///
56    /// # Arguments
57    /// * `single`          — `[B, N, d_single]`
58    /// * `residue_indices` — `[B, N]` integer residue positions
59    /// * `chain_ids`       — `[B, N]` integer chain identifiers
60    ///
61    /// # Returns
62    /// `[B, N, N, d_pair]`
63    pub fn init_pair(
64        &self,
65        single: &Tensor,
66        residue_indices: &Tensor,
67        chain_ids: &Tensor,
68    ) -> Result<Tensor> {
69        self.pair_init.forward(single, residue_indices, chain_ids)
70    }
71
72    /// Run all 24 Pairformer layers.
73    ///
74    /// # Arguments
75    /// * `single` — `[B, N, d_single]`  (passed through; single update is TODO)
76    /// * `pair`   — `[B, N, N, d_pair]` (from `init_pair`)
77    ///
78    /// # Returns
79    /// `(single [B, N, d_single], pair [B, N, N, d_pair])`
80    pub fn forward(&self, single: &Tensor, pair: &Tensor) -> Result<(Tensor, Tensor)> {
81        let mut z = pair.clone();
82        for block in &self.blocks {
83            z = block.forward(&z)?;
84        }
85        let single = self.single_norm.forward(single)?;
86        Ok((single, z))
87    }
88}