Skip to main content

ferritin_plms/esmfold2/layers/
pairformer.rs

1//! Pairformer block — core iterative refinement layer of the ESMFold2 FoldingTrunk.
2//!
3//! Each block refines the pair representation `[B, N, N, d_pair]` by applying:
4//! 1. Row-wise triangle attention  (starting-node pair bias)
5//! 2. Column-wise triangle attention (ending-node pair bias)
6//! 3. Triangle multiplicative update — outgoing  (z_ij += Σ_k a_ik ⊙ b_jk)
7//! 4. Triangle multiplicative update — incoming  (z_ij += Σ_k a_ki ⊙ b_kj)
8//! 5. Pair transition FFN
9//!
10//! Weight layout (rooted at `folding_trunk.blocks.{i}`):
11//! ```text
12//! tri_attn_row.norm.*       — pre-norm LayerNorm
13//! tri_attn_row.q_proj.*     — Q (no bias)
14//! tri_attn_row.k_proj.*     — K (no bias)
15//! tri_attn_row.v_proj.*     — V (no bias)
16//! tri_attn_row.pair_bias.*  — pair bias → n_heads (no bias)
17//! tri_attn_row.gate.*       — sigmoid gate → n_heads*d_head (no bias)
18//! tri_attn_row.out_proj.*   — output → d_pair (no bias)
19//! tri_attn_col.*            — identical layout
20//! tri_mult_out.norm.*       — input LayerNorm
21//! tri_mult_out.left_proj.*  — left projection → c_hidden (no bias)
22//! tri_mult_out.right_proj.* — right projection → c_hidden (no bias)
23//! tri_mult_out.left_gate.*  — left sigmoid gate → c_hidden (no bias)
24//! tri_mult_out.right_gate.* — right sigmoid gate → c_hidden (no bias)
25//! tri_mult_out.out_norm.*   — pre-output LayerNorm on c_hidden
26//! tri_mult_out.out_proj.*   — c_hidden → d_pair (no bias)
27//! tri_mult_out.out_gate.*   — output sigmoid gate → d_pair (no bias)
28//! tri_mult_in.*             — identical layout
29//! pair_trans.norm.*         — transition pre-norm
30//! pair_trans.fc1.*          — d_pair → 4*d_pair (no bias)
31//! pair_trans.fc2.*          — 4*d_pair → d_pair (no bias)
32//! ```
33
34use candle_core::{D, Result, Tensor};
35use candle_nn::{self as nn, LayerNorm, LayerNormConfig, Module, VarBuilder};
36
37// ── Triangle Multiplicative Update ───────────────────────────────────────────
38
39struct TriangleMult {
40    norm: LayerNorm,
41    left_proj: nn::Linear,
42    right_proj: nn::Linear,
43    left_gate: nn::Linear,
44    right_gate: nn::Linear,
45    out_norm: LayerNorm,
46    out_proj: nn::Linear,
47    out_gate: nn::Linear,
48    outgoing: bool,
49}
50
51impl TriangleMult {
52    fn load(vb: VarBuilder, d_pair: usize, c_hidden: usize, outgoing: bool) -> Result<Self> {
53        Ok(Self {
54            norm: nn::layer_norm(d_pair, LayerNormConfig::from(1e-5), vb.pp("norm"))?,
55            left_proj: nn::linear_no_bias(d_pair, c_hidden, vb.pp("left_proj"))?,
56            right_proj: nn::linear_no_bias(d_pair, c_hidden, vb.pp("right_proj"))?,
57            left_gate: nn::linear_no_bias(d_pair, c_hidden, vb.pp("left_gate"))?,
58            right_gate: nn::linear_no_bias(d_pair, c_hidden, vb.pp("right_gate"))?,
59            out_norm: nn::layer_norm(c_hidden, LayerNormConfig::from(1e-5), vb.pp("out_norm"))?,
60            out_proj: nn::linear_no_bias(c_hidden, d_pair, vb.pp("out_proj"))?,
61            out_gate: nn::linear_no_bias(d_pair, d_pair, vb.pp("out_gate"))?,
62            outgoing,
63        })
64    }
65
66    fn forward(&self, z: &Tensor) -> Result<Tensor> {
67        let z_n = self.norm.forward(z)?;
68
69        // Gated projections: [B, N, N, c_hidden]
70        let left_g = nn::ops::sigmoid(&self.left_gate.forward(&z_n)?)?;
71        let left = (self.left_proj.forward(&z_n)? * left_g)?;
72        let right_g = nn::ops::sigmoid(&self.right_gate.forward(&z_n)?)?;
73        let right = (self.right_proj.forward(&z_n)? * right_g)?;
74
75        // Permute to [B, c, N_i, N_k] for batch matmul.
76        // Outgoing: left[b,i,k,c] → permute(0,3,1,2) → [B,c,i,k]
77        // Incoming: left[b,k,i,c] treated as [B,c,i,k] → permute(0,3,2,1) swaps the N dims
78        let (left_p, right_p) = if self.outgoing {
79            (
80                left.permute((0, 3, 1, 2))?.contiguous()?,
81                right.permute((0, 3, 1, 2))?.contiguous()?,
82            )
83        } else {
84            (
85                left.permute((0, 3, 2, 1))?.contiguous()?,
86                right.permute((0, 3, 2, 1))?.contiguous()?,
87            )
88        };
89
90        // p[b,c,i,j] = Σ_k left_p[b,c,i,k] * right_p[b,c,j,k]
91        let p = left_p.matmul(&right_p.transpose(D::Minus2, D::Minus1)?.contiguous()?)?; // [B, c, N, N]
92        let p = p.permute((0, 2, 3, 1))?; // [B, N, N, c]
93        let p = self.out_norm.forward(&p)?;
94
95        let out_g = nn::ops::sigmoid(&self.out_gate.forward(&z_n)?)?; // [B, N, N, d_pair]
96        let out = (out_g * self.out_proj.forward(&p)?)?;
97        z + &out
98    }
99}
100
101// ── Triangle Attention ────────────────────────────────────────────────────────
102
103struct TriangleAttention {
104    norm: LayerNorm,
105    q_proj: nn::Linear,
106    k_proj: nn::Linear,
107    v_proj: nn::Linear,
108    pair_bias: nn::Linear,
109    gate: nn::Linear,
110    out_proj: nn::Linear,
111    n_heads: usize,
112    d_head: usize,
113    row_wise: bool,
114}
115
116impl TriangleAttention {
117    fn load(vb: VarBuilder, d_pair: usize, n_heads: usize, row_wise: bool) -> Result<Self> {
118        let d_head = d_pair / n_heads;
119        Ok(Self {
120            norm: nn::layer_norm(d_pair, LayerNormConfig::from(1e-5), vb.pp("norm"))?,
121            q_proj: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("q_proj"))?,
122            k_proj: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("k_proj"))?,
123            v_proj: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("v_proj"))?,
124            pair_bias: nn::linear_no_bias(d_pair, n_heads, vb.pp("pair_bias"))?,
125            gate: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("gate"))?,
126            out_proj: nn::linear_no_bias(n_heads * d_head, d_pair, vb.pp("out_proj"))?,
127            n_heads,
128            d_head,
129            row_wise,
130        })
131    }
132
133    // Operates on z [B, n1, n2, d]: treats n1 as independent rows, attends along n2.
134    fn forward_inner(&self, z: &Tensor) -> Result<Tensor> {
135        let (b, n1, n2, _) = z.dims4()?;
136        let (h, dh) = (self.n_heads, self.d_head);
137
138        let z_n = self.norm.forward(z)?;
139
140        let q = self.q_proj.forward(&z_n)?; // [B, n1, n2, H*dh]
141        let k = self.k_proj.forward(&z_n)?;
142        let v = self.v_proj.forward(&z_n)?;
143        let bias = self.pair_bias.forward(&z_n)?; // [B, n1, n2, H]
144        let gate = nn::ops::sigmoid(&self.gate.forward(&z_n)?)?; // [B, n1, n2, H*dh]
145
146        // Merge (B, n1) → B*n1 and split heads: [B,n1,n2,H*dh] → [B*n1, H, n2, dh]
147        let bn1 = b * n1;
148        let to_heads = |t: Tensor| -> Result<Tensor> {
149            t.reshape((bn1, n2, h, dh))?
150                .permute((0, 2, 1, 3))?
151                .contiguous()
152        };
153        let q = to_heads(q)?;
154        let k = to_heads(k)?;
155        let v = to_heads(v)?;
156
157        // Pair bias: [B,n1,n2,H] → [B*n1,H,1,n2] (broadcast over query positions)
158        let bias = bias
159            .reshape((bn1, n2, h))?
160            .permute((0, 2, 1))? // [B*n1, H, n2]
161            .unsqueeze(2)?; // [B*n1, H, 1, n2]
162
163        let scale = (dh as f64).sqrt();
164        let scores = (q.matmul(&k.transpose(D::Minus2, D::Minus1)?.contiguous()?)? / scale)?;
165        // scores: [B*n1, H, n2, n2]; bias: [B*n1, H, 1, n2] — broadcast over query positions
166        let scores = scores.broadcast_add(&bias)?;
167        let attn = nn::ops::softmax(&scores, D::Minus1)?;
168        let out = attn.matmul(&v)?; // [B*n1, H, n2, dh]
169
170        // [B*n1, H, n2, dh] → [B, n1, n2, H*dh]
171        let out = out
172            .permute((0, 2, 1, 3))?
173            .contiguous()? // [B*n1, n2, H, dh]
174            .reshape((b, n1, n2, h * dh))?;
175
176        let out = (gate * out)?;
177        self.out_proj.forward(&out)
178    }
179
180    fn forward(&self, z: &Tensor) -> Result<Tensor> {
181        let delta = if self.row_wise {
182            self.forward_inner(z)?
183        } else {
184            // Column-wise: transpose the two sequence dims before/after
185            let z_t = z.permute((0, 2, 1, 3))?;
186            self.forward_inner(&z_t)?.permute((0, 2, 1, 3))?
187        };
188        z + &delta
189    }
190}
191
192// ── Pair Transition FFN ───────────────────────────────────────────────────────
193
194struct PairTransition {
195    norm: LayerNorm,
196    fc1: nn::Linear,
197    fc2: nn::Linear,
198}
199
200impl PairTransition {
201    fn load(vb: VarBuilder, d_pair: usize) -> Result<Self> {
202        Ok(Self {
203            norm: nn::layer_norm(d_pair, LayerNormConfig::from(1e-5), vb.pp("norm"))?,
204            fc1: nn::linear_no_bias(d_pair, 4 * d_pair, vb.pp("fc1"))?,
205            fc2: nn::linear_no_bias(4 * d_pair, d_pair, vb.pp("fc2"))?,
206        })
207    }
208
209    fn forward(&self, z: &Tensor) -> Result<Tensor> {
210        let h = self.norm.forward(z)?;
211        let h = self.fc1.forward(&h)?.relu()?;
212        let h = self.fc2.forward(&h)?;
213        z + &h
214    }
215}
216
217// ── PairformerBlock ───────────────────────────────────────────────────────────
218
219/// One Pairformer layer: refines the pair representation [B, N, N, d_pair].
220///
221/// Used by both the 24-layer FoldingTrunk and the 4-layer ConfidenceHead trunk.
222/// Triangle multiplication hidden dim (`c_hidden`) is fixed at 128.
223pub struct PairformerBlock {
224    tri_attn_row: TriangleAttention,
225    tri_attn_col: TriangleAttention,
226    tri_mult_out: TriangleMult,
227    tri_mult_in: TriangleMult,
228    pair_trans: PairTransition,
229}
230
231/// Hidden dim for triangle multiplicative updates (standard: 128).
232const C_HIDDEN_MULT: usize = 128;
233
234impl PairformerBlock {
235    /// Load one Pairformer block from a `VarBuilder` rooted at `blocks.{i}`.
236    ///
237    /// `n_heads` must evenly divide `d_pair`; for the trunk, `n_heads=8`, `d_pair=256`.
238    pub fn load(vb: VarBuilder, d_pair: usize, n_heads: usize) -> Result<Self> {
239        Ok(Self {
240            tri_attn_row: TriangleAttention::load(vb.pp("tri_attn_row"), d_pair, n_heads, true)?,
241            tri_attn_col: TriangleAttention::load(vb.pp("tri_attn_col"), d_pair, n_heads, false)?,
242            tri_mult_out: TriangleMult::load(vb.pp("tri_mult_out"), d_pair, C_HIDDEN_MULT, true)?,
243            tri_mult_in: TriangleMult::load(vb.pp("tri_mult_in"), d_pair, C_HIDDEN_MULT, false)?,
244            pair_trans: PairTransition::load(vb.pp("pair_trans"), d_pair)?,
245        })
246    }
247
248    /// Run one Pairformer layer.
249    ///
250    /// Input/output: `[B, N, N, d_pair]`.
251    pub fn forward(&self, z: &Tensor) -> Result<Tensor> {
252        let z = self.tri_attn_row.forward(z)?;
253        let z = self.tri_attn_col.forward(&z)?;
254        let z = self.tri_mult_out.forward(&z)?;
255        let z = self.tri_mult_in.forward(&z)?;
256        self.pair_trans.forward(&z)
257    }
258}
259
260// ── Tests ─────────────────────────────────────────────────────────────────────
261
262#[cfg(test)]
263mod tests {
264    use super::*;
265    use candle_core::{Device,DType, Tensor};
266
267    const B: usize = 1;
268    const N: usize = 8;
269    const D_PAIR: usize = 32; // small dims for fast tests
270    const N_HEADS: usize = 4;
271    const C_HIDDEN: usize = 16;
272
273    fn zeros_pair(device: &Device) -> Tensor {
274        Tensor::zeros(&[B, N, N, D_PAIR], DType::F32, device).unwrap()
275    }
276
277    #[test]
278    fn test_triangle_mult_outgoing_shape() {
279        let device = Device::Cpu;
280        let vb = VarBuilder::zeros(DType::F32, &device);
281        let m = TriangleMult::load(vb, D_PAIR, C_HIDDEN, true).unwrap();
282        let out = m.forward(&zeros_pair(&device)).unwrap();
283        assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
284    }
285
286    #[test]
287    fn test_triangle_mult_incoming_shape() {
288        let device = Device::Cpu;
289        let vb = VarBuilder::zeros(DType::F32, &device);
290        let m = TriangleMult::load(vb, D_PAIR, C_HIDDEN, false).unwrap();
291        let out = m.forward(&zeros_pair(&device)).unwrap();
292        assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
293    }
294
295    #[test]
296    fn test_triangle_attn_row_shape() {
297        let device = Device::Cpu;
298        let vb = VarBuilder::zeros(DType::F32, &device);
299        let m = TriangleAttention::load(vb, D_PAIR, N_HEADS, true).unwrap();
300        let out = m.forward(&zeros_pair(&device)).unwrap();
301        assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
302    }
303
304    #[test]
305    fn test_triangle_attn_col_shape() {
306        let device = Device::Cpu;
307        let vb = VarBuilder::zeros(DType::F32, &device);
308        let m = TriangleAttention::load(vb, D_PAIR, N_HEADS, false).unwrap();
309        let out = m.forward(&zeros_pair(&device)).unwrap();
310        assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
311    }
312
313    #[test]
314    fn test_pair_transition_shape() {
315        let device = Device::Cpu;
316        let vb = VarBuilder::zeros(DType::F32, &device);
317        let m = PairTransition::load(vb, D_PAIR).unwrap();
318        let out = m.forward(&zeros_pair(&device)).unwrap();
319        assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
320    }
321
322    #[test]
323    fn test_pairformer_block_shape() {
324        let device = Device::Cpu;
325        let vb = VarBuilder::zeros(DType::F32, &device);
326        let block = PairformerBlock::load(vb, D_PAIR, N_HEADS).unwrap();
327        let out = block.forward(&zeros_pair(&device)).unwrap();
328        assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
329    }
330
331    #[test]
332    fn test_pairformer_block_batch_shape() {
333        let device = Device::Cpu;
334        let vb = VarBuilder::zeros(DType::F32, &device);
335        let block = PairformerBlock::load(vb, D_PAIR, N_HEADS).unwrap();
336        let z = Tensor::zeros(&[2, N, N, D_PAIR], DType::F32, &device).unwrap();
337        let out = block.forward(&z).unwrap();
338        assert_eq!(out.dims(), &[2, N, N, D_PAIR]);
339    }
340
341    #[test]
342    fn test_triangle_mult_math_outgoing() {
343        // Verify outgoing vs incoming produce the same shape but different values
344        // when z is asymmetric (upper triangle != lower triangle).
345        let device = Device::Cpu;
346        let n = 4;
347        let d = 8;
348        // Use random-ish values via arange
349        let vals: Vec<f32> = (0..n * n * d).map(|i| i as f32 * 0.01).collect();
350        let z = Tensor::from_vec(vals, &[1, n, n, d], &device).unwrap();
351
352        let vb_out = VarBuilder::zeros(DType::F32, &device);
353        let vb_in = VarBuilder::zeros(DType::F32, &device);
354        let m_out = TriangleMult::load(vb_out, d, 4, true).unwrap();
355        let m_in = TriangleMult::load(vb_in, d, 4, false).unwrap();
356
357        let out_out = m_out.forward(&z).unwrap();
358        let out_in = m_in.forward(&z).unwrap();
359        assert_eq!(out_out.dims(), &[1, n, n, d]);
360        assert_eq!(out_in.dims(), &[1, n, n, d]);
361        // With zero weights both will be zero + residual = z; shapes are equal
362        assert_eq!(out_out.dims(), out_in.dims());
363    }
364}