Skip to main content

ferritin_plms/esm3/layers/
encode_inputs.rs

1//! ESM3 multi-track input embedding (EncodeInputs).
2//!
3//! Embeds up to 8 input tracks into a single `(B, L, d_model)` tensor by summing
4//! per-track contributions. All tracks are optional; missing tracks contribute zero.
5
6use crate::esm3::models::esm3::ESM3Config;
7use candle_core::{D, Module, Result, Tensor};
8use candle_nn::{self as nn, VarBuilder};
9
10// ── EmbeddingBag ─────────────────────────────────────────────────────────────
11
12/// Sum-mode embedding bag: looks up embeddings for multiple indices per position
13/// and sums them, zeroing out the padding index.
14pub struct EmbeddingBag {
15    embed: nn::Embedding,
16    padding_idx: u32,
17}
18
19impl EmbeddingBag {
20    pub fn load(
21        vb: VarBuilder,
22        vocab_size: usize,
23        embed_dim: usize,
24        padding_idx: u32,
25    ) -> Result<Self> {
26        Ok(Self {
27            embed: nn::embedding(vocab_size, embed_dim, vb)?,
28            padding_idx,
29        })
30    }
31
32    /// `indices`: `(*, K)` — K annotation IDs per position; `padding_idx` entries → zero.
33    /// Returns `(*, embed_dim)` via sum over K.
34    pub fn forward(&self, indices: &Tensor) -> Result<Tensor> {
35        // Look up: (*, K) → (*, K, embed_dim)
36        let embedded = self.embed.forward(indices)?;
37
38        // Mask out padding: where index == padding_idx, contribution = 0
39        let pad = self.padding_idx;
40        let mask = indices
41            .ne(pad)?
42            .unsqueeze(D::Minus1)?
43            .broadcast_as(embedded.shape())?
44            .to_dtype(embedded.dtype())?;
45        let masked = (embedded * mask)?;
46
47        // Sum over K dimension (second-to-last)
48        masked.sum(D::Minus2)
49    }
50}
51
52// ── RBF encoding ─────────────────────────────────────────────────────────────
53
54/// Radial Basis Function encoding of scalar values into `n_bins` features.
55///
56/// `values`: `(*)` float tensor in `[v_min, v_max]`.
57/// Returns `(*, n_bins)`.
58fn rbf(values: &Tensor, v_min: f64, v_max: f64, n_bins: usize) -> Result<Tensor> {
59    let device = values.device();
60    let dtype = values.dtype();
61
62    // Evenly spaced centers in [v_min, v_max]
63    let centers: Vec<f32> = (0..n_bins)
64        .map(|i| (v_min + (v_max - v_min) * i as f64 / (n_bins - 1) as f64) as f32)
65        .collect();
66    let centers = Tensor::new(centers.as_slice(), device)?
67        .to_dtype(dtype)?
68        .reshape((1, 1, n_bins))?; // (1, 1, n_bins) for broadcasting
69
70    let width = (v_max - v_min) / (n_bins - 1) as f64;
71    let denom = 2.0 * width * width;
72
73    let v = values.unsqueeze(D::Minus1)?; // (*, 1)
74    let diff = v.broadcast_sub(&centers)?; // (*, n_bins)
75    diff.sqr()?.affine(-1.0 / denom, 0.0)?.exp()
76}
77
78// ── EncodeInputs ──────────────────────────────────────────────────────────────
79
80pub struct EncodeInputs {
81    d_model: usize,
82    // Sequence track
83    sequence_embed: nn::Embedding,
84    // pLDDT tracks (projected from 16-bin RBF)
85    plddt_projection: nn::Linear,
86    structure_per_res_plddt_projection: nn::Linear,
87    // Structure token track
88    structure_tokens_embed: nn::Embedding,
89    // Secondary structure and SASA
90    ss8_embed: nn::Embedding,
91    sasa_embed: nn::Embedding,
92    // Function annotation tracks (8 separate embeddings concatenated)
93    function_embeds: Vec<nn::Embedding>,
94    n_function_tracks: usize,
95    // Residue (InterPro) annotation track (EmbeddingBag, sum mode)
96    residue_embed: EmbeddingBag,
97}
98
99impl EncodeInputs {
100    pub fn load(vb: VarBuilder, config: &ESM3Config) -> Result<Self> {
101        let d = config.d_model;
102        let n_rbf = config.n_rbf_bins;
103
104        let sequence_embed = nn::embedding(config.d_sequence_vocab, d, vb.pp("sequence_embed"))?;
105
106        let plddt_projection = nn::linear_no_bias(n_rbf, d, vb.pp("plddt_projection"))?;
107        let structure_per_res_plddt_projection =
108            nn::linear_no_bias(n_rbf, d, vb.pp("structure_per_res_plddt_projection"))?;
109
110        // Structure vocab: 4096 codes + 5 special tokens
111        let structure_tokens_embed = nn::embedding(
112            config.d_structure_vocab + 5,
113            d,
114            vb.pp("structure_tokens_embed"),
115        )?;
116
117        let ss8_embed = nn::embedding(config.d_ss8_vocab, d, vb.pp("ss8_embed"))?;
118        let sasa_embed = nn::embedding(config.d_sasa_vocab, d, vb.pp("sasa_embed"))?;
119
120        // 8 function-track embeddings, each produces d_model // 8 features
121        let func_dim = d / config.n_function_tracks;
122        let mut function_embeds = Vec::with_capacity(config.n_function_tracks);
123        for i in 0..config.n_function_tracks {
124            function_embeds.push(nn::embedding(
125                config.d_function_vocab,
126                func_dim,
127                vb.pp(format!("function_embed.{}", i)),
128            )?);
129        }
130
131        let residue_embed = EmbeddingBag::load(
132            vb.pp("residue_embed"),
133            config.d_residue_vocab,
134            d,
135            0, // padding_idx
136        )?;
137
138        Ok(Self {
139            d_model: d,
140            sequence_embed,
141            plddt_projection,
142            structure_per_res_plddt_projection,
143            structure_tokens_embed,
144            ss8_embed,
145            sasa_embed,
146            function_embeds,
147            n_function_tracks: config.n_function_tracks,
148            residue_embed,
149        })
150    }
151
152    /// Embed all input tracks and sum them.
153    ///
154    /// All arguments are optional; present tracks contribute their embedding,
155    /// absent tracks contribute zero.
156    ///
157    /// - `sequence_tokens`:           `(B, L)` u32 sequence token IDs.
158    /// - `structure_tokens`:          `(B, L)` u32 structure (VQ-VAE) tokens.
159    /// - `ss8_tokens`:                `(B, L)` u32 secondary-structure tokens.
160    /// - `sasa_tokens`:               `(B, L)` u32 SASA-bin tokens.
161    /// - `function_tokens`:           `(B, L, n_tracks)` u32 function annotation tokens.
162    /// - `residue_annotation_tokens`: `(B, L, K)` u32 InterPro annotation IDs.
163    /// - `average_plddt`:             `(B, L)` f32 average per-structure pLDDT in [0,1].
164    /// - `per_res_plddt`:             `(B, L)` f32 per-residue pLDDT in [0,1].
165    ///
166    /// Returns `(B, L, d_model)`.
167    pub fn forward(
168        &self,
169        sequence_tokens: Option<&Tensor>,
170        structure_tokens: Option<&Tensor>,
171        ss8_tokens: Option<&Tensor>,
172        sasa_tokens: Option<&Tensor>,
173        function_tokens: Option<&Tensor>,
174        residue_annotation_tokens: Option<&Tensor>,
175        average_plddt: Option<&Tensor>,
176        per_res_plddt: Option<&Tensor>,
177    ) -> Result<Tensor> {
178        let n_rbf = 16usize;
179
180        // Accumulate embeddings into x; first present track initialises x.
181        let mut x: Option<Tensor> = None;
182        let mut add = |t: Tensor| -> Result<()> {
183            x = Some(match x.take() {
184                None => t,
185                Some(acc) => acc.add(&t)?,
186            });
187            Ok(())
188        };
189
190        if let Some(seq) = sequence_tokens {
191            add(self.sequence_embed.forward(seq)?)?;
192        }
193
194        if let Some(st) = structure_tokens {
195            add(self.structure_tokens_embed.forward(st)?)?;
196        }
197
198        if let Some(ss8) = ss8_tokens {
199            add(self.ss8_embed.forward(ss8)?)?;
200        }
201
202        if let Some(sasa) = sasa_tokens {
203            add(self.sasa_embed.forward(sasa)?)?;
204        }
205
206        if let Some(plddt) = average_plddt {
207            let enc = rbf(plddt, 0.0, 1.0, n_rbf)?;
208            add(self.plddt_projection.forward(&enc)?)?;
209        }
210
211        if let Some(per_res) = per_res_plddt {
212            let enc = rbf(per_res, 0.0, 1.0, n_rbf)?;
213            add(self.structure_per_res_plddt_projection.forward(&enc)?)?;
214        }
215
216        if let Some(func) = function_tokens {
217            // func: (B, L, n_tracks); each track uses its own embedding
218            let mut parts: Vec<Tensor> = Vec::with_capacity(self.n_function_tracks);
219            for i in 0..self.n_function_tracks {
220                let track = func.narrow(D::Minus1, i, 1)?.squeeze(D::Minus1)?; // (B, L)
221                parts.push(self.function_embeds[i].forward(&track)?); // (B, L, func_dim)
222            }
223            let func_emb = Tensor::cat(&parts, D::Minus1)?; // (B, L, d_model)
224            add(func_emb)?;
225        }
226
227        if let Some(res) = residue_annotation_tokens {
228            add(self.residue_embed.forward(res)?)?;
229        }
230
231        // If no tracks provided, return zeros. In practice at least sequence_tokens is present.
232        match x {
233            Some(t) => Ok(t),
234            None => candle_core::bail!("EncodeInputs: at least one input track must be provided"),
235        }
236    }
237}