1use 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
29pub 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, 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 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 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
206impl 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 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 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}