Skip to main content

runtime/benchmark/
unillm_backend.rs

1//! UniLLM inference backend for benchmarking
2
3use 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
17/// UniLLM inference backend
18pub struct UniLLMBackend {
19    model: Option<LlamaModelV2>,
20    tokenizer: Option<Tokenizer>,
21    device: Device,
22}
23
24impl UniLLMBackend {
25    /// Create new UniLLM backend
26    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        // Load weights
50        let loader = UnifiedWeightLoader::new();
51        let weights = loader.load_weights(path)?;
52
53        // Create config from GGUF metadata
54        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        // Create tokenizer
61        self.tokenizer = Some(Tokenizer::from_model_weights(&weights)?);
62
63        // Create model
64        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        // Encode prompt
80        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        // Initialize KV cache
85        let num_layers = model.config().num_hidden_layers;
86        let mut cache = KVCache::new(num_layers);
87
88        // Track timing
89        let start = Instant::now();
90        let mut first_token_time: Option<std::time::Duration> = None;
91
92        // === PREFILL PHASE ===
93        // Process entire prompt at once
94        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        // Record time to first token (after prefill)
105        first_token_time = Some(start.elapsed());
106
107        // Get logits and sample first new token
108        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        // Greedy sampling
123        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        // Check EOS
135        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        // === DECODE PHASE ===
149        // Generate tokens one at a time using cache
150        for _ in 1..max_tokens {
151            // Create input tensor for SINGLE new token
152            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            // Forward with cache - only processes the new token!
165            let outputs = model.forward_with_cache(&inputs, Some(&mut cache))?;
166
167            // Get logits
168            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            // Greedy sampling
177            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            // Check EOS
189            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        // Decode output
200        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}