Skip to main content

ferritin_plms/esmfold2/layers/
confidence_head.rs

1//! Confidence head: predicts pLDDT, pAE, pDE, and distogram.
2//!
3//! Architecture:
4//! 1. 4-layer Pairformer trunk refines (single, pair)
5//! 2. Four linear heads project to logit bins:
6//!    - pLDDT:     single → 50 bins  → softmax → weighted mean → scalar [0,1]
7//!    - pAE:       pair   → 64 bins  → softmax → weighted mean → matrix [Å]
8//!    - pDE:       pair   → 64 bins  (predicted distance error)
9//!    - distogram: pair   → 39 bins
10//!
11//! Weight layout (rooted at `confidence_head`):
12//! ```text
13//! trunk.blocks.{0..3}.*   — 4-layer Pairformer (8 heads, d_pair=256)
14//! plddt_head.weight        — linear d_single → num_plddt_bins (no bias)
15//! pae_head.weight          — linear d_pair   → num_pae_bins   (no bias)
16//! pde_head.weight          — linear d_pair   → num_pde_bins   (no bias)
17//! distogram_head.weight    — linear d_pair   → distogram_bins (no bias)
18//! ```
19
20use super::pairformer::PairformerBlock;
21use candle_core::{Result, Tensor};
22use candle_nn::{self as nn, Module, VarBuilder, ops::softmax};
23
24// ── Helpers ────────────────────────────────────────────────────────────────
25
26/// Convert bin logits to a scalar via softmax + weighted average of bin centres.
27///
28/// `logits`    — `[..., n_bins]`
29/// `min_val`   — value of the first bin centre
30/// `max_val`   — value of the last bin centre
31///
32/// Returns `[...]` with the same leading dims, values in `[min_val, max_val]`.
33pub fn bins_to_scalar(logits: &Tensor, min_val: f64, max_val: f64) -> Result<Tensor> {
34    let n_bins = logits.dim(candle_core::D::Minus1)?;
35    let probs = softmax(logits, candle_core::D::Minus1)?;
36
37    // Build bin centres tensor on the same device
38    let step = (max_val - min_val) / (n_bins - 1) as f64;
39    let centres: Vec<f32> = (0..n_bins)
40        .map(|i| (min_val + i as f64 * step) as f32)
41        .collect();
42    let centres = Tensor::from_vec(centres, n_bins, logits.device())?.to_dtype(logits.dtype())?;
43
44    // Weighted sum over last dim
45    (probs * centres.broadcast_as(logits.shape())?)?.sum(candle_core::D::Minus1)
46}
47
48/// Convert pLDDT logits `[B, N, 50]` → scalar pLDDT `[B, N]` in `[0, 1]`.
49pub fn plddt_from_logits(logits: &Tensor) -> Result<Tensor> {
50    bins_to_scalar(logits, 0.0, 1.0)
51}
52
53// ── ConfidenceHead ─────────────────────────────────────────────────────────
54
55/// Output of the confidence head.
56pub struct ConfidenceOutput {
57    /// pLDDT logits `[B, N_tok, num_plddt_bins]`.
58    pub plddt_logits: Tensor,
59    /// pLDDT scalar per token `[B, N_tok]`, range [0, 1].
60    pub plddt: Tensor,
61    /// pAE logits `[B, N_tok, N_tok, num_pae_bins]`. `None` if not produced.
62    pub pae_logits: Option<Tensor>,
63    /// pDE logits `[B, N_tok, N_tok, num_pde_bins]`. `None` if not produced.
64    pub pde_logits: Option<Tensor>,
65    /// Distogram logits `[B, N_tok, N_tok, distogram_bins]`. `None` if not produced.
66    pub distogram_logits: Option<Tensor>,
67}
68
69/// 4-layer Pairformer confidence head.
70pub struct ConfidenceHead {
71    trunk: Vec<PairformerBlock>,
72    plddt_head: nn::Linear,
73    pae_head: nn::Linear,
74    pde_head: nn::Linear,
75    distogram_head: nn::Linear,
76    num_plddt_bins: usize,
77}
78
79/// Number of attention heads in the confidence head Pairformer trunk (same as folding trunk).
80const CONFIDENCE_N_HEADS: usize = 8;
81
82impl ConfidenceHead {
83    /// Load the confidence head from a `VarBuilder` rooted at `confidence_head.*`.
84    pub fn load(
85        vb: VarBuilder,
86        d_single: usize,
87        d_pair: usize,
88        num_plddt_bins: usize,
89        num_pae_bins: usize,
90        num_pde_bins: usize,
91        distogram_bins: usize,
92    ) -> Result<Self> {
93        let n_trunk_layers = 4;
94        let trunk = (0..n_trunk_layers)
95            .map(|i| {
96                PairformerBlock::load(
97                    vb.pp(format!("trunk.blocks.{i}")),
98                    d_pair,
99                    CONFIDENCE_N_HEADS,
100                )
101            })
102            .collect::<Result<Vec<_>>>()?;
103        Ok(Self {
104            trunk,
105            plddt_head: nn::linear_no_bias(d_single, num_plddt_bins, vb.pp("plddt_head"))?,
106            pae_head: nn::linear_no_bias(d_pair, num_pae_bins, vb.pp("pae_head"))?,
107            pde_head: nn::linear_no_bias(d_pair, num_pde_bins, vb.pp("pde_head"))?,
108            distogram_head: nn::linear_no_bias(d_pair, distogram_bins, vb.pp("distogram_head"))?,
109            num_plddt_bins,
110        })
111    }
112
113    /// Compute confidence predictions from trunk representations.
114    ///
115    /// # Arguments
116    /// * `single` — `[B, N_tok, d_single]`
117    /// * `pair`   — `[B, N_tok, N_tok, d_pair]`
118    pub fn forward(&self, single: &Tensor, pair: &Tensor) -> Result<ConfidenceOutput> {
119        // Refine pair with 4-layer Pairformer trunk (single passes through)
120        let mut pair = pair.clone();
121        for block in &self.trunk {
122            pair = block.forward(&pair)?;
123        }
124
125        // pLDDT: [B, N, d_single] → [B, N, 50] → scalar [B, N]
126        let plddt_logits = self.plddt_head.forward(single)?;
127        let plddt = plddt_from_logits(&plddt_logits)?;
128
129        // pAE: [B, N, N, d_pair] → [B, N, N, 64]
130        let pae_logits = self.pae_head.forward(&pair)?;
131
132        // pDE: [B, N, N, d_pair] → [B, N, N, 64]
133        let pde_logits = self.pde_head.forward(&pair)?;
134
135        // Distogram: [B, N, N, d_pair] → [B, N, N, 39]
136        let distogram_logits = self.distogram_head.forward(&pair)?;
137
138        Ok(ConfidenceOutput {
139            plddt_logits,
140            plddt,
141            pae_logits: Some(pae_logits),
142            pde_logits: Some(pde_logits),
143            distogram_logits: Some(distogram_logits),
144        })
145    }
146}
147
148// ── Tests ──────────────────────────────────────────────────────────────────
149
150#[cfg(test)]
151mod tests {
152    use super::*;
153    use candle_core::{Device, DType, Tensor};
154
155    const B: usize = 1;
156    const N: usize = 12;
157    const D_SINGLE: usize = 384;
158    const D_PAIR: usize = 256;
159
160    fn make_head() -> ConfidenceHead {
161        let device = Device::Cpu;
162        let vb = VarBuilder::zeros(DType::F32, &device);
163        ConfidenceHead::load(vb, D_SINGLE, D_PAIR, 50, 64, 64, 39).unwrap()
164    }
165
166    #[test]
167    fn test_confidence_head_output_shapes() {
168        let head = make_head();
169        let device = Device::Cpu;
170        let single = Tensor::zeros(&[B, N, D_SINGLE], DType::F32, &device).unwrap();
171        let pair = Tensor::zeros(&[B, N, N, D_PAIR], DType::F32, &device).unwrap();
172        let out = head.forward(&single, &pair).unwrap();
173
174        assert_eq!(out.plddt_logits.dims(), &[B, N, 50]);
175        assert_eq!(out.plddt.dims(), &[B, N]);
176        assert_eq!(out.pae_logits.unwrap().dims(), &[B, N, N, 64]);
177        assert_eq!(out.pde_logits.unwrap().dims(), &[B, N, N, 64]);
178        assert_eq!(out.distogram_logits.unwrap().dims(), &[B, N, N, 39]);
179    }
180
181    #[test]
182    fn test_plddt_range_with_uniform_logits() {
183        // Uniform logits → softmax uniform → weighted mean = midpoint = 0.5
184        let device = Device::Cpu;
185        let logits = Tensor::zeros(&[1, 4, 50], DType::F32, &device).unwrap();
186        let plddt = plddt_from_logits(&logits).unwrap();
187        let vals = plddt.flatten_all().unwrap().to_vec1::<f32>().unwrap();
188        for v in &vals {
189            assert!(
190                (*v - 0.5).abs() < 1e-4,
191                "uniform logits → pLDDT ≈ 0.5, got {v}"
192            );
193        }
194    }
195
196    #[test]
197    fn test_bins_to_scalar_extremes() {
198        let device = Device::Cpu;
199        // All logit mass on first bin → min value
200        let mut logits_vec = vec![0.0f32; 10];
201        logits_vec[0] = 100.0; // very large → softmax ≈ 1.0 at bin 0
202        let logits = Tensor::from_vec(logits_vec, &[1, 1, 10], &device).unwrap();
203        let scalar = bins_to_scalar(&logits, 0.0, 1.0).unwrap();
204        let val = scalar.flatten_all().unwrap().to_vec1::<f32>().unwrap()[0];
205        assert!(val < 0.01, "mass on first bin → scalar ≈ 0.0, got {val}");
206
207        // All mass on last bin → max value
208        let mut logits_vec = vec![0.0f32; 10];
209        logits_vec[9] = 100.0;
210        let logits = Tensor::from_vec(logits_vec, &[1, 1, 10], &device).unwrap();
211        let scalar = bins_to_scalar(&logits, 0.0, 1.0).unwrap();
212        let val = scalar.flatten_all().unwrap().to_vec1::<f32>().unwrap()[0];
213        assert!(val > 0.99, "mass on last bin → scalar ≈ 1.0, got {val}");
214    }
215}