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}