Skip to main content

ferritin_plms/amplify/
amplify_runner.rs

1//! Amplify RUnner
2//!
3//! Class for loading and running the AMPLIFY models
4
5use super::super::types::{ContactMap, PseudoProbability};
6use super::amplify::{AMPLIFY, AmplifyOutput};
7use super::config::AMPLIFYConfig;
8use crate::plm_runner::PlmRunner;
9use anyhow::{Error as E, Result, anyhow};
10use candle_core::{D, DType, Device, Tensor};
11use candle_nn::VarBuilder;
12use candle_nn::ops;
13use hf_hub::HFClientSync;
14use tokenizers::Tokenizer;
15
16const AMPLIFY_DTYPE: DType = DType::F32;
17
18pub enum AmplifyModels {
19    AMP120M,
20    AMP350M,
21}
22impl AmplifyModels {
23    pub fn get_model_files(model: Self) -> (&'static str, &'static str) {
24        match model {
25            AmplifyModels::AMP120M => ("chandar-lab/AMPLIFY_120M", "main"),
26            AmplifyModels::AMP350M => ("chandar-lab/AMPLIFY_350M", "main"),
27        }
28    }
29}
30
31pub struct AmplifyRunner {
32    model: AMPLIFY,
33    tokenizer: Tokenizer,
34}
35impl AmplifyRunner {
36    pub fn load_model(modeltype: AmplifyModels, device: Device) -> Result<AmplifyRunner> {
37        let (model_id, revision) = AmplifyModels::get_model_files(modeltype);
38        let (owner, name) = model_id.split_once('/').unwrap_or(("", model_id));
39        let client = HFClientSync::new()?;
40        let repo = client.model(owner, name);
41        let (config_filename, tokenizer_filename, weights_filename) = {
42            let config = repo
43                .download_file()
44                .filename("config.json")
45                .revision(revision)
46                .send()?;
47            let tokenizer = repo
48                .download_file()
49                .filename("tokenizer.json")
50                .revision(revision)
51                .send()?;
52            let weights = repo
53                .download_file()
54                .filename("model.safetensors")
55                .revision(revision)
56                .send()?;
57            (config, tokenizer, weights)
58        };
59        let config_str = std::fs::read_to_string(config_filename)?;
60        let config_str = config_str
61            .replace("SwiGLU", "swiglu")
62            .replace("Swiglu", "swiglu");
63        let config: AMPLIFYConfig = serde_json::from_str(&config_str)?;
64        let tokenizer = Tokenizer::from_file(tokenizer_filename).map_err(E::msg)?;
65        let vb = unsafe {
66            VarBuilder::from_mmaped_safetensors(&[weights_filename], AMPLIFY_DTYPE, &device)?
67        };
68        let model = AMPLIFY::load(vb, &config)?;
69        Ok(AmplifyRunner { model, tokenizer })
70    }
71    pub fn run_forward(&self, prot_sequence: &str) -> Result<AmplifyOutput> {
72        let device = self.model.get_device();
73        let tokens = self
74            .tokenizer
75            .encode(prot_sequence.to_string(), false)
76            .map_err(E::msg)?
77            .get_ids()
78            .to_vec();
79        let token_ids = Tensor::new(&tokens[..], device)?.unsqueeze(0)?;
80        let encoded = self.model.forward(&token_ids, None, false, true)?;
81        Ok(encoded)
82    }
83    pub fn get_best_prediction(
84        &self,
85        prot_sequence: &str,
86    ) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
87        let model_output: AmplifyOutput = self.run_forward(prot_sequence)?;
88        let predictions = model_output.logits.argmax(D::Minus1)?;
89        let indices: Vec<u32> = predictions.to_vec2()?[0].to_vec();
90        let decoded = self.tokenizer.decode(indices.as_slice(), true)?;
91        let decoded = decoded.replace(" ", "");
92        Ok(decoded)
93    }
94    pub fn get_pseudo_probabilities(&self, prot_sequence: &str) -> Result<Vec<PseudoProbability>> {
95        let model_output: AmplifyOutput = self.run_forward(prot_sequence)?;
96        let predictions = model_output.logits;
97        let outputs = self.extract_logits(&predictions)?;
98        Ok(outputs)
99    }
100    pub fn get_contact_map(&self, prot_sequence: &str) -> Result<Vec<ContactMap>> {
101        let model_output: AmplifyOutput = self.run_forward(prot_sequence)?;
102        let contact_map_tensor = model_output.get_contact_map()?;
103        let averaged = contact_map_tensor.clone().unwrap().max_keepdim(D::Minus1)?;
104        let (position1, position2, val) = averaged.dims3()?;
105        let data = averaged.to_vec3::<f32>()?;
106
107        let mut contacts = Vec::new();
108        for i in 0..position1 {
109            for j in 0..position2 {
110                for k in 0..val {
111                    contacts.push(ContactMap {
112                        position_1: i,
113                        amino_acid_1: self
114                            .tokenizer
115                            .decode(&[i as u32], true)
116                            .ok()
117                            .and_then(|s| s.chars().next())
118                            .unwrap_or('?'),
119                        position_2: j,
120                        amino_acid_2: self
121                            .tokenizer
122                            .decode(&[j as u32], true)
123                            .ok()
124                            .and_then(|s| s.chars().next())
125                            .unwrap_or('?'),
126                        contact_estimate: data[i][j][k],
127                        layer: 1,
128                    });
129                }
130            }
131        }
132        Ok(contacts)
133    }
134    // Softmax and simplify
135    fn extract_logits(&self, tensor: &Tensor) -> Result<Vec<PseudoProbability>> {
136        let tensor = ops::softmax(tensor, D::Minus1)?;
137        let data = tensor.to_vec3::<f32>()?;
138        let (_, seq_len, vocab_size) = tensor.dims3()?;
139        let mut logit_positions = Vec::with_capacity(seq_len * vocab_size);
140        for seq_pos in 0..seq_len {
141            for vocab_idx in 0..vocab_size {
142                let score = data[0][seq_pos][vocab_idx];
143                let amino_acid_char = self
144                    .tokenizer
145                    .decode(&[vocab_idx as u32], false)
146                    .map_err(|e| anyhow!("Failed to decode: {}", e))?
147                    .chars()
148                    .next()
149                    .ok_or_else(|| anyhow!("Empty decoded string"))?;
150                logit_positions.push(PseudoProbability {
151                    position: seq_pos,
152                    amino_acid: amino_acid_char,
153                    pseudo_prob: score,
154                });
155            }
156        }
157        Ok(logit_positions)
158    }
159}
160
161impl PlmRunner for AmplifyRunner {
162    /// Run AMPLIFY and return the last-layer hidden states as per-residue embeddings.
163    ///
164    /// Shape: `(1, L, hidden_size)` where `L` includes BOS and EOS tokens.
165    fn embed(&self, sequence: &str) -> Result<Tensor> {
166        let device = self.model.get_device();
167        let tokens = self
168            .tokenizer
169            .encode(sequence.to_string(), false)
170            .map_err(E::msg)?
171            .get_ids()
172            .to_vec();
173        let token_ids = Tensor::new(&tokens[..], device)?.unsqueeze(0)?;
174        let output = self.model.forward(&token_ids, None, true, false)?;
175        let mut hidden_states = output
176            .hidden_states
177            .ok_or_else(|| anyhow!("AMPLIFY forward() returned no hidden states"))?;
178        hidden_states
179            .pop()
180            .ok_or_else(|| anyhow!("AMPLIFY returned empty hidden states list"))
181    }
182
183    fn model_name(&self) -> &str {
184        "amplify"
185    }
186}