runtime/benchmark/
unillm_backend.rs1use std::path::Path;
4use std::time::Instant;
5
6use anyhow::Result;
7
8use crate::kv_cache::KVCache;
9use crate::model_core::{Model, ModelInputs, ModelOutputs};
10use crate::models_v2::llama::{LlamaConfig, LlamaModelV2};
11use crate::tensor_core::{Device, Tensor};
12use crate::tokenizer::Tokenizer;
13use crate::weight_loader_core::UnifiedWeightLoader;
14
15use super::{GenerationResult, InferenceBackend};
16
17pub struct UniLLMBackend {
19 model: Option<LlamaModelV2>,
20 tokenizer: Option<Tokenizer>,
21 device: Device,
22}
23
24impl UniLLMBackend {
25 pub fn new() -> Self {
27 Self {
28 model: None,
29 tokenizer: None,
30 device: Device::CPU,
31 }
32 }
33}
34
35impl Default for UniLLMBackend {
36 fn default() -> Self {
37 Self::new()
38 }
39}
40
41impl InferenceBackend for UniLLMBackend {
42 fn name(&self) -> &str {
43 "UniLLM"
44 }
45
46 fn load_model(&mut self, path: &Path) -> Result<f64> {
47 let start = Instant::now();
48
49 let loader = UnifiedWeightLoader::new();
51 let weights = loader.load_weights(path)?;
52
53 let config = if let Some(ref gguf_config) = weights.gguf_config {
55 LlamaConfig::from_gguf_config(gguf_config)
56 } else {
57 LlamaConfig::default()
58 };
59
60 self.tokenizer = Some(Tokenizer::from_model_weights(&weights)?);
62
63 self.model = Some(LlamaModelV2::from_weights(config, weights)?);
65
66 Ok(start.elapsed().as_secs_f64() * 1000.0)
67 }
68
69 fn generate(&mut self, prompt: &str, max_tokens: usize) -> Result<GenerationResult> {
70 let model = self
71 .model
72 .as_ref()
73 .ok_or_else(|| anyhow::anyhow!("Model not loaded"))?;
74 let tokenizer = self
75 .tokenizer
76 .as_ref()
77 .ok_or_else(|| anyhow::anyhow!("Tokenizer not loaded"))?;
78
79 let prompt_tokens_vec: Vec<u32> = tokenizer.encode_with_special_tokens(prompt, true, false);
81 let prompt_tokens = prompt_tokens_vec.len();
82 let mut tokens = prompt_tokens_vec.clone();
83
84 let num_layers = model.config().num_hidden_layers;
86 let mut cache = KVCache::new(num_layers);
87
88 let start = Instant::now();
90 let mut first_token_time: Option<std::time::Duration> = None;
91
92 let prompt_i64: Vec<i64> = prompt_tokens_vec.iter().map(|&t| t as i64).collect();
95 let prompt_tensor = Tensor::from_i64_slice(&prompt_i64, &[1, prompt_tokens], &self.device)?;
96 let inputs = ModelInputs::Text {
97 input_ids: prompt_tensor,
98 attention_mask: None,
99 position_ids: None,
100 };
101
102 let outputs = model.forward_with_cache(&inputs, Some(&mut cache))?;
103
104 first_token_time = Some(start.elapsed());
106
107 let logits = match outputs {
109 ModelOutputs::Logits { logits, .. } => logits,
110 _ => return Err(anyhow::anyhow!("Expected logits output")),
111 };
112
113 let logits_candle = logits.to_candle()?;
114 let shape = logits_candle.dims();
115 let seq_len = if shape.len() == 3 { shape[1] } else { shape[0] };
116 let last_logits = if shape.len() == 3 {
117 logits_candle.narrow(1, seq_len - 1, 1)?.squeeze(1)?.squeeze(0)?
118 } else {
119 logits_candle.narrow(0, seq_len - 1, 1)?.squeeze(0)?
120 };
121
122 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
124 let mut max_idx = 0;
125 let mut max_val = logits_vec[0];
126 for (idx, &val) in logits_vec.iter().enumerate() {
127 if val > max_val {
128 max_val = val;
129 max_idx = idx;
130 }
131 }
132 let mut next_token = max_idx as u32;
133
134 if next_token == tokenizer.eos_token_id() {
136 let total_time = start.elapsed();
137 return Ok(GenerationResult {
138 output_text: String::new(),
139 tokens_generated: 0,
140 prompt_tokens,
141 time_to_first_token_ms: first_token_time.map(|t| t.as_secs_f64() * 1000.0).unwrap_or(0.0),
142 total_time_ms: total_time.as_secs_f64() * 1000.0,
143 });
144 }
145
146 tokens.push(next_token);
147
148 for _ in 1..max_tokens {
151 let input_tensor = Tensor::from_i64_slice(
153 &[next_token as i64],
154 &[1, 1],
155 &self.device
156 )?;
157
158 let inputs = ModelInputs::Text {
159 input_ids: input_tensor,
160 attention_mask: None,
161 position_ids: None,
162 };
163
164 let outputs = model.forward_with_cache(&inputs, Some(&mut cache))?;
166
167 let logits = match outputs {
169 ModelOutputs::Logits { logits, .. } => logits,
170 _ => return Err(anyhow::anyhow!("Expected logits output")),
171 };
172
173 let logits_candle = logits.to_candle()?;
174 let last_logits = logits_candle.squeeze(0)?.squeeze(0)?;
175
176 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
178 let mut max_idx = 0;
179 let mut max_val = logits_vec[0];
180 for (idx, &val) in logits_vec.iter().enumerate() {
181 if val > max_val {
182 max_val = val;
183 max_idx = idx;
184 }
185 }
186 next_token = max_idx as u32;
187
188 if next_token == tokenizer.eos_token_id() {
190 break;
191 }
192
193 tokens.push(next_token);
194 }
195
196 let total_time = start.elapsed();
197 let tokens_generated = tokens.len() - prompt_tokens;
198
199 let output_text = tokenizer.decode(&tokens[prompt_tokens..]);
201
202 Ok(GenerationResult {
203 output_text,
204 tokens_generated,
205 prompt_tokens,
206 time_to_first_token_ms: first_token_time
207 .map(|t| t.as_secs_f64() * 1000.0)
208 .unwrap_or(0.0),
209 total_time_ms: total_time.as_secs_f64() * 1000.0,
210 })
211 }
212
213 fn memory_usage(&self) -> u64 {
214 super::runner::get_process_memory()
215 }
216
217 fn unload(&mut self) {
218 self.model = None;
219 self.tokenizer = None;
220 }
221}
222
223#[cfg(test)]
224mod tests {
225 use super::*;
226
227 #[test]
228 fn test_backend_creation() {
229 let backend = UniLLMBackend::new();
230 assert_eq!(backend.name(), "UniLLM");
231 }
232}