Skip to main content

catgrad_llm/
run.rs

1//! A stripped-down version of ModelRunner from catgrad examples, intended for serving
2use catgrad::{
3    backend::cpu::{
4        eval::{Builder, EvalState},
5        ndarray::{NdArray, TaggedNdArray},
6    },
7    core::nn::layers::{argmax, cast, reshape},
8    core::{Dtype, NdArrayType, Shape, Var},
9};
10use minijinja::{Environment, context};
11use minijinja_contrib::pycompat::unknown_method_callback;
12use std::collections::HashMap;
13use std::path::PathBuf;
14use std::rc::Rc;
15use tokenizers::tokenizer::{Result, Tokenizer};
16
17use crate::models::gemma::Model as GemmaModel;
18use crate::models::gpt2::Model as GPT2Model;
19use crate::models::llama::Model as LlamaModel;
20use crate::models::olmo::Model as OlmoModel;
21use crate::models::phi::Model as PhiModel;
22use crate::models::qwen::Model as QwenModel;
23use crate::models::utils::{Cache, Config, ModelBuilder};
24
25use crate::utils::{get_model_files, read_safetensors_multiple};
26
27use crate::serve;
28
29/// Load model
30pub struct ModelLoader {
31    config: Config,
32    model_paths: Vec<PathBuf>,
33    tokenizer_path: PathBuf,
34    tokenizer_config_path: PathBuf,
35    use_kv_cache: bool,
36}
37
38fn read_to_value<V: for<'a> serde::Deserialize<'a>>(path: PathBuf) -> Result<V> {
39    let config_str = &std::fs::read_to_string(path).map_err(|e| serve::Error::IO(e.to_string()))?;
40    let result: V =
41        serde_json::from_str(config_str).map_err(|e| serve::Error::IO(e.to_string()))?;
42    Ok(result)
43}
44
45impl ModelLoader {
46    pub fn new(model_name: &str, use_kv_cache: bool) -> serve::Result<Self> {
47        let (model_paths, config_path, tokenizer_path, tokenizer_config_path) =
48            get_model_files(model_name);
49
50        let config: Config = read_to_value(config_path)?;
51
52        Ok(Self {
53            config,
54            model_paths,
55            tokenizer_path,
56            tokenizer_config_path,
57            use_kv_cache,
58        })
59    }
60}
61
62pub struct ModelTokenizer {
63    pub tokenizer: Tokenizer,
64    pub chat_template: String,
65}
66
67impl ModelTokenizer {
68    fn new(tokenizer_path: PathBuf, tokenizer_config_path: PathBuf) -> serve::Result<Self> {
69        let tokenizer = Tokenizer::from_file(tokenizer_path)?;
70
71        let tokenizer_config: serde_json::Value = read_to_value(tokenizer_config_path)?;
72        let chat_template = tokenizer_config
73            .get("chat_template")
74            .and_then(|v| v.as_str())
75            .unwrap_or("")
76            .to_string();
77
78        Ok(Self {
79            tokenizer,
80            chat_template,
81        })
82    }
83
84    fn render_context(&self, messages: &[serve::Message]) -> String {
85        let mut env = Environment::new();
86        env.set_unknown_method_callback(unknown_method_callback);
87        env.add_template("chat", &self.chat_template).unwrap();
88        let tmpl = env.get_template("chat").unwrap();
89        let message_context: Vec<_> = messages
90            .iter()
91            .map(|msg| context!(role => msg.role, content => msg.content))
92            .collect();
93        tmpl.render(context!(
94            messages => message_context,
95            add_generation_prompt => true,
96            enable_thinking => false
97        ))
98        .expect("template failed to render")
99    }
100}
101
102pub struct ModelRunner {
103    pub tensors: Rc<HashMap<String, TaggedNdArray>>,
104    pub state: Option<EvalState>,
105    pub model: Box<dyn ModelBuilder>,
106    pub use_kv_cache: bool,
107    pub config: Config,
108    pub context: Vec<i32>,
109}
110
111impl ModelRunner {
112    pub fn new(
113        model_paths: Vec<PathBuf>,
114        config: Config,
115        use_kv_cache: bool,
116    ) -> Result<ModelRunner> {
117        env_logger::init();
118
119        let arch = &config.architectures[0];
120
121        let mut model: Box<dyn ModelBuilder> = match arch.as_str() {
122            "LlamaForCausalLM" => Box::new(LlamaModel {}),
123            "Olmo2ForCausalLM" => Box::new(OlmoModel {}),
124            "Qwen3ForCausalLM" => Box::new(QwenModel {}),
125            "Gemma3ForCausalLM" => Box::new(GemmaModel {}),
126            "Phi3ForCausalLM" => Box::new(PhiModel {}),
127            "GPT2LMHeadModel" => Box::new(GPT2Model {}),
128            _ => return Err("Unknown architecture {arch}".into()),
129        };
130
131        let mut tensors = read_safetensors_multiple(model_paths);
132        model.post_load(&mut tensors);
133
134        Ok(Self {
135            tensors: Rc::new(tensors),
136            state: None, // TODO?
137            model,
138            use_kv_cache,
139            config,
140            context: vec![],
141        })
142    }
143
144    fn next_token(&self, builder: &Builder, logits: Var) -> Var {
145        let batches = logits.label.shape.0[0];
146        let am = argmax(builder, logits);
147        let am = reshape(builder, Shape(vec![batches, 1]), am);
148        cast(builder, Dtype::I32, am)
149    }
150
151    fn build(&mut self, tokens: usize) {
152        let batches = 1;
153        let in_type = NdArrayType::new(Shape(vec![batches, tokens]), Dtype::I32);
154
155        let state = EvalState::build(|builder| {
156            let x = Var::new(builder.clone(), in_type.clone());
157            let positions = x.label.shape.0[1];
158            let mut cache = Cache::init(builder, &self.config, positions, self.use_kv_cache);
159            let result = self
160                .model
161                .build(builder, &self.config, &mut cache, 0, x.clone());
162            let new_token = self.next_token(builder, result);
163            (vec![x], vec![new_token])
164        });
165
166        self.state = Some(state);
167        self.state
168            .as_mut()
169            .unwrap()
170            .set_parameters(Rc::clone(&self.tensors));
171    }
172
173    // Make a forward pass given a list of tokens
174    fn run(&mut self, x: &NdArray<i32>) -> TaggedNdArray {
175        let [result] = self
176            .state
177            .as_mut()
178            .unwrap()
179            .eval_with(vec![x.clone().into()])[..]
180        else {
181            panic!("unexpected result")
182        };
183
184        result.clone()
185    }
186
187    pub fn generate(&mut self) -> Option<i32> {
188        // TODO: store tokens as growable NdArray instead so we don't have to clone the context for
189        // NdArray to take ownership.
190        // Blocked by https://github.com/hellas-ai/catgrad/issues/117.
191        let tokens = self.context.clone();
192        let num_tokens = tokens.len();
193        let batches = 1;
194        let input = NdArray::new(tokens, Shape(vec![batches, num_tokens / batches]));
195        self.build(num_tokens);
196        let result = self.run(&input);
197
198        let token = result.data()[0] as i32;
199        if self.config.get_eos_token_ids().contains(&token) {
200            return None;
201        }
202        Some(token)
203    }
204}
205
206////////////////////////////////////////////////////////////////////////////////
207// Trait impls
208
209impl Iterator for ModelRunner {
210    type Item = i32;
211
212    fn next(&mut self) -> Option<Self::Item> {
213        let next_token = self.generate();
214        if let Some(token) = next_token {
215            self.context.push(token);
216        }
217        next_token
218    }
219}
220
221impl serve::LM<i32> for ModelRunner {
222    fn set_context(&mut self, context: Vec<i32>) {
223        self.context = context;
224    }
225}
226
227impl serve::Tokenizer<i32> for ModelTokenizer {
228    fn encode(&self, content: String) -> serve::Result<Vec<i32>> {
229        let tokens = self.tokenizer.encode(content, true)?;
230        Ok(tokens.get_ids().iter().map(|&x| x as i32).collect())
231    }
232
233    fn decode(&self, tokens: Vec<i32>) -> serve::Result<String> {
234        // TODO: efficiency?
235        // TODO: support u32 in interpreter to remove try_into().unwrap().
236        let tokens_u32: Vec<u32> = tokens.into_iter().map(|i| i.try_into().unwrap()).collect();
237        Ok(self.tokenizer.decode(&tokens_u32, false)?)
238    }
239}
240
241impl serve::ChatTokenizer<i32> for ModelTokenizer {
242    fn encode_messages(&self, messages: Vec<serve::Message>) -> serve::Result<Vec<i32>> {
243        // initialize context
244        let content = self.render_context(&messages);
245        use serve::Tokenizer;
246        self.encode(content)
247    }
248}
249
250impl serve::Loader<i32, ModelRunner, ModelTokenizer> for ModelLoader {
251    fn load_runner(&self) -> serve::Result<ModelRunner> {
252        Ok(ModelRunner::new(
253            self.model_paths.clone(),
254            self.config.clone(),
255            self.use_kv_cache,
256        )?)
257    }
258
259    fn load_tokenizer(&self) -> serve::Result<ModelTokenizer> {
260        ModelTokenizer::new(
261            self.tokenizer_path.clone(),
262            self.tokenizer_config_path.clone(),
263        )
264    }
265}