ferritin_plms/esmfold2/layers/
confidence_head.rs1use super::pairformer::PairformerBlock;
21use candle_core::{Result, Tensor};
22use candle_nn::{self as nn, Module, VarBuilder, ops::softmax};
23
24pub 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 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 (probs * centres.broadcast_as(logits.shape())?)?.sum(candle_core::D::Minus1)
46}
47
48pub fn plddt_from_logits(logits: &Tensor) -> Result<Tensor> {
50 bins_to_scalar(logits, 0.0, 1.0)
51}
52
53pub struct ConfidenceOutput {
57 pub plddt_logits: Tensor,
59 pub plddt: Tensor,
61 pub pae_logits: Option<Tensor>,
63 pub pde_logits: Option<Tensor>,
65 pub distogram_logits: Option<Tensor>,
67}
68
69pub 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
79const CONFIDENCE_N_HEADS: usize = 8;
81
82impl ConfidenceHead {
83 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 pub fn forward(&self, single: &Tensor, pair: &Tensor) -> Result<ConfidenceOutput> {
119 let mut pair = pair.clone();
121 for block in &self.trunk {
122 pair = block.forward(&pair)?;
123 }
124
125 let plddt_logits = self.plddt_head.forward(single)?;
127 let plddt = plddt_from_logits(&plddt_logits)?;
128
129 let pae_logits = self.pae_head.forward(&pair)?;
131
132 let pde_logits = self.pde_head.forward(&pair)?;
134
135 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#[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 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 let mut logits_vec = vec![0.0f32; 10];
201 logits_vec[0] = 100.0; 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 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}