ferritin_plms/esm3/models/
vqvae.rs1use crate::esm3::layers::transformer_stack::TransformerStack;
9use crate::esm3::models::esm3::ESM3Config;
10use crate::esm3::utils::affine3d::Affine3D;
11use candle_core::{Module, Result, Tensor};
12use candle_nn::{self as nn, VarBuilder};
13
14#[derive(Debug, Clone)]
17pub struct VqVaeConfig {
18 pub enc_d_model: usize, pub enc_n_heads: usize, pub enc_v_heads: usize, pub enc_n_layers: usize, pub d_codebook: usize, pub n_codes: usize, pub dec_d_model: usize, pub dec_n_heads: usize, pub dec_n_layers: usize, }
31
32impl Default for VqVaeConfig {
33 fn default() -> Self {
34 Self {
35 enc_d_model: 1024,
36 enc_n_heads: 1,
37 enc_v_heads: 128,
38 enc_n_layers: 2,
39 d_codebook: 128,
40 n_codes: 4096,
41 dec_d_model: 1280,
42 dec_n_heads: 20,
43 dec_n_layers: 30,
44 }
45 }
46}
47
48impl VqVaeConfig {
49 fn encoder_esm3_config(&self) -> ESM3Config {
53 ESM3Config {
54 d_model: self.enc_d_model,
55 n_heads: self.enc_n_heads,
56 n_layers: self.enc_n_layers,
57 n_layers_geom: self.enc_n_layers, v_head_transformer: self.enc_v_heads,
59 expansion_ratio: 8.0 / 3.0,
60 scale_residue: false,
61 mask_and_zero_frameless: true,
62 qk_layernorm: true,
63 bias: false,
64 d_sequence_vocab: 0,
66 d_structure_vocab: self.n_codes,
67 d_ss8_vocab: 0,
68 d_sasa_vocab: 0,
69 n_function_tracks: 0,
70 d_function_vocab: 0,
71 d_residue_vocab: 0,
72 n_rbf_bins: 0,
73 }
74 }
75}
76
77struct VqCodebook {
82 embeddings: Tensor, }
84
85impl VqCodebook {
86 pub fn load(vb: VarBuilder, n_codes: usize, d_codebook: usize) -> Result<Self> {
87 let embeddings = vb.get((n_codes, d_codebook), "embeddings")?;
88 Ok(Self { embeddings })
89 }
90
91 pub fn quantize(&self, z: &Tensor) -> Result<Tensor> {
96 let (b, l, d) = z.dims3()?;
97 let z_flat = z.reshape((b * l, d))?; let z_sq = z_flat.sqr()?.sum_keepdim(1)?; let z_et = z_flat.matmul(&self.embeddings.transpose(0, 1)?)?; let e_sq = self.embeddings.sqr()?.sum_keepdim(1)?.transpose(0, 1)?; let distances = z_sq
106 .broadcast_sub(&z_et.affine(2.0, 0.0)?)?
107 .broadcast_add(&e_sq)?;
108
109 let indices = distances.argmin(1)?; indices.reshape((b, l))
111 }
112}
113
114pub struct StructureTokenEncoder {
124 transformer: TransformerStack,
125 pre_vq_proj: nn::Linear,
126 codebook: VqCodebook,
127 config: VqVaeConfig,
128}
129
130impl StructureTokenEncoder {
131 pub fn load(vb: VarBuilder, config: VqVaeConfig) -> Result<Self> {
132 let enc_cfg = config.encoder_esm3_config();
133 let transformer = TransformerStack::load(vb.pp("encoder"), &enc_cfg)?;
134 let pre_vq_proj =
135 nn::linear_no_bias(config.enc_d_model, config.d_codebook, vb.pp("pre_vq_proj"))?;
136 let codebook = VqCodebook::load(vb.pp("codebook"), config.n_codes, config.d_codebook)?;
137 Ok(Self {
138 transformer,
139 pre_vq_proj,
140 codebook,
141 config,
142 })
143 }
144
145 pub fn encode(
153 &self,
154 coords: &Tensor,
155 sequence_id: Option<&Tensor>,
156 chain_id: Option<&Tensor>,
157 ) -> Result<Tensor> {
158 let b = coords.dim(0)?;
159 let l = coords.dim(1)?;
160 let device = coords.device();
161 let dtype = coords.dtype();
162
163 let (affine, affine_mask) = Affine3D::build_affine3d_from_coordinates(coords)?;
164
165 let x = Tensor::zeros((b, l, self.config.enc_d_model), dtype, device)?;
167
168 let (x, _pre_norm) = self.transformer.forward(
169 &x,
170 sequence_id,
171 Some(&affine),
172 Some(&affine_mask),
173 chain_id,
174 )?;
175
176 let z = self.pre_vq_proj.forward(&x)?;
177 self.codebook.quantize(&z)
178 }
179}
180
181pub struct StructureTokenDecoder;
188
189impl StructureTokenDecoder {
190 pub fn stub() -> Self {
191 Self
192 }
193}
194
195#[cfg(test)]
198mod tests {
199 use super::*;
200 use candle_core::{Device, Tensor};
201
202 #[test]
203 fn test_vq_codebook_quantize_shape() -> Result<()> {
204 let device = &Device::Cpu;
205 let n_codes = 16usize;
206 let d = 8usize;
207 let b = 2usize;
208 let l = 5usize;
209
210 let embeddings = Tensor::randn(0f32, 1f32, (n_codes, d), device)?;
211 let codebook = VqCodebook { embeddings };
212
213 let z = Tensor::randn(0f32, 1f32, (b, l, d), device)?;
214 let tokens = codebook.quantize(&z)?;
215
216 assert_eq!(tokens.shape().dims(), &[b, l]);
217 Ok(())
218 }
219
220 #[test]
221 fn test_vq_codebook_nearest_neighbour() -> Result<()> {
222 let device = &Device::Cpu;
223 let embeddings = Tensor::new(&[[1f32, 0.], [-1., 0.]], device)?;
225 let codebook = VqCodebook { embeddings };
226
227 let z = Tensor::new(&[[[0.9f32, 0.1]]], device)?; let tokens = codebook.quantize(&z)?;
230 assert_eq!(tokens.to_vec2::<u32>()?, vec![vec![0u32]]);
231
232 let z = Tensor::new(&[[[-0.8f32, 0.1]]], device)?;
234 let tokens = codebook.quantize(&z)?;
235 assert_eq!(tokens.to_vec2::<u32>()?, vec![vec![1u32]]);
236 Ok(())
237 }
238
239 #[test]
240 fn test_vqvae_config_encoder_esm3_config() {
241 let cfg = VqVaeConfig::default();
242 let enc = cfg.encoder_esm3_config();
243 assert_eq!(enc.d_model, 1024);
244 assert_eq!(enc.n_heads, 1);
245 assert_eq!(enc.n_layers, 2);
246 assert_eq!(
247 enc.n_layers_geom, 2,
248 "all layers should use geometric attention"
249 );
250 assert_eq!(enc.v_head_transformer, 128);
251 assert!(!enc.scale_residue);
252 assert!(enc.mask_and_zero_frameless);
253 }
254}