ferritin_plms/amplify/
amplify_runner.rs1use 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 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 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}