ferritin_plms/esm3/layers/
encode_inputs.rs1use crate::esm3::models::esm3::ESM3Config;
7use candle_core::{D, Module, Result, Tensor};
8use candle_nn::{self as nn, VarBuilder};
9
10pub 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 pub fn forward(&self, indices: &Tensor) -> Result<Tensor> {
35 let embedded = self.embed.forward(indices)?;
37
38 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 masked.sum(D::Minus2)
49 }
50}
51
52fn 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 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))?; 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)?; let diff = v.broadcast_sub(¢ers)?; diff.sqr()?.affine(-1.0 / denom, 0.0)?.exp()
76}
77
78pub struct EncodeInputs {
81 d_model: usize,
82 sequence_embed: nn::Embedding,
84 plddt_projection: nn::Linear,
86 structure_per_res_plddt_projection: nn::Linear,
87 structure_tokens_embed: nn::Embedding,
89 ss8_embed: nn::Embedding,
91 sasa_embed: nn::Embedding,
92 function_embeds: Vec<nn::Embedding>,
94 n_function_tracks: usize,
95 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 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 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, )?;
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 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 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 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)?; parts.push(self.function_embeds[i].forward(&track)?); }
223 let func_emb = Tensor::cat(&parts, D::Minus1)?; add(func_emb)?;
225 }
226
227 if let Some(res) = residue_annotation_tokens {
228 add(self.residue_embed.forward(res)?)?;
229 }
230
231 match x {
233 Some(t) => Ok(t),
234 None => candle_core::bail!("EncodeInputs: at least one input track must be provided"),
235 }
236 }
237}