1use std::any::Any;
2use std::sync::atomic::AtomicUsize;
3use std::sync::Arc;
4use std::{fmt::Debug, str::FromStr};
5
6use anyhow::Result;
7use candle_core::{DType, Device, Tensor, D};
8use candle_nn::Conv2dConfig;
9use image::{ColorType, DynamicImage};
10use itertools::Itertools;
11use mistralrs_quant::log::once_log_info;
12use mistralrs_quant::ShardedVarBuilder;
13
14#[cfg(feature = "pyo3_macros")]
15use pyo3::pyclass;
16
17use regex::Regex;
18use serde::Deserialize;
19
20use self::minicpmo::{MiniCpmOConfig, MiniCpmOModel, MiniCpmOProcessor};
21
22use super::{DeviceMappedModelLoader, NonMappedSubModel, NormalLoadingMetadata};
23use crate::amoe::AnyMoeBaseModelMixin;
24use crate::attention::ATTENTION_CHUNK_SIZE;
25use crate::device_map::DeviceMapper;
26use crate::layers::Conv3dConfig;
27use crate::matformer::MatformerSliceConfig;
28use crate::paged_attention::{AttentionImplementation, ModelConfigLike, ModelConfigMetadata};
29use crate::pipeline::isq::IsqModelLoader;
30use crate::pipeline::loaders::AutoDeviceMapParams;
31use crate::pipeline::text_models_inputs_processor::{FlashParams, PagedAttentionInputMetadata};
32use crate::pipeline::{
33 EitherCache, IsqModel, Modalities, MultimodalPromptPrefixer, Processor, ProcessorCreator,
34 SupportedModality,
35};
36use crate::utils::varbuilder_utils::DeviceForLoadTensor;
37use crate::vision_models::clip::ClipConfig;
38use crate::vision_models::gemma3::config::Gemma3Config;
39use crate::vision_models::gemma3::{Gemma3Model, Gemma3Processor};
40use crate::vision_models::gemma3n::config::{Gemma3nConfig, IntermediateSize};
41use crate::vision_models::gemma3n::{Gemma3nModel, Gemma3nProcessor};
42use crate::vision_models::gemma4::config::Gemma4Config;
43use crate::vision_models::gemma4::{Gemma4Model, Gemma4Processor};
44use crate::vision_models::idefics2::{Config as Idefics2Config, Idefics2};
45use crate::vision_models::idefics2_input_processor::Idefics2Processor;
46use crate::vision_models::idefics3::{Idefics3Config, Idefics3Model, Idefics3Processor};
47use crate::vision_models::image_processor::ImagePreProcessor;
48use crate::vision_models::inputs_processor::Phi4MMProcessor;
49use crate::vision_models::llama4::{
50 self, Llama4Config, Llama4ImageProcessor, Llama4Model, Llama4Processor,
51};
52use crate::vision_models::llava::config::Config as LLaVAConfig;
53use crate::vision_models::llava15::Model as LLaVA;
54use crate::vision_models::llava_inputs_processor::{self, LLaVAProcessor};
55use crate::vision_models::llava_next::Model as LLaVANext;
56use crate::vision_models::llava_next_inputs_processor::{self, LLaVANextProcessor};
57use crate::vision_models::mistral3::{Mistral3Config, Mistral3Model, Mistral3Processor};
58use crate::vision_models::mllama::{MLlamaConfig, MLlamaModel, MLlamaProcessor};
59use crate::vision_models::phi3::{Config as Phi3Config, Model as Phi3, PHI3V_CLIP_CONFIG};
60use crate::vision_models::phi3_inputs_processor::Phi3Processor;
61use crate::vision_models::phi4::{Phi4MMConfig, Phi4MMModel, PHI4_MM_VISION_CFG};
62use crate::vision_models::preprocessor_config::PreProcessorConfig;
63use crate::vision_models::processor_config::ProcessorConfig;
64use crate::vision_models::qwen2_5_vl::{
65 Config as Qwen2_5VLConfig, Qwen2_5VLModel, Qwen2_5VLProcessor,
66};
67use crate::vision_models::qwen2vl::{Config as Qwen2VLConfig, Qwen2VLModel, Qwen2VLProcessor};
68use crate::vision_models::qwen3_5::{Config as Qwen3_5Config, Qwen3_5Model, Qwen3_5Processor};
69use crate::vision_models::qwen3_5_moe::{
70 Config as Qwen3_5MoeConfig, Qwen3_5MoeModel, Qwen3_5MoeProcessor,
71};
72use crate::vision_models::qwen3_vl::{Config as Qwen3VLConfig, Qwen3VLModel, Qwen3VLProcessor};
73use crate::vision_models::qwen3_vl_moe::{
74 Config as Qwen3VLMoEConfig, Qwen3VLMoEModel, Qwen3VLMoEProcessor,
75};
76use crate::vision_models::voxtral::config::VoxtralConfig;
77use crate::vision_models::voxtral::{VoxtralModel, VoxtralProcessor};
78use crate::vision_models::{minicpmo, phi4};
79
80pub trait MultimodalModel: IsqModel + AnyMoeBaseModelMixin {
81 #[allow(clippy::too_many_arguments)]
83 fn forward(
84 &self,
85 input_ids: &Tensor,
86 pixel_values: Option<Tensor>,
87 seqlen_offsets: &[usize],
88 context_lens: Vec<(usize, usize)>,
89 position_ids: Vec<usize>,
90 model_specific_args: Box<dyn Any>, metadata: Option<(Vec<(Tensor, Tensor)>, &PagedAttentionInputMetadata)>,
92 flash_params: &FlashParams,
93 ) -> candle_core::Result<Tensor>;
94 fn device(&self) -> &Device;
95 fn cache(&self) -> &EitherCache;
96 fn cache_mut(&mut self) -> &mut EitherCache;
97 fn max_seq_len(&self) -> usize;
98 fn config(&self) -> &ModelConfigMetadata;
99 fn model_config(&self) -> Arc<dyn ModelConfigLike + Send + Sync> {
100 Arc::new(self.config().clone())
101 }
102 fn default_model_specific_args(&self, input_ids: &Tensor) -> Box<dyn Any>;
104 fn encoder_cache_counters(&self) -> Option<(Arc<AtomicUsize>, Arc<AtomicUsize>)> {
106 None
107 }
108 fn reset_model_specific_state(&self) {}
111}
112
113pub trait MultimodalModelLoader: IsqModelLoader + Send + Sync + DeviceMappedModelLoader {
114 fn load(
115 &self,
116 config: &str,
117 vb: ShardedVarBuilder,
118 normal_loading_metadata: NormalLoadingMetadata,
119 attention_mechanism: AttentionImplementation,
120 ) -> Result<Box<dyn MultimodalModel + Send + Sync>>;
121 fn is_gptx(&self, config: &str) -> bool;
122 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>>;
123 fn get_processor(
124 &self,
125 model_config: &str,
126 processor_config: Option<ProcessorConfig>,
127 preprocessor_config: PreProcessorConfig,
128 max_edge: Option<u32>,
129 ) -> Arc<dyn Processor + Send + Sync>;
130 fn supports_paged_attention(&self, config: &str) -> bool;
131 fn supports_prefix_cacher(&self, _config: &str) -> bool {
132 false
134 }
135 fn modalities(&self, config: &str) -> Result<Modalities>;
136 fn prefixer(&self, config: &str) -> Arc<dyn MultimodalPromptPrefixer>;
137 fn default_chat_template(&self, _config: &str) -> Option<String> {
142 None
143 }
144 fn default_bos_eos(&self, _config: &str) -> Option<(String, String)> {
148 None
149 }
150 fn get_device_for_tensor(
151 &self,
152 config: &str,
153 _mapper: &dyn DeviceMapper,
154 loading_isq: bool,
155 ) -> Result<Arc<dyn Fn(String) -> DeviceForLoadTensor + Send + Sync + 'static>> {
156 if loading_isq {
157 Ok(Arc::new(|_| DeviceForLoadTensor::Base))
158 } else {
159 let re = Regex::new(r"\.layers\.(\d+)\.").unwrap();
160 let num_layers = self.model_config(config)?.num_layers();
161 let closure = move |name: String| {
162 if let Some(captures) = re.captures(&name) {
163 captures
164 .get(1)
165 .and_then(|m| m.as_str().parse::<usize>().ok())
166 .map(|l| l.min(num_layers))
167 .map(DeviceForLoadTensor::Idx)
168 .unwrap_or(DeviceForLoadTensor::Base)
169 } else {
170 DeviceForLoadTensor::Base
171 }
172 };
173
174 Ok(Arc::new(closure))
175 }
176 }
177}
178
179#[cfg_attr(feature = "pyo3_macros", pyclass(eq, eq_int))]
180#[derive(Clone, Debug, Deserialize, serde::Serialize, PartialEq)]
181pub enum MultimodalLoaderType {
183 #[serde(rename = "phi3v")]
184 Phi3V,
185 #[serde(rename = "idefics2")]
186 Idefics2,
187 #[serde(rename = "llava_next")]
188 LLaVANext,
189 #[serde(rename = "llava")]
190 LLaVA,
191 #[serde(rename = "vllama")]
192 VLlama,
193 #[serde(rename = "qwen2vl")]
194 Qwen2VL,
195 #[serde(rename = "idefics3")]
196 Idefics3,
197 #[serde(rename = "minicpmo")]
198 MiniCpmO,
199 #[serde(rename = "phi4mm")]
200 Phi4MM,
201 #[serde(rename = "qwen2_5vl")]
202 Qwen2_5VL,
203 #[serde(rename = "gemma3")]
204 Gemma3,
205 #[serde(rename = "mistral3")]
206 Mistral3,
207 #[serde(rename = "llama4")]
208 Llama4,
209 #[serde(rename = "gemma3n")]
210 Gemma3n,
211 #[serde(rename = "qwen3vl")]
212 Qwen3VL,
213 #[serde(rename = "qwen3vlmoe")]
214 Qwen3VLMoE,
215 #[serde(rename = "qwen3_5")]
216 Qwen3_5,
217 #[serde(rename = "qwen3_5moe")]
218 Qwen3_5Moe,
219 #[serde(rename = "voxtral")]
220 Voxtral,
221 #[serde(rename = "gemma4")]
222 Gemma4,
223}
224
225impl MultimodalLoaderType {
227 pub fn from_causal_lm_name(name: &str) -> Result<Self> {
228 match name {
229 "Phi3VForCausalLM" => Ok(Self::Phi3V),
230 "Idefics2ForConditionalGeneration" => Ok(Self::Idefics2),
231 "LlavaNextForConditionalGeneration" => Ok(Self::LLaVANext),
232 "LlavaForConditionalGeneration" => Ok(Self::LLaVA),
233 "MllamaForConditionalGeneration" => Ok(Self::VLlama),
234 "Qwen2VLForConditionalGeneration" => Ok(Self::Qwen2VL),
235 "Idefics3ForConditionalGeneration" => Ok(Self::Idefics3),
236 "MiniCPMO" => Ok(Self::MiniCpmO),
237 "Phi4MMForCausalLM" => Ok(Self::Phi4MM),
238 "Qwen2_5_VLForConditionalGeneration" => Ok(Self::Qwen2_5VL),
239 "Gemma3ForConditionalGeneration" | "Gemma3ForCausalLM" => Ok(Self::Gemma3),
240 "Mistral3ForConditionalGeneration" => Ok(Self::Mistral3),
241 "Llama4ForConditionalGeneration" => Ok(Self::Llama4),
242 "Gemma3nForConditionalGeneration" => Ok(Self::Gemma3n),
243 "Gemma4ForConditionalGeneration" => Ok(Self::Gemma4),
244 "Qwen3VLForConditionalGeneration" => Ok(Self::Qwen3VL),
245 "Qwen3VLMoeForConditionalGeneration" => Ok(Self::Qwen3VLMoE),
246 "Qwen3_5ForConditionalGeneration" => Ok(Self::Qwen3_5),
247 "Qwen3_5MoeForConditionalGeneration" => Ok(Self::Qwen3_5Moe),
248 "VoxtralForConditionalGeneration"
249 | "VoxtralRealtimeForConditionalGeneration" => Ok(Self::Voxtral),
250 other => anyhow::bail!(
251 "Unsupported Hugging Face Transformers -CausalLM model class `{other}`. Please raise an issue."
252 ),
253 }
254 }
255}
256
257impl FromStr for MultimodalLoaderType {
258 type Err = String;
259 fn from_str(s: &str) -> Result<Self, Self::Err> {
260 match s {
261 "phi3v" => Ok(Self::Phi3V),
262 "idefics2" => Ok(Self::Idefics2),
263 "llava_next" => Ok(Self::LLaVANext),
264 "llava" => Ok(Self::LLaVA),
265 "vllama" => Ok(Self::VLlama),
266 "qwen2vl" => Ok(Self::Qwen2VL),
267 "idefics3" => Ok(Self::Idefics3),
268 "minicpmo" => Ok(Self::MiniCpmO),
269 "phi4mm" => Ok(Self::Phi4MM),
270 "qwen2_5vl" => Ok(Self::Qwen2_5VL),
271 "gemma3" => Ok(Self::Gemma3),
272 "mistral3" => Ok(Self::Mistral3),
273 "llama4" => Ok(Self::Llama4),
274 "gemma3n" => Ok(Self::Gemma3n),
275 "gemma4" => Ok(Self::Gemma4),
276 "qwen3vl" => Ok(Self::Qwen3VL),
277 "qwen3vlmoe" => Ok(Self::Qwen3VLMoE),
278 "qwen3_5" => Ok(Self::Qwen3_5),
279 "qwen3_5moe" => Ok(Self::Qwen3_5Moe),
280 "voxtral" => Ok(Self::Voxtral),
281 a => Err(format!("Unknown architecture `{a}`. Possible architectures: `phi3v`, `idefics2`, `llava_next`, `llava`, `vllama`, `qwen2vl`, `idefics3`, `minicpmo`, `phi4mm`, `qwen2_5vl`, `gemma3`, `mistral3`, `llama4`, `gemma3n`, `gemma4`, `qwen3vl`, `qwen3vlmoe`, `qwen3_5`, `qwen3_5moe`, `voxtral`.")),
282 }
283 }
284}
285
286impl std::fmt::Display for MultimodalLoaderType {
287 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
288 let name = match self {
289 MultimodalLoaderType::Phi3V => "phi3v",
290 MultimodalLoaderType::Idefics2 => "idefics2",
291 MultimodalLoaderType::LLaVANext => "llava_next",
292 MultimodalLoaderType::LLaVA => "llava",
293 MultimodalLoaderType::VLlama => "vllama",
294 MultimodalLoaderType::Qwen2VL => "qwen2vl",
295 MultimodalLoaderType::Idefics3 => "idefics3",
296 MultimodalLoaderType::MiniCpmO => "minicpmo",
297 MultimodalLoaderType::Phi4MM => "phi4mm",
298 MultimodalLoaderType::Qwen2_5VL => "qwen2_5vl",
299 MultimodalLoaderType::Gemma3 => "gemma3",
300 MultimodalLoaderType::Mistral3 => "mistral3",
301 MultimodalLoaderType::Llama4 => "llama4",
302 MultimodalLoaderType::Gemma3n => "gemma3n",
303 MultimodalLoaderType::Qwen3VL => "qwen3vl",
304 MultimodalLoaderType::Qwen3VLMoE => "qwen3vlmoe",
305 MultimodalLoaderType::Qwen3_5 => "qwen3_5",
306 MultimodalLoaderType::Qwen3_5Moe => "qwen3_5moe",
307 MultimodalLoaderType::Voxtral => "voxtral",
308 MultimodalLoaderType::Gemma4 => "gemma4",
309 };
310 write!(f, "{name}")
311 }
312}
313
314#[derive(Deserialize)]
315struct AutoMultimodalLoaderConfig {
316 #[serde(default)]
317 architectures: Vec<String>,
318 #[serde(default)]
320 multimodal: Option<serde_json::Value>,
321}
322
323pub struct AutoMultimodalLoader;
325
326impl AutoMultimodalLoader {
327 fn get_loader(config: &str) -> Result<Box<dyn MultimodalModelLoader>> {
328 let auto_cfg: AutoMultimodalLoaderConfig = serde_json::from_str(config)?;
329
330 if auto_cfg.multimodal.is_some() && auto_cfg.architectures.is_empty() {
332 once_log_info("Automatic loader type determined to be `voxtral`");
333 return Ok(Box::new(VoxtralLoader));
334 }
335
336 if auto_cfg.architectures.len() != 1 {
337 anyhow::bail!("Expected exactly one architecture in config");
338 }
339
340 let name = &auto_cfg.architectures[0];
341 let tp = MultimodalLoaderType::from_causal_lm_name(name)?;
342
343 once_log_info(format!("Automatic loader type determined to be `{tp}`"));
344
345 Ok(match tp {
347 MultimodalLoaderType::Phi3V => Box::new(Phi3VLoader),
348 MultimodalLoaderType::Idefics2 => Box::new(Idefics2Loader),
349 MultimodalLoaderType::LLaVANext => Box::new(LLaVANextLoader),
350 MultimodalLoaderType::LLaVA => Box::new(LLaVALoader),
351 MultimodalLoaderType::VLlama => Box::new(VLlamaLoader),
352 MultimodalLoaderType::Qwen2VL => Box::new(Qwen2VLLoader),
353 MultimodalLoaderType::Idefics3 => Box::new(Idefics3Loader),
354 MultimodalLoaderType::MiniCpmO => Box::new(MiniCpmOLoader),
355 MultimodalLoaderType::Phi4MM => Box::new(Phi4MMLoader),
356 MultimodalLoaderType::Qwen2_5VL => Box::new(Qwen2_5VLLoader),
357 MultimodalLoaderType::Gemma3 => Box::new(Gemma3Loader),
358 MultimodalLoaderType::Mistral3 => Box::new(Mistral3Loader),
359 MultimodalLoaderType::Llama4 => Box::new(VLlama4Loader),
360 MultimodalLoaderType::Gemma3n => Box::new(Gemma3nLoader),
361 MultimodalLoaderType::Qwen3VL => Box::new(Qwen3VLLoader),
362 MultimodalLoaderType::Qwen3VLMoE => Box::new(Qwen3VLMoELoader),
363 MultimodalLoaderType::Qwen3_5 => Box::new(Qwen3_5Loader),
364 MultimodalLoaderType::Qwen3_5Moe => Box::new(Qwen3_5MoeLoader),
365 MultimodalLoaderType::Voxtral => Box::new(VoxtralLoader),
366 MultimodalLoaderType::Gemma4 => Box::new(Gemma4Loader),
367 })
368 }
369}
370
371impl MultimodalModelLoader for AutoMultimodalLoader {
372 fn load(
373 &self,
374 config: &str,
375 vb: ShardedVarBuilder,
376 normal_loading_metadata: NormalLoadingMetadata,
377 attention_mechanism: AttentionImplementation,
378 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
379 Self::get_loader(config)?.load(config, vb, normal_loading_metadata, attention_mechanism)
380 }
381
382 fn is_gptx(&self, config: &str) -> bool {
383 Self::get_loader(config)
384 .expect("AutoMultimodalLoader get_loader")
385 .is_gptx(config)
386 }
387
388 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
389 Self::get_loader(config)?.get_config_repr(config)
390 }
391
392 fn get_processor(
393 &self,
394 model_config: &str,
395 proc_cfg: Option<ProcessorConfig>,
396 preproc_cfg: PreProcessorConfig,
397 max_edge: Option<u32>,
398 ) -> Arc<dyn Processor + Send + Sync> {
399 Self::get_loader(model_config)
400 .expect("AutoMultimodalLoader get_loader")
401 .get_processor(model_config, proc_cfg, preproc_cfg, max_edge)
402 }
403
404 fn supports_paged_attention(&self, config: &str) -> bool {
405 Self::get_loader(config)
406 .expect("AutoMultimodalLoader")
407 .supports_paged_attention(config)
408 }
409
410 fn modalities(&self, config: &str) -> Result<Modalities> {
411 Self::get_loader(config)?.modalities(config)
412 }
413
414 fn supports_prefix_cacher(&self, config: &str) -> bool {
415 Self::get_loader(config)
416 .expect("AutoMultimodalLoader")
417 .supports_prefix_cacher(config)
418 }
419
420 fn prefixer(&self, config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
421 Self::get_loader(config)
422 .expect("AutoMultimodalLoader")
423 .prefixer(config)
424 }
425
426 fn default_chat_template(&self, config: &str) -> Option<String> {
427 Self::get_loader(config).ok()?.default_chat_template(config)
428 }
429
430 fn default_bos_eos(&self, config: &str) -> Option<(String, String)> {
431 Self::get_loader(config).ok()?.default_bos_eos(config)
432 }
433
434 fn get_device_for_tensor(
435 &self,
436 config: &str,
437 mapper: &dyn DeviceMapper,
438 loading_isq: bool,
439 ) -> Result<Arc<dyn Fn(String) -> DeviceForLoadTensor + Send + Sync + 'static>> {
440 Self::get_loader(config)?.get_device_for_tensor(config, mapper, loading_isq)
441 }
442}
443
444impl IsqModelLoader for AutoMultimodalLoader {
445 fn isq_layer_regexes(&self, config: &str) -> Result<Vec<Regex>> {
446 Self::get_loader(config)?.isq_layer_regexes(config)
447 }
448 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
449 Self::get_loader(config)?.immediate_isq_predicates(config)
450 }
451 fn isq_layer_regexes_moqe(&self, config: &str) -> Result<Vec<Regex>> {
452 Self::get_loader(config)?.isq_layer_regexes_moqe(config)
453 }
454 fn immediate_isq_predicates_moqe(&self, config: &str) -> Result<Vec<Regex>> {
455 Self::get_loader(config)?.immediate_isq_predicates_moqe(config)
456 }
457}
458
459impl DeviceMappedModelLoader for AutoMultimodalLoader {
460 fn mapped_max_act_size_elems(
461 &self,
462 config: &str,
463 params: &AutoDeviceMapParams,
464 ) -> Result<usize> {
465 Self::get_loader(config)?.mapped_max_act_size_elems(config, params)
466 }
467 fn non_mapped_max_act_size_elems(
468 &self,
469 config: &str,
470 params: &AutoDeviceMapParams,
471 ) -> Result<usize> {
472 Self::get_loader(config)?.non_mapped_max_act_size_elems(config, params)
473 }
474 fn non_mapped_size_in_bytes(
475 &self,
476 config: &str,
477 dtype: DType,
478 weight_pack_factor: usize,
479 _matformer_config: Option<&MatformerSliceConfig>,
480 ) -> Result<usize> {
481 Self::get_loader(config)?.non_mapped_size_in_bytes(
482 config,
483 dtype,
484 weight_pack_factor,
485 _matformer_config,
486 )
487 }
488 fn layer_sizes_in_bytes(
489 &self,
490 config: &str,
491 dtype: DType,
492 weight_pack_factor: usize,
493 _matformer_config: Option<&MatformerSliceConfig>,
494 ) -> Result<Vec<usize>> {
495 Self::get_loader(config)?.layer_sizes_in_bytes(
496 config,
497 dtype,
498 weight_pack_factor,
499 _matformer_config,
500 )
501 }
502 fn num_layers(&self, config: &str) -> Result<usize> {
503 Self::get_loader(config)?.num_layers(config)
504 }
505 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
506 Self::get_loader(config)?.model_config(config)
507 }
508}
509
510macro_rules! bias_if {
511 ($cond:expr, $size:expr) => {
512 if $cond {
513 $size
514 } else {
515 0
516 }
517 };
518}
519
520fn get_clip_vit_num_elems(cfg: &ClipConfig) -> usize {
521 let pre_layer_norm = cfg.hidden_size;
522 let final_layer_norm = cfg.hidden_size;
523
524 let num_patches = (cfg.image_size / cfg.patch_size).pow(2);
525 let num_positions = num_patches + 1;
526
527 let class_embedding = cfg.hidden_size;
528
529 let position_ids = num_positions;
530 let position_embedding = num_positions * cfg.hidden_size;
531
532 let conv2dconfig = Conv2dConfig {
533 stride: cfg.patch_size,
534 ..Default::default()
535 };
536 let patch_embedding =
537 cfg.num_channels * cfg.hidden_size / conv2dconfig.groups * cfg.patch_size * cfg.patch_size;
538
539 let encoder_layer_elems = {
540 let layer_norm1 = cfg.hidden_size;
541 let layer_norm2 = cfg.hidden_size;
542
543 let q_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
544 let k_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
545 let v_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
546 let o_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
547
548 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
549 let fc2 = cfg.intermediate_size * cfg.hidden_size + cfg.hidden_size;
550
551 layer_norm1 + layer_norm2 + q_proj + k_proj + v_proj + o_proj + fc1 + fc2
552 };
553
554 pre_layer_norm
555 + final_layer_norm
556 + class_embedding
557 + position_ids
558 + position_embedding
559 + patch_embedding
560 + cfg.num_hidden_layers * encoder_layer_elems
561}
562
563pub struct Phi3VLoader;
569
570pub struct Phi3VPrefixer;
571
572impl MultimodalPromptPrefixer for Phi3VPrefixer {
573 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
574 format!(
576 "{}{prompt}",
577 image_indexes
578 .into_iter()
579 .map(|image_index| format!("<|image_{}|>", image_index + 1))
580 .join("")
581 )
582 }
583}
584
585impl MultimodalModelLoader for Phi3VLoader {
586 fn load(
587 &self,
588 config: &str,
589 vb: ShardedVarBuilder,
590 normal_loading_metadata: NormalLoadingMetadata,
591 attention_mechanism: AttentionImplementation,
592 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
593 let cfg: crate::vision_models::phi3::Config = serde_json::from_str(config)?;
594 Ok(Box::new(Phi3::new(
595 &cfg,
596 vb,
597 self.is_gptx(config),
598 normal_loading_metadata,
599 attention_mechanism,
600 )?))
601 }
602 fn is_gptx(&self, _config: &str) -> bool {
603 true
604 }
605 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
606 let cfg: crate::vision_models::phi3::Config = serde_json::from_str(config)?;
607 Ok(Box::new(cfg))
608 }
609 fn get_processor(
610 &self,
611 _model_config: &str,
612 processor_config: Option<ProcessorConfig>,
613 preprocessor_config: PreProcessorConfig,
614 _max_edge: Option<u32>,
615 ) -> Arc<dyn Processor + Send + Sync> {
616 Phi3Processor::new_processor(processor_config, preprocessor_config)
617 }
618 fn supports_paged_attention(&self, _config: &str) -> bool {
619 true
620 }
621 fn supports_prefix_cacher(&self, _config: &str) -> bool {
622 true
623 }
624 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
625 Arc::new(Phi3VPrefixer)
626 }
627 fn modalities(&self, _config: &str) -> Result<Modalities> {
628 Ok(Modalities {
629 input: vec![SupportedModality::Text, SupportedModality::Vision],
630 output: vec![SupportedModality::Text],
631 })
632 }
633}
634
635impl IsqModelLoader for Phi3VLoader {
636 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
637 Ok(vec![
638 Regex::new(r"lm_head\.(weight|bias)$")?,
639 Regex::new(r"layers\.(\d+)\.self_attn\.qkv_proj\.(weight|bias)$")?,
641 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
642 Regex::new(r"layers\.(\d+)\.mlp\.gate_up_proj\.(weight|bias)$")?,
644 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
645 ])
646 }
647 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
648 self.isq_layer_regexes(config)
649 }
650}
651
652impl DeviceMappedModelLoader for Phi3VLoader {
653 fn mapped_max_act_size_elems(
654 &self,
655 config: &str,
656 params: &AutoDeviceMapParams,
657 ) -> Result<usize> {
658 let AutoDeviceMapParams::Multimodal {
660 max_seq_len,
661 max_batch_size,
662 max_image_shape: _,
663 max_num_images,
664 } = params
665 else {
666 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
667 };
668
669 let cfg: Phi3Config = serde_json::from_str(config)?;
670
671 let vcfg = &PHI3V_CLIP_CONFIG;
672
673 let num_patches = (vcfg.image_size / vcfg.patch_size).pow(2);
674 let img_seq_len = (num_patches + 1) * max_num_images;
675
676 let max_text_attn = {
677 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
679 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
680 };
681
682 Ok(max_text_attn)
683 }
684
685 fn non_mapped_max_act_size_elems(
686 &self,
687 config: &str,
688 params: &AutoDeviceMapParams,
689 ) -> Result<usize> {
690 let AutoDeviceMapParams::Multimodal {
692 max_seq_len: _,
693 max_batch_size,
694 max_image_shape: _,
695 max_num_images,
696 } = params
697 else {
698 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
699 };
700
701 let cfg: Phi3Config = serde_json::from_str(config)?;
702
703 let vcfg = &PHI3V_CLIP_CONFIG;
704
705 let num_patches = (vcfg.image_size / vcfg.patch_size).pow(2);
706 let img_seq_len = num_patches + 1;
707
708 let max_vision_attn = {
709 (max_batch_size * max_num_images) * cfg.num_attention_heads * img_seq_len * img_seq_len
710 };
711
712 Ok(max_vision_attn)
713 }
714
715 fn non_mapped_size_in_bytes(
716 &self,
717 config: &str,
718 dtype: DType,
719 weight_pack_factor: usize,
720 _matformer_config: Option<&MatformerSliceConfig>,
721 ) -> Result<usize> {
722 let cfg: Phi3Config = serde_json::from_str(config)?;
723 let elems = {
724 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
725 let lm_head = if !cfg.tie_word_embeddings || weight_pack_factor != 1 {
727 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
728 } else {
729 0
730 };
731 let norm = cfg.hidden_size;
732
733 let image_embed = {
734 let projection_cls = cfg
735 .embd_layer
736 .projection_cls
737 .clone()
738 .unwrap_or("linear".to_string());
739 let with_learnable_separator =
740 cfg.embd_layer.with_learnable_separator.unwrap_or(false);
741 let use_hd_transform = cfg.embd_layer.use_hd_transform.unwrap_or(false);
742 let image_dim_out = cfg.img_processor.image_dim_out;
743
744 let proj = match (projection_cls.as_str(), use_hd_transform) {
745 ("linear", _) => image_dim_out * cfg.hidden_size + cfg.hidden_size,
746 ("mlp", true) => {
747 let a = (image_dim_out * 4) * cfg.hidden_size + cfg.hidden_size;
748 let b = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
749 a + b
750 }
751 ("mlp", false) => {
752 let a = image_dim_out * cfg.hidden_size + cfg.hidden_size;
753 let b = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
754 a + b
755 }
756 _ => {
757 anyhow::bail!("projection_cls=`{projection_cls}` not implemented.");
758 }
759 };
760
761 let (glb_gn, sub_gn) = if with_learnable_separator {
762 let glb_gn = image_dim_out * 4;
763 let sub_gn = image_dim_out * 4;
764 (glb_gn, sub_gn)
765 } else {
766 (0, 0)
767 };
768
769 let clip_vit = get_clip_vit_num_elems(&PHI3V_CLIP_CONFIG);
770
771 proj + glb_gn + sub_gn + clip_vit
772 };
773
774 embed_tokens + lm_head + norm + image_embed
775 };
776
777 Ok(elems * dtype.size_in_bytes())
778 }
779
780 fn layer_sizes_in_bytes(
781 &self,
782 config: &str,
783 dtype: DType,
784 weight_pack_factor: usize,
785 _matformer_config: Option<&MatformerSliceConfig>,
786 ) -> Result<Vec<usize>> {
787 let cfg: Phi3Config = serde_json::from_str(config)?;
788 let per_layer_elems = {
789 let input_layernorm = cfg.hidden_size;
790 let post_attention_layernorm = cfg.hidden_size;
791
792 let size_in = cfg.hidden_size;
793 let head_dim = cfg.head_dim();
794 let op_size =
795 cfg.num_attention_heads * head_dim + 2 * cfg.num_key_value_heads * head_dim;
796 let qkv_proj = size_in * op_size / weight_pack_factor;
797 let o_proj = (cfg.num_attention_heads * head_dim) * size_in / weight_pack_factor;
798
799 let h_size = cfg.hidden_size;
800 let i_size = cfg.intermediate_size;
801 let gate_up_proj = h_size * (2 * i_size) / weight_pack_factor;
802 let down_proj = h_size * i_size / weight_pack_factor;
803
804 input_layernorm
805 + post_attention_layernorm
806 + qkv_proj
807 + o_proj
808 + gate_up_proj
809 + down_proj
810 };
811 Ok(vec![
812 per_layer_elems * dtype.size_in_bytes();
813 cfg.num_hidden_layers
814 ])
815 }
816
817 fn num_layers(&self, config: &str) -> Result<usize> {
818 let cfg: Phi3Config = serde_json::from_str(config)?;
819 Ok(cfg.num_hidden_layers)
820 }
821
822 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
823 let cfg: Phi3Config = serde_json::from_str(config)?;
824
825 let cfg = ModelConfigMetadata {
826 max_seq_len: cfg.max_position_embeddings,
827 num_layers: cfg.num_hidden_layers,
828 hidden_size: cfg.hidden_size,
829 num_kv_heads: cfg.num_key_value_heads,
830 num_attn_heads: cfg.num_attention_heads,
831 sliding_window: cfg.sliding_window,
832 k_head_dim: cfg.head_dim(),
833 v_head_dim: cfg.head_dim(),
834 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
835 };
836
837 Ok(Box::new(cfg))
838 }
839
840 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
841 Some(vec![NonMappedSubModel::Vision])
842 }
843}
844
845pub struct Idefics2Loader;
851
852pub struct Idefics2Prefixer;
853
854impl MultimodalPromptPrefixer for Idefics2Prefixer {
855 fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
856 prompt.to_string()
858 }
859}
860
861impl MultimodalModelLoader for Idefics2Loader {
862 fn load(
863 &self,
864 config: &str,
865 vb: ShardedVarBuilder,
866 normal_loading_metadata: NormalLoadingMetadata,
867 attention_mechanism: AttentionImplementation,
868 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
869 let cfg: crate::vision_models::idefics2::Config = serde_json::from_str(config)?;
870 Ok(Box::new(Idefics2::new(
871 &cfg,
872 vb,
873 self.is_gptx(config),
874 normal_loading_metadata,
875 attention_mechanism,
876 )?))
877 }
878 fn is_gptx(&self, _config: &str) -> bool {
879 true
880 }
881 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
882 let cfg: crate::vision_models::idefics2::Config = serde_json::from_str(config)?;
883 Ok(Box::new(cfg))
884 }
885 fn get_processor(
886 &self,
887 _model_config: &str,
888 processor_config: Option<ProcessorConfig>,
889 preprocessor_config: PreProcessorConfig,
890 max_edge: Option<u32>,
891 ) -> Arc<dyn Processor + Send + Sync> {
892 Arc::new(Idefics2Processor::new(
893 processor_config.unwrap(),
894 preprocessor_config,
895 max_edge,
896 ))
897 }
898 fn supports_paged_attention(&self, _config: &str) -> bool {
899 true
900 }
901 fn supports_prefix_cacher(&self, _config: &str) -> bool {
902 true
903 }
904 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
905 Arc::new(Idefics2Prefixer)
906 }
907 fn modalities(&self, _config: &str) -> Result<Modalities> {
908 Ok(Modalities {
909 input: vec![SupportedModality::Text, SupportedModality::Vision],
910 output: vec![SupportedModality::Text],
911 })
912 }
913}
914
915impl IsqModelLoader for Idefics2Loader {
916 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
917 Ok(vec![
918 Regex::new(r"lm_head\.(weight|bias)$")?,
919 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
921 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
922 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
923 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
924 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
926 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
927 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
928 ])
929 }
930 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
931 Ok(vec![
932 Regex::new(r"lm_head\.(weight|bias)$")?,
933 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
935 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
936 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
937 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
938 Regex::new(r"model\.text_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
940 Regex::new(r"model\.text_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
941 Regex::new(r"model\.text_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
942 ])
943 }
944}
945
946impl DeviceMappedModelLoader for Idefics2Loader {
947 fn mapped_max_act_size_elems(
948 &self,
949 config: &str,
950 params: &AutoDeviceMapParams,
951 ) -> Result<usize> {
952 let AutoDeviceMapParams::Multimodal {
953 max_seq_len,
954 max_batch_size,
955 max_image_shape: _,
956 max_num_images,
957 } = params
958 else {
959 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
960 };
961
962 let cfg: Idefics2Config = serde_json::from_str(config)?;
963
964 let num_patches = (cfg.vision_config.image_size / cfg.vision_config.patch_size).pow(2);
965 let img_seq_len = (num_patches + 1) * max_num_images;
966
967 let max_text_attn = {
968 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
970 max_batch_size * cfg.text_config.num_attention_heads * max_seq_len * max_seq_len
971 };
972
973 Ok(max_text_attn)
974 }
975
976 fn non_mapped_max_act_size_elems(
977 &self,
978 config: &str,
979 params: &AutoDeviceMapParams,
980 ) -> Result<usize> {
981 let AutoDeviceMapParams::Multimodal {
982 max_seq_len: _,
983 max_batch_size,
984 max_image_shape: _,
985 max_num_images,
986 } = params
987 else {
988 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
989 };
990
991 let cfg: Idefics2Config = serde_json::from_str(config)?;
992
993 let num_patches = (cfg.vision_config.image_size / cfg.vision_config.patch_size).pow(2);
994 let img_seq_len = num_patches + 1;
995
996 let max_vision_attn = {
997 let images_factor = 5;
999
1000 (max_batch_size * images_factor * max_num_images)
1001 * cfg.vision_config.num_attention_heads
1002 * img_seq_len
1003 * img_seq_len
1004 };
1005
1006 Ok(max_vision_attn)
1007 }
1008
1009 fn non_mapped_size_in_bytes(
1010 &self,
1011 config: &str,
1012 dtype: DType,
1013 weight_pack_factor: usize,
1014 _matformer_config: Option<&MatformerSliceConfig>,
1015 ) -> Result<usize> {
1016 let cfg: Idefics2Config = serde_json::from_str(config)?;
1017 let text_elems = {
1018 let tie_word_embeddings = cfg.tie_word_embeddings;
1019 let cfg = &cfg.text_config;
1020
1021 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
1022 let lm_head = if !tie_word_embeddings {
1023 cfg.hidden_size * cfg.vocab_size
1024 } else {
1025 0
1026 };
1027 let norm = cfg.hidden_size;
1028 embed_tokens + lm_head + norm
1029 };
1030
1031 let connector_elems = {
1032 let tcfg = &cfg.text_config;
1033 let vcfg = &cfg.vision_config;
1034 let gate_proj = vcfg.hidden_size * tcfg.intermediate_size;
1035 let up_proj = vcfg.hidden_size * tcfg.intermediate_size;
1036 let down_proj = tcfg.intermediate_size * tcfg.hidden_size;
1037
1038 let perceiver_elems = {
1039 let tcfg = &cfg.text_config;
1040 let pcfg = &cfg.perceiver_config;
1041
1042 let n_latents = pcfg.resampler_n_latents;
1043 let hidden_size = tcfg.hidden_size;
1044 let depth = pcfg.resampler_depth;
1045
1046 let norm = tcfg.hidden_size;
1047 let latents = n_latents * hidden_size;
1048
1049 let layer_elems = {
1050 let input_latents_norm = hidden_size;
1051 let input_context_norm = hidden_size;
1052 let post_attn_norm = hidden_size;
1053
1054 let num_heads = pcfg.resampler_n_heads;
1055 let head_dim = pcfg.resampler_head_dim;
1056 let num_key_value_heads = pcfg.num_key_value_heads;
1057
1058 let q_proj = hidden_size * num_heads * head_dim;
1059 let k_proj = hidden_size * num_key_value_heads * head_dim;
1060 let v_proj = hidden_size * num_key_value_heads * head_dim;
1061 let o_proj = num_heads * head_dim * hidden_size;
1062
1063 let gate_proj = hidden_size * hidden_size * 4;
1064 let up_proj = hidden_size * hidden_size * 4;
1065 let down_proj = hidden_size * 4 * hidden_size;
1066
1067 input_latents_norm
1068 + input_context_norm
1069 + post_attn_norm
1070 + q_proj
1071 + k_proj
1072 + v_proj
1073 + o_proj
1074 + gate_proj
1075 + up_proj
1076 + down_proj
1077 };
1078
1079 norm + latents + layer_elems * depth
1080 };
1081
1082 gate_proj + up_proj + down_proj + perceiver_elems
1083 };
1084
1085 let vision_transformer = {
1086 let cfg = &cfg.vision_config;
1087
1088 let post_layernorm = cfg.hidden_size;
1089
1090 let conv_config = Conv2dConfig {
1091 stride: cfg.patch_size,
1092 ..Default::default()
1093 };
1094 let patch_embedding = cfg.num_channels * cfg.hidden_size / conv_config.groups
1095 * cfg.patch_size
1096 * cfg.patch_size;
1097
1098 let num_patches_per_side = cfg.image_size / cfg.patch_size;
1099 let num_patches = num_patches_per_side.pow(2);
1100 let position_embedding = num_patches * cfg.hidden_size;
1101
1102 let layer_elems = {
1103 let layer_norm_1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
1104 let layer_norm_2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
1105
1106 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
1107 let fc2 = cfg.intermediate_size * cfg.hidden_size + cfg.hidden_size;
1108
1109 let q_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
1110 let k_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
1111 let v_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
1112 let o_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
1113
1114 layer_norm_1 + layer_norm_2 + fc1 + fc2 + q_proj + k_proj + v_proj + o_proj
1115 };
1116
1117 post_layernorm + patch_embedding + position_embedding + layer_elems
1118 };
1119
1120 let elems = text_elems + connector_elems + vision_transformer;
1121
1122 Ok(elems * dtype.size_in_bytes())
1123 }
1124
1125 fn layer_sizes_in_bytes(
1126 &self,
1127 config: &str,
1128 dtype: DType,
1129 weight_pack_factor: usize,
1130 _matformer_config: Option<&MatformerSliceConfig>,
1131 ) -> Result<Vec<usize>> {
1132 let cfg: Idefics2Config = serde_json::from_str(config)?;
1133 let cfg = cfg.text_config;
1134 let per_layer_elems = {
1135 let input_layernorm = cfg.hidden_size;
1136 let post_attention_layernorm = cfg.hidden_size;
1137
1138 let size_in = cfg.hidden_size;
1139 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
1140 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
1141 let q_proj = size_in * size_q / weight_pack_factor;
1142 let k_proj = size_in * size_kv / weight_pack_factor;
1143 let v_proj = size_in * size_kv / weight_pack_factor;
1144 let o_proj = size_q * size_in / weight_pack_factor;
1145
1146 let h_size = cfg.hidden_size;
1147 let i_size = cfg.intermediate_size;
1148 let gate_proj = h_size * i_size / weight_pack_factor;
1149 let up_proj = h_size * i_size / weight_pack_factor;
1150 let down_proj = i_size * h_size / weight_pack_factor;
1151
1152 input_layernorm
1153 + post_attention_layernorm
1154 + q_proj
1155 + k_proj
1156 + v_proj
1157 + o_proj
1158 + gate_proj
1159 + up_proj
1160 + down_proj
1161 };
1162 Ok(vec![
1163 per_layer_elems * dtype.size_in_bytes();
1164 cfg.num_hidden_layers
1165 ])
1166 }
1167
1168 fn num_layers(&self, config: &str) -> Result<usize> {
1169 let cfg: Idefics2Config = serde_json::from_str(config)?;
1170 Ok(cfg.text_config.num_hidden_layers)
1171 }
1172 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
1173 let cfg: Idefics2Config = serde_json::from_str(config)?;
1174 let cfg = &cfg.text_config;
1175
1176 let cfg = ModelConfigMetadata {
1177 max_seq_len: cfg.max_position_embeddings,
1178 num_layers: cfg.num_hidden_layers,
1179 hidden_size: cfg.hidden_size,
1180 num_kv_heads: cfg.num_key_value_heads,
1181 num_attn_heads: cfg.num_attention_heads,
1182 sliding_window: cfg.sliding_window,
1183 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
1184 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
1185 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
1186 };
1187
1188 Ok(Box::new(cfg))
1189 }
1190
1191 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
1192 Some(vec![NonMappedSubModel::Vision])
1193 }
1194}
1195
1196pub struct LLaVANextLoader;
1202
1203pub struct LLaVANextPrefixer;
1204
1205impl MultimodalPromptPrefixer for LLaVANextPrefixer {
1206 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
1207 format!("{}{prompt}", "<image>".repeat(image_indexes.len()))
1208 }
1209}
1210
1211impl MultimodalModelLoader for LLaVANextLoader {
1212 fn load(
1213 &self,
1214 config: &str,
1215 vb: ShardedVarBuilder,
1216 normal_loading_metadata: NormalLoadingMetadata,
1217 attention_mechanism: AttentionImplementation,
1218 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
1219 let cfg: crate::vision_models::llava::config::Config = serde_json::from_str(config)?;
1220 Ok(Box::new(LLaVANext::new(
1221 &cfg,
1222 vb,
1223 self.is_gptx(config),
1224 normal_loading_metadata,
1225 attention_mechanism,
1226 )?))
1227 }
1228 fn is_gptx(&self, _config: &str) -> bool {
1229 false
1230 }
1231 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
1232 let cfg: crate::vision_models::llava::config::Config = serde_json::from_str(config)?;
1233 Ok(Box::new(cfg))
1234 }
1235 fn get_processor(
1236 &self,
1237 model_config: &str,
1238 _processor_config: Option<ProcessorConfig>,
1239 _preprocessor_config: PreProcessorConfig,
1240 _max_edge: Option<u32>,
1241 ) -> Arc<dyn Processor + Send + Sync> {
1242 Arc::new(LLaVANextProcessor::new(model_config))
1243 }
1244 fn supports_paged_attention(&self, _config: &str) -> bool {
1245 true
1246 }
1247 fn supports_prefix_cacher(&self, _config: &str) -> bool {
1248 true
1249 }
1250 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
1251 Arc::new(LLaVANextPrefixer)
1252 }
1253 fn modalities(&self, _config: &str) -> Result<Modalities> {
1254 Ok(Modalities {
1255 input: vec![SupportedModality::Text, SupportedModality::Vision],
1256 output: vec![SupportedModality::Text],
1257 })
1258 }
1259}
1260
1261impl IsqModelLoader for LLaVANextLoader {
1262 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
1263 Ok(vec![
1264 Regex::new(r"lm_head\.(weight|bias)$")?,
1265 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
1267 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
1268 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
1269 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
1270 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
1272 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
1273 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
1274 ])
1275 }
1276 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
1277 Ok(vec![
1278 Regex::new(r"lm_head\.(weight|bias)$")?,
1279 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
1281 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
1282 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
1283 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
1284 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
1286 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
1287 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
1288 ])
1289 }
1290}
1291
1292impl DeviceMappedModelLoader for LLaVANextLoader {
1293 fn mapped_max_act_size_elems(
1294 &self,
1295 config: &str,
1296 params: &AutoDeviceMapParams,
1297 ) -> Result<usize> {
1298 let AutoDeviceMapParams::Multimodal {
1299 max_seq_len,
1300 max_batch_size,
1301 max_image_shape,
1302 max_num_images,
1303 } = params
1304 else {
1305 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
1306 };
1307
1308 let config: LLaVAConfig = serde_json::from_str(config)?;
1309
1310 #[allow(clippy::cast_possible_truncation)]
1311 let img_seq_len =
1312 llava_next_inputs_processor::LLaVANextInputProcessor::get_num_image_tokens(
1313 &config,
1314 (max_image_shape.0 as u32, max_image_shape.1 as u32),
1315 );
1316 let img_seq_len = img_seq_len * max_num_images;
1317
1318 let max_text_attn = {
1319 let cfg = &config.text_config;
1320 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
1322
1323 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
1324 };
1325
1326 Ok(max_text_attn)
1327 }
1328
1329 fn non_mapped_max_act_size_elems(
1330 &self,
1331 config: &str,
1332 params: &AutoDeviceMapParams,
1333 ) -> Result<usize> {
1334 let AutoDeviceMapParams::Multimodal {
1335 max_seq_len: _,
1336 max_batch_size,
1337 max_image_shape,
1338 max_num_images,
1339 } = params
1340 else {
1341 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
1342 };
1343
1344 let config: LLaVAConfig = serde_json::from_str(config)?;
1345
1346 #[allow(clippy::cast_possible_truncation)]
1347 let img_seq_len =
1348 llava_next_inputs_processor::LLaVANextInputProcessor::get_num_image_tokens(
1349 &config,
1350 (max_image_shape.0 as u32, max_image_shape.1 as u32),
1351 );
1352
1353 let max_vision_attn = {
1354 (max_batch_size * max_num_images)
1355 * config.vision_config.num_attention_heads
1356 * img_seq_len
1357 * img_seq_len
1358 };
1359
1360 Ok(max_vision_attn)
1361 }
1362
1363 fn non_mapped_size_in_bytes(
1364 &self,
1365 config: &str,
1366 dtype: DType,
1367 weight_pack_factor: usize,
1368 _matformer_config: Option<&MatformerSliceConfig>,
1369 ) -> Result<usize> {
1370 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1371 let text_elems = {
1372 let cfg = &cfg.text_config;
1373 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
1374 let lm_head = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
1375 let norm = cfg.hidden_size;
1376 embed_tokens + lm_head + norm
1377 };
1378
1379 let image_newline = cfg.text_config.hidden_size;
1380 let mmproj = {
1381 let linear_1 = cfg.vision_config.hidden_size * cfg.text_config.hidden_size
1382 + cfg.text_config.hidden_size;
1383 let linear_2 = cfg.text_config.hidden_size * cfg.text_config.hidden_size
1384 + cfg.text_config.hidden_size;
1385
1386 linear_1 + linear_2
1387 };
1388 let vision_tower = get_clip_vit_num_elems(&cfg.to_clip_config());
1389
1390 let elems = text_elems + image_newline + mmproj + vision_tower;
1391 Ok(elems * dtype.size_in_bytes())
1392 }
1393
1394 fn layer_sizes_in_bytes(
1395 &self,
1396 config: &str,
1397 dtype: DType,
1398 weight_pack_factor: usize,
1399 _matformer_config: Option<&MatformerSliceConfig>,
1400 ) -> Result<Vec<usize>> {
1401 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1402 let per_layer_elems = {
1403 let cfg = &cfg.text_config;
1404 let input_layernorm = cfg.hidden_size;
1405 let post_attention_layernorm = cfg.hidden_size;
1406
1407 let size_in = cfg.hidden_size;
1408 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
1409 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
1410 let q_proj = size_in * size_q / weight_pack_factor;
1411 let k_proj = size_in * size_kv / weight_pack_factor;
1412 let v_proj = size_in * size_kv / weight_pack_factor;
1413 let o_proj = size_q * size_in / weight_pack_factor;
1414
1415 let h_size = cfg.hidden_size;
1416 let i_size = cfg.intermediate_size;
1417 let gate_proj = h_size * i_size / weight_pack_factor;
1418 let up_proj = h_size * i_size / weight_pack_factor;
1419 let down_proj = i_size * h_size / weight_pack_factor;
1420
1421 input_layernorm
1422 + post_attention_layernorm
1423 + q_proj
1424 + k_proj
1425 + v_proj
1426 + o_proj
1427 + gate_proj
1428 + up_proj
1429 + down_proj
1430 };
1431 Ok(vec![
1432 per_layer_elems * dtype.size_in_bytes();
1433 cfg.text_config.num_hidden_layers
1434 ])
1435 }
1436
1437 fn num_layers(&self, config: &str) -> Result<usize> {
1438 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1439 Ok(cfg.text_config.num_hidden_layers)
1440 }
1441
1442 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
1443 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1444 let cfg = &cfg.text_config;
1445
1446 let cfg = ModelConfigMetadata {
1447 max_seq_len: cfg.max_position_embeddings,
1448 num_layers: cfg.num_hidden_layers,
1449 hidden_size: cfg.hidden_size,
1450 num_kv_heads: cfg.num_key_value_heads,
1451 num_attn_heads: cfg.num_attention_heads,
1452 sliding_window: cfg.sliding_window,
1453 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
1454 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
1455 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
1456 };
1457
1458 Ok(Box::new(cfg))
1459 }
1460
1461 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
1462 Some(vec![NonMappedSubModel::Vision])
1463 }
1464}
1465
1466pub struct LLaVALoader;
1472
1473pub struct LLaVAPrefixer;
1474
1475impl MultimodalPromptPrefixer for LLaVAPrefixer {
1476 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
1477 format!("{}{prompt}", "<image>".repeat(image_indexes.len()))
1478 }
1479}
1480
1481impl MultimodalModelLoader for LLaVALoader {
1482 fn load(
1483 &self,
1484 config: &str,
1485 vb: ShardedVarBuilder,
1486 normal_loading_metadata: NormalLoadingMetadata,
1487 attention_mechanism: AttentionImplementation,
1488 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
1489 let cfg: crate::vision_models::llava::config::Config = serde_json::from_str(config)?;
1490 Ok(Box::new(LLaVA::new(
1491 &cfg,
1492 vb,
1493 self.is_gptx(config),
1494 normal_loading_metadata,
1495 attention_mechanism,
1496 )?))
1497 }
1498 fn is_gptx(&self, _config: &str) -> bool {
1499 false
1500 }
1501 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
1502 let cfg: crate::vision_models::llava::config::Config = serde_json::from_str(config)?;
1503 Ok(Box::new(cfg))
1504 }
1505 fn get_processor(
1506 &self,
1507 model_config: &str,
1508 _processor_config: Option<ProcessorConfig>,
1509 _preprocessor_config: PreProcessorConfig,
1510 _max_edge: Option<u32>,
1511 ) -> Arc<dyn Processor + Send + Sync> {
1512 Arc::new(LLaVAProcessor::new(model_config))
1513 }
1514 fn supports_paged_attention(&self, _config: &str) -> bool {
1515 true
1516 }
1517 fn supports_prefix_cacher(&self, _config: &str) -> bool {
1518 true
1519 }
1520 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
1521 Arc::new(LLaVAPrefixer)
1522 }
1523 fn modalities(&self, _config: &str) -> Result<Modalities> {
1524 Ok(Modalities {
1525 input: vec![SupportedModality::Text, SupportedModality::Vision],
1526 output: vec![SupportedModality::Text],
1527 })
1528 }
1529}
1530
1531impl IsqModelLoader for LLaVALoader {
1532 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
1533 Ok(vec![
1534 Regex::new(r"lm_head\.(weight|bias)$")?,
1535 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
1537 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
1538 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
1539 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
1540 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
1542 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
1543 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
1544 ])
1545 }
1546 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
1547 Ok(vec![
1548 Regex::new(r"lm_head\.(weight|bias)$")?,
1549 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
1551 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
1552 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
1553 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
1554 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
1556 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
1557 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
1558 ])
1559 }
1560}
1561
1562impl DeviceMappedModelLoader for LLaVALoader {
1563 fn mapped_max_act_size_elems(
1564 &self,
1565 config: &str,
1566 params: &AutoDeviceMapParams,
1567 ) -> Result<usize> {
1568 let AutoDeviceMapParams::Multimodal {
1569 max_seq_len,
1570 max_batch_size,
1571 max_image_shape: _,
1572 max_num_images,
1573 } = params
1574 else {
1575 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
1576 };
1577
1578 let config: LLaVAConfig = serde_json::from_str(config)?;
1579
1580 let img_seq_len =
1581 llava_inputs_processor::LLaVAInputProcessor::get_num_image_tokens(&config);
1582 let img_seq_len = img_seq_len * max_num_images;
1583
1584 let max_text_attn = {
1585 let cfg = &config.text_config;
1586 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
1588
1589 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
1590 };
1591
1592 Ok(max_text_attn)
1593 }
1594
1595 fn non_mapped_max_act_size_elems(
1596 &self,
1597 config: &str,
1598 params: &AutoDeviceMapParams,
1599 ) -> Result<usize> {
1600 let AutoDeviceMapParams::Multimodal {
1601 max_seq_len: _,
1602 max_batch_size,
1603 max_image_shape: _,
1604 max_num_images,
1605 } = params
1606 else {
1607 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
1608 };
1609
1610 let config: LLaVAConfig = serde_json::from_str(config)?;
1611
1612 let img_seq_len =
1613 llava_inputs_processor::LLaVAInputProcessor::get_num_image_tokens(&config);
1614
1615 let max_vision_attn = {
1616 (max_batch_size * max_num_images)
1617 * config.vision_config.num_attention_heads
1618 * img_seq_len
1619 * img_seq_len
1620 };
1621
1622 Ok(max_vision_attn)
1623 }
1624
1625 fn non_mapped_size_in_bytes(
1626 &self,
1627 config: &str,
1628 dtype: DType,
1629 weight_pack_factor: usize,
1630 _matformer_config: Option<&MatformerSliceConfig>,
1631 ) -> Result<usize> {
1632 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1633 let text_elems = {
1634 let cfg = &cfg.text_config;
1635 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
1636 let lm_head = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
1637 let norm = cfg.hidden_size;
1638 embed_tokens + lm_head + norm
1639 };
1640
1641 let image_newline = cfg.text_config.hidden_size;
1642 let mmproj = {
1643 let linear_1 = cfg.vision_config.hidden_size * cfg.text_config.hidden_size
1644 + cfg.text_config.hidden_size;
1645 let linear_2 = cfg.text_config.hidden_size * cfg.text_config.hidden_size
1646 + cfg.text_config.hidden_size;
1647
1648 linear_1 + linear_2
1649 };
1650 let vision_tower = get_clip_vit_num_elems(&cfg.to_clip_config());
1651
1652 let elems = text_elems + image_newline + mmproj + vision_tower;
1653 Ok(elems * dtype.size_in_bytes())
1654 }
1655
1656 fn layer_sizes_in_bytes(
1657 &self,
1658 config: &str,
1659 dtype: DType,
1660 weight_pack_factor: usize,
1661 _matformer_config: Option<&MatformerSliceConfig>,
1662 ) -> Result<Vec<usize>> {
1663 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1664 let per_layer_elems = {
1665 let cfg = &cfg.text_config;
1666 let input_layernorm = cfg.hidden_size;
1667 let post_attention_layernorm = cfg.hidden_size;
1668
1669 let size_in = cfg.hidden_size;
1670 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
1671 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
1672 let q_proj = size_in * size_q / weight_pack_factor;
1673 let k_proj = size_in * size_kv / weight_pack_factor;
1674 let v_proj = size_in * size_kv / weight_pack_factor;
1675 let o_proj = size_q * size_in / weight_pack_factor;
1676
1677 let h_size = cfg.hidden_size;
1678 let i_size = cfg.intermediate_size;
1679 let gate_proj = h_size * i_size / weight_pack_factor;
1680 let up_proj = h_size * i_size / weight_pack_factor;
1681 let down_proj = i_size * h_size / weight_pack_factor;
1682
1683 input_layernorm
1684 + post_attention_layernorm
1685 + q_proj
1686 + k_proj
1687 + v_proj
1688 + o_proj
1689 + gate_proj
1690 + up_proj
1691 + down_proj
1692 };
1693 Ok(vec![
1694 per_layer_elems * dtype.size_in_bytes();
1695 cfg.text_config.num_hidden_layers
1696 ])
1697 }
1698
1699 fn num_layers(&self, config: &str) -> Result<usize> {
1700 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1701 Ok(cfg.text_config.num_hidden_layers)
1702 }
1703
1704 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
1705 let cfg: LLaVAConfig = serde_json::from_str(config)?;
1706 let cfg = &cfg.text_config;
1707
1708 let cfg = ModelConfigMetadata {
1709 max_seq_len: cfg.max_position_embeddings,
1710 num_layers: cfg.num_hidden_layers,
1711 hidden_size: cfg.hidden_size,
1712 num_kv_heads: cfg.num_key_value_heads,
1713 num_attn_heads: cfg.num_attention_heads,
1714 sliding_window: cfg.sliding_window,
1715 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
1716 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
1717 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
1718 };
1719
1720 Ok(Box::new(cfg))
1721 }
1722
1723 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
1724 Some(vec![NonMappedSubModel::Vision])
1725 }
1726}
1727
1728pub struct VLlamaLoader;
1734
1735pub struct VLlamaPrefixer;
1736
1737impl MultimodalPromptPrefixer for VLlamaPrefixer {
1738 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
1739 format!("{}{prompt}", "<|image|>".repeat(image_indexes.len()))
1740 }
1741}
1742
1743impl MultimodalModelLoader for VLlamaLoader {
1744 fn load(
1745 &self,
1746 config: &str,
1747 vb: ShardedVarBuilder,
1748 normal_loading_metadata: NormalLoadingMetadata,
1749 attention_mechanism: AttentionImplementation,
1750 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
1751 let cfg: crate::vision_models::mllama::MLlamaConfig = serde_json::from_str(config)?;
1752 Ok(Box::new(MLlamaModel::new(
1753 &cfg,
1754 vb,
1755 self.is_gptx(config),
1756 normal_loading_metadata,
1757 attention_mechanism,
1758 )?))
1759 }
1760 fn is_gptx(&self, _config: &str) -> bool {
1761 true
1762 }
1763 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
1764 let cfg: crate::vision_models::mllama::MLlamaConfig = serde_json::from_str(config)?;
1765 Ok(Box::new(cfg))
1766 }
1767 fn get_processor(
1768 &self,
1769 _model_config: &str,
1770 _processor_config: Option<ProcessorConfig>,
1771 _preprocessor_config: PreProcessorConfig,
1772 _max_edge: Option<u32>,
1773 ) -> Arc<dyn Processor + Send + Sync> {
1774 Arc::new(MLlamaProcessor::new())
1775 }
1776 fn supports_paged_attention(&self, _config: &str) -> bool {
1777 true
1778 }
1779 fn supports_prefix_cacher(&self, _config: &str) -> bool {
1780 true
1781 }
1782 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
1783 Arc::new(VLlamaPrefixer)
1784 }
1785 fn modalities(&self, _config: &str) -> Result<Modalities> {
1786 Ok(Modalities {
1787 input: vec![SupportedModality::Text, SupportedModality::Vision],
1788 output: vec![SupportedModality::Text],
1789 })
1790 }
1791}
1792
1793impl IsqModelLoader for VLlamaLoader {
1794 fn isq_layer_regexes(&self, config: &str) -> Result<Vec<Regex>> {
1795 let config: MLlamaConfig = serde_json::from_str(config)?;
1796 let cross_attn_layers = &config.text_config.cross_attention_layers;
1797 let transformer_layers =
1798 (0..config.text_config.num_hidden_layers).filter(|i| !cross_attn_layers.contains(i));
1799 let mut text_regexes = Vec::new();
1800 for layer in transformer_layers {
1801 text_regexes.extend(vec![
1802 Regex::new(&format!(
1804 r"language_model.model.layers\.{layer}\.self_attn\.q_proj\.(weight|bias)$"
1805 ))?,
1806 Regex::new(&format!(
1807 r"language_model.model.layers\.{layer}\.self_attn\.k_proj\.(weight|bias)$"
1808 ))?,
1809 Regex::new(&format!(
1810 r"language_model.model.layers\.{layer}\.self_attn\.v_proj\.(weight|bias)$"
1811 ))?,
1812 Regex::new(&format!(
1813 r"language_model.model.layers\.{layer}\.self_attn\.o_proj\.(weight|bias)$"
1814 ))?,
1815 Regex::new(&format!(
1817 r"language_model.model.layers\.{layer}\.mlp\.gate_proj\.(weight|bias)$"
1818 ))?,
1819 Regex::new(&format!(
1820 r"language_model.model.layers\.{layer}\.mlp\.up_proj\.(weight|bias)$"
1821 ))?,
1822 Regex::new(&format!(
1823 r"language_model.model.layers\.{layer}\.mlp\.down_proj\.(weight|bias)$"
1824 ))?,
1825 ]);
1826 }
1827 let vision_regexes = vec![
1828 Regex::new(
1830 r"vision_model.transformer.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$",
1831 )?,
1832 Regex::new(
1833 r"vision_model.transformer.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$",
1834 )?,
1835 Regex::new(
1836 r"vision_model.transformer.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$",
1837 )?,
1838 Regex::new(
1839 r"vision_model.transformer.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$",
1840 )?,
1841 Regex::new(
1843 r"vision_model.global_transformer.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$",
1844 )?,
1845 Regex::new(
1846 r"vision_model.global_transformer.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$",
1847 )?,
1848 Regex::new(
1849 r"vision_model.global_transformer.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$",
1850 )?,
1851 Regex::new(
1852 r"vision_model.global_transformer.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$",
1853 )?,
1854 Regex::new(r"layers\.(\d+)\.mlp\.fc1\.(weight|bias)$")?,
1856 Regex::new(r"layers\.(\d+)\.mlp\.fc2\.(weight|bias)$")?,
1857 ];
1858
1859 Ok([text_regexes, vision_regexes].concat())
1860 }
1861 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
1862 self.isq_layer_regexes(config)
1863 }
1864}
1865
1866impl DeviceMappedModelLoader for VLlamaLoader {
1867 fn mapped_max_act_size_elems(
1868 &self,
1869 config: &str,
1870 params: &AutoDeviceMapParams,
1871 ) -> Result<usize> {
1872 let AutoDeviceMapParams::Multimodal {
1873 max_seq_len,
1874 max_batch_size,
1875 max_image_shape: _,
1876 max_num_images,
1877 } = params
1878 else {
1879 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
1880 };
1881
1882 let config: MLlamaConfig = serde_json::from_str(config)?;
1883
1884 let img_seq_len = {
1885 let cfg = &config.vision_config;
1886 let num_patches = (cfg.image_size / cfg.patch_size).pow(2) + 1;
1887 let num_padding_patches = (8 - (num_patches as isize % 8)) % 8;
1888 cfg.max_num_tiles * (num_patches as isize + num_padding_patches) as usize
1889 };
1890 let img_seq_len = img_seq_len * max_num_images;
1891
1892 let max_cross_text_attn = {
1893 let cfg = &config.text_config;
1894 max_batch_size * cfg.num_attention_heads * img_seq_len * img_seq_len
1895 };
1896
1897 let max_self_text_attn = {
1898 let cfg = &config.text_config;
1899 max_batch_size * cfg.num_attention_heads * max_seq_len.min(&ATTENTION_CHUNK_SIZE).pow(2)
1900 };
1901
1902 Ok(max_self_text_attn.max(max_cross_text_attn))
1903 }
1904
1905 fn non_mapped_max_act_size_elems(
1906 &self,
1907 config: &str,
1908 params: &AutoDeviceMapParams,
1909 ) -> Result<usize> {
1910 let AutoDeviceMapParams::Multimodal {
1911 max_seq_len: _,
1912 max_batch_size,
1913 max_image_shape: _,
1914 max_num_images,
1915 } = params
1916 else {
1917 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
1918 };
1919
1920 let config: MLlamaConfig = serde_json::from_str(config)?;
1921
1922 let img_seq_len = {
1923 let cfg = &config.vision_config;
1924 let num_patches = (cfg.image_size / cfg.patch_size).pow(2) + 1;
1925 let num_padding_patches = (8 - (num_patches as isize % 8)) % 8;
1926 cfg.max_num_tiles * (num_patches as isize + num_padding_patches) as usize
1927 };
1928 let max_vision_attn = {
1929 let cfg = &config.vision_config;
1930 (max_batch_size * max_num_images) * cfg.num_attention_heads * img_seq_len * img_seq_len
1931 };
1932
1933 Ok(max_vision_attn)
1934 }
1935
1936 fn non_mapped_size_in_bytes(
1937 &self,
1938 config: &str,
1939 dtype: DType,
1940 weight_pack_factor: usize,
1941 _matformer_config: Option<&MatformerSliceConfig>,
1942 ) -> Result<usize> {
1943 let config: MLlamaConfig = serde_json::from_str(config)?;
1944 let text_elems = {
1945 let cfg = &config.text_config;
1946 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
1947 let lm_head = if !cfg.tie_word_embeddings || weight_pack_factor != 1 {
1949 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
1950 } else {
1951 0
1952 };
1953 let norm = cfg.hidden_size;
1954 embed_tokens + lm_head + norm
1955 };
1956
1957 let vision_elems = {
1958 let cfg = &config.vision_config;
1959
1960 let conv_cfg = Conv2dConfig {
1961 stride: cfg.patch_size,
1962 ..Default::default()
1963 };
1964 let patch_embedding = cfg.num_channels * cfg.hidden_size / conv_cfg.groups
1965 * cfg.patch_size
1966 * cfg.patch_size;
1967
1968 let class_embedding = cfg.hidden_size;
1969
1970 let gated_positional_embedding = {
1971 let num_patches = (cfg.image_size / cfg.patch_size).pow(2) + 1;
1972 let embedding = num_patches * cfg.hidden_size;
1973 let tile_embedding = (cfg.max_aspect_ratio_id() + 1)
1974 * (cfg.max_num_tiles * num_patches * cfg.hidden_size);
1975
1976 embedding + tile_embedding
1977 };
1978
1979 let pre_tile_positional_embedding =
1980 (cfg.max_aspect_ratio_id() + 1) * (cfg.max_num_tiles * cfg.hidden_size);
1981 let post_tile_positional_embedding =
1982 (cfg.max_aspect_ratio_id() + 1) * (cfg.max_num_tiles * cfg.hidden_size);
1983
1984 let layernorm_pre = cfg.hidden_size;
1985 let layernorm_post = cfg.hidden_size;
1986
1987 let encoder_layer = {
1988 let input_layernorm = cfg.hidden_size + cfg.hidden_size;
1989 let post_attention_layernorm = cfg.hidden_size + cfg.hidden_size;
1990
1991 let head_dim = cfg.hidden_size / cfg.num_attention_heads;
1992 let q_proj =
1993 cfg.hidden_size * cfg.num_attention_heads * head_dim / weight_pack_factor;
1994 let k_proj =
1995 cfg.hidden_size * cfg.num_attention_heads * head_dim / weight_pack_factor;
1996 let v_proj =
1997 cfg.hidden_size * cfg.num_attention_heads * head_dim / weight_pack_factor;
1998 let o_proj =
1999 cfg.hidden_size * cfg.num_attention_heads * head_dim / weight_pack_factor;
2000
2001 let fc1 = (cfg.hidden_size * cfg.intermediate_size) / weight_pack_factor
2002 + cfg.intermediate_size;
2003 let fc2 = (cfg.intermediate_size * cfg.hidden_size) / weight_pack_factor
2004 + cfg.hidden_size;
2005
2006 input_layernorm
2007 + post_attention_layernorm
2008 + q_proj
2009 + k_proj
2010 + v_proj
2011 + o_proj
2012 + fc1
2013 + fc2
2014 };
2015
2016 patch_embedding
2017 + class_embedding
2018 + gated_positional_embedding
2019 + pre_tile_positional_embedding
2020 + post_tile_positional_embedding
2021 + layernorm_pre
2022 + layernorm_post
2023 + encoder_layer * (cfg.num_hidden_layers + cfg.num_global_layers)
2024 };
2025
2026 let elems = text_elems + vision_elems;
2027 Ok(elems * dtype.size_in_bytes())
2028 }
2029
2030 fn layer_sizes_in_bytes(
2031 &self,
2032 config: &str,
2033 dtype: DType,
2034 weight_pack_factor: usize,
2035 _matformer_config: Option<&MatformerSliceConfig>,
2036 ) -> Result<Vec<usize>> {
2037 let config: MLlamaConfig = serde_json::from_str(config)?;
2038 let cfg = &config.text_config;
2039
2040 let mut layer_sizes = Vec::new();
2041
2042 for i in 0..cfg.num_hidden_layers {
2043 let weight_pack_factor = if cfg.cross_attention_layers.contains(&i) {
2044 1
2046 } else {
2047 weight_pack_factor
2048 };
2049
2050 let per_layer_elems = {
2051 let input_layernorm = cfg.hidden_size;
2052 let post_attention_layernorm = cfg.hidden_size;
2053
2054 let size_in = cfg.hidden_size;
2055 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
2056 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
2057 let q_proj = size_in * size_q / weight_pack_factor;
2058 let k_proj = size_in * size_kv / weight_pack_factor;
2059 let v_proj = size_in * size_kv / weight_pack_factor;
2060 let o_proj = size_q * size_in / weight_pack_factor;
2061
2062 let h_size = cfg.hidden_size;
2063 let i_size = cfg.intermediate_size;
2064 let gate_proj = h_size * i_size / weight_pack_factor;
2065 let up_proj = h_size * i_size / weight_pack_factor;
2066 let down_proj = i_size * h_size / weight_pack_factor;
2067
2068 input_layernorm
2069 + post_attention_layernorm
2070 + q_proj
2071 + k_proj
2072 + v_proj
2073 + o_proj
2074 + gate_proj
2075 + up_proj
2076 + down_proj
2077 };
2078
2079 layer_sizes.push(per_layer_elems * dtype.size_in_bytes());
2080 }
2081
2082 Ok(layer_sizes)
2083 }
2084
2085 fn num_layers(&self, config: &str) -> Result<usize> {
2086 let config: MLlamaConfig = serde_json::from_str(config)?;
2087 Ok(config.text_config.num_hidden_layers)
2088 }
2089
2090 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
2091 let cfg: MLlamaConfig = serde_json::from_str(config)?;
2092 let cfg = &cfg.text_config;
2093
2094 let cfg = ModelConfigMetadata {
2095 max_seq_len: cfg.max_position_embeddings,
2096 num_layers: cfg.num_hidden_layers,
2097 hidden_size: cfg.hidden_size,
2098 num_kv_heads: cfg.num_key_value_heads,
2099 num_attn_heads: cfg.num_attention_heads,
2100 sliding_window: None,
2101 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
2102 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
2103 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
2104 };
2105
2106 Ok(Box::new(cfg))
2107 }
2108
2109 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
2110 Some(vec![NonMappedSubModel::Vision])
2111 }
2112}
2113
2114pub struct Qwen2VLLoader;
2120
2121pub struct Qwen2VLPrefixer;
2122
2123impl MultimodalPromptPrefixer for Qwen2VLPrefixer {
2124 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
2125 format!(
2126 "{}{prompt}",
2127 format!(
2128 "{}{}{}",
2129 Qwen2VLProcessor::VISION_START,
2130 Qwen2VLProcessor::IMAGE_PAD,
2131 Qwen2VLProcessor::VISION_END
2132 )
2133 .repeat(image_indexes.len())
2134 )
2135 }
2136}
2137
2138impl MultimodalModelLoader for Qwen2VLLoader {
2139 fn load(
2140 &self,
2141 config: &str,
2142 vb: ShardedVarBuilder,
2143 normal_loading_metadata: NormalLoadingMetadata,
2144 attention_mechanism: AttentionImplementation,
2145 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
2146 let cfg: Qwen2VLConfig = serde_json::from_str(config)?;
2147 Ok(Box::new(Qwen2VLModel::new(
2148 &cfg,
2149 vb,
2150 self.is_gptx(config),
2151 normal_loading_metadata,
2152 attention_mechanism,
2153 )?))
2154 }
2155 fn is_gptx(&self, _config: &str) -> bool {
2156 true
2157 }
2158 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
2159 let config: Qwen2VLConfig = serde_json::from_str(config)?;
2160 Ok(Box::new(config))
2161 }
2162 fn get_processor(
2163 &self,
2164 _model_config: &str,
2165 _processor_config: Option<ProcessorConfig>,
2166 _preprocessor_config: PreProcessorConfig,
2167 max_edge: Option<u32>,
2168 ) -> Arc<dyn Processor + Send + Sync> {
2169 Arc::new(Qwen2VLProcessor::new(max_edge))
2170 }
2171 fn supports_paged_attention(&self, _config: &str) -> bool {
2172 false
2173 }
2174 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
2175 Arc::new(Qwen2VLPrefixer)
2176 }
2177 fn modalities(&self, _config: &str) -> Result<Modalities> {
2178 Ok(Modalities {
2179 input: vec![SupportedModality::Text, SupportedModality::Vision],
2180 output: vec![SupportedModality::Text],
2181 })
2182 }
2183}
2184
2185impl IsqModelLoader for Qwen2VLLoader {
2186 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
2187 Ok(vec![
2188 Regex::new(r"lm_head\.(weight|bias)$")?,
2189 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
2191 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
2192 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
2193 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
2194 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
2196 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
2197 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
2198 ])
2199 }
2200 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
2201 self.isq_layer_regexes(config)
2202 }
2203}
2204
2205impl DeviceMappedModelLoader for Qwen2VLLoader {
2206 fn mapped_max_act_size_elems(
2207 &self,
2208 config: &str,
2209 params: &AutoDeviceMapParams,
2210 ) -> Result<usize> {
2211 let AutoDeviceMapParams::Multimodal {
2212 max_seq_len,
2213 max_batch_size,
2214 max_image_shape,
2215 max_num_images,
2216 } = params
2217 else {
2218 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
2219 };
2220
2221 let cfg: Qwen2VLConfig = serde_json::from_str(config)?;
2222
2223 let img_seq_len = {
2225 let cfg = &cfg.vision_config;
2226 let grid_t = 1;
2228 let grid_h = (max_image_shape.0 / cfg.patch_size) / cfg.spatial_merge_size;
2230 let grid_w = (max_image_shape.1 / cfg.patch_size) / cfg.spatial_merge_size;
2231 grid_t * grid_h * grid_w * max_num_images
2232 };
2233
2234 let max_text_attn = {
2235 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
2237 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
2238 };
2239
2240 Ok(max_text_attn)
2241 }
2242
2243 fn non_mapped_max_act_size_elems(
2244 &self,
2245 config: &str,
2246 params: &AutoDeviceMapParams,
2247 ) -> Result<usize> {
2248 let AutoDeviceMapParams::Multimodal {
2249 max_seq_len: _,
2250 max_batch_size,
2251 max_image_shape,
2252 max_num_images,
2253 } = params
2254 else {
2255 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
2256 };
2257
2258 let cfg: Qwen2VLConfig = serde_json::from_str(config)?;
2259
2260 let img_seq_len = {
2262 let cfg = &cfg.vision_config;
2263 let grid_t = 1;
2265 let grid_h = max_image_shape.0 / cfg.patch_size;
2266 let grid_w = max_image_shape.1 / cfg.patch_size;
2267 grid_t * grid_h * grid_w
2268 };
2269
2270 let max_vision_attn = {
2271 let cfg = &cfg.vision_config;
2272 (max_batch_size * max_num_images) * cfg.num_heads * img_seq_len * img_seq_len
2273 };
2274
2275 Ok(max_vision_attn)
2276 }
2277
2278 fn non_mapped_size_in_bytes(
2279 &self,
2280 config: &str,
2281 dtype: DType,
2282 weight_pack_factor: usize,
2283 _matformer_config: Option<&MatformerSliceConfig>,
2284 ) -> Result<usize> {
2285 let cfg: Qwen2VLConfig = serde_json::from_str(config)?;
2286 let text_elems = {
2287 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
2288 let lm_head = if !cfg.tie_word_embeddings || weight_pack_factor != 1 {
2290 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
2291 } else {
2292 0
2293 };
2294 let norm = cfg.hidden_size;
2295 embed_tokens + lm_head + norm
2296 };
2297
2298 let patch_merger = {
2299 let cfg = &cfg.vision_config;
2300 let hidden_size = cfg.embed_dim * cfg.spatial_merge_size.pow(2);
2301
2302 let mlp0 = hidden_size * hidden_size + hidden_size;
2303 let mlp2 = hidden_size * cfg.hidden_size + cfg.hidden_size;
2304
2305 let ln_q = cfg.embed_dim + bias_if!(true, cfg.embed_dim);
2306
2307 mlp0 + mlp2 + ln_q
2308 };
2309
2310 let patch_embed = {
2311 let cfg = &cfg.vision_config;
2312 let conv_cfg = Conv3dConfig {
2313 stride: cfg.patch_size,
2314 ..Default::default()
2315 };
2316 let kernel_sizes = [cfg.temporal_patch_size, cfg.patch_size, cfg.patch_size];
2317 cfg.in_channels * cfg.embed_dim / conv_cfg.groups
2318 * kernel_sizes[0]
2319 * kernel_sizes[1]
2320 * kernel_sizes[2]
2321 };
2322
2323 let encoder_layer = {
2324 let cfg = &cfg.vision_config;
2325 let norm1 = cfg.embed_dim + bias_if!(true, cfg.embed_dim);
2326 let norm2 = cfg.embed_dim + bias_if!(true, cfg.embed_dim);
2327
2328 #[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
2329 let mlp_hidden_dim = (cfg.embed_dim as f64 * cfg.mlp_ratio) as usize;
2330 let fc1 = cfg.embed_dim * mlp_hidden_dim + mlp_hidden_dim;
2331 let fc2 = cfg.embed_dim * mlp_hidden_dim + cfg.embed_dim;
2332
2333 let qkv = cfg.embed_dim * cfg.embed_dim * 3 + cfg.embed_dim * 3;
2334 let out = cfg.embed_dim * cfg.embed_dim + cfg.embed_dim;
2335
2336 norm1 + norm2 + fc1 + fc2 + qkv + out
2337 };
2338
2339 let elems =
2340 text_elems + patch_merger + patch_embed + encoder_layer * cfg.vision_config.depth;
2341
2342 Ok(elems * dtype.size_in_bytes())
2343 }
2344
2345 fn layer_sizes_in_bytes(
2346 &self,
2347 config: &str,
2348 dtype: DType,
2349 weight_pack_factor: usize,
2350 _matformer_config: Option<&MatformerSliceConfig>,
2351 ) -> Result<Vec<usize>> {
2352 let cfg: Qwen2VLConfig = serde_json::from_str(config)?;
2353 let per_layer_elems = {
2354 let input_layernorm = cfg.hidden_size;
2355 let post_attention_layernorm = cfg.hidden_size;
2356
2357 let size_in = cfg.hidden_size;
2358 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
2359 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
2360 let q_proj = size_in * size_q / weight_pack_factor + size_q;
2361 let k_proj = size_in * size_kv / weight_pack_factor + size_kv;
2362 let v_proj = size_in * size_kv / weight_pack_factor + size_kv;
2363 let o_proj = size_q * size_in / weight_pack_factor;
2364
2365 let h_size = cfg.hidden_size;
2366 let i_size = cfg.intermediate_size;
2367 let gate_proj = h_size * i_size / weight_pack_factor;
2368 let up_proj = h_size * i_size / weight_pack_factor;
2369 let down_proj = i_size * h_size / weight_pack_factor;
2370
2371 input_layernorm
2372 + post_attention_layernorm
2373 + q_proj
2374 + k_proj
2375 + v_proj
2376 + o_proj
2377 + gate_proj
2378 + up_proj
2379 + down_proj
2380 };
2381 Ok(vec![
2382 per_layer_elems * dtype.size_in_bytes();
2383 cfg.num_hidden_layers
2384 ])
2385 }
2386
2387 fn num_layers(&self, config: &str) -> Result<usize> {
2388 let cfg: Qwen2VLConfig = serde_json::from_str(config)?;
2389 Ok(cfg.num_hidden_layers)
2390 }
2391
2392 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
2393 let cfg: Qwen2VLConfig = serde_json::from_str(config)?;
2394
2395 let cfg = ModelConfigMetadata {
2396 max_seq_len: cfg.max_position_embeddings,
2397 num_layers: cfg.num_hidden_layers,
2398 hidden_size: cfg.hidden_size,
2399 num_kv_heads: cfg.num_key_value_heads,
2400 num_attn_heads: cfg.num_attention_heads,
2401 sliding_window: cfg.sliding_window,
2402 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
2403 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
2404 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
2405 };
2406
2407 Ok(Box::new(cfg))
2408 }
2409
2410 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
2411 Some(vec![NonMappedSubModel::Vision])
2412 }
2413}
2414
2415pub struct Idefics3Loader;
2421
2422pub struct Idefics3Prefixer;
2423
2424impl MultimodalPromptPrefixer for Idefics3Prefixer {
2425 fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
2426 prompt.to_string()
2428 }
2429}
2430
2431impl MultimodalModelLoader for Idefics3Loader {
2432 fn load(
2433 &self,
2434 config: &str,
2435 vb: ShardedVarBuilder,
2436 normal_loading_metadata: NormalLoadingMetadata,
2437 attention_mechanism: AttentionImplementation,
2438 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
2439 let cfg: crate::vision_models::idefics3::Idefics3Config = serde_json::from_str(config)?;
2440 Ok(Box::new(Idefics3Model::new(
2441 &cfg,
2442 vb,
2443 self.is_gptx(config),
2444 normal_loading_metadata,
2445 attention_mechanism,
2446 )?))
2447 }
2448 fn is_gptx(&self, _config: &str) -> bool {
2449 true
2450 }
2451 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
2452 let cfg: crate::vision_models::idefics3::Idefics3Config = serde_json::from_str(config)?;
2453 Ok(Box::new(cfg))
2454 }
2455 fn get_processor(
2456 &self,
2457 _model_config: &str,
2458 processor_config: Option<ProcessorConfig>,
2459 preprocessor_config: PreProcessorConfig,
2460 max_edge: Option<u32>,
2461 ) -> Arc<dyn Processor + Send + Sync> {
2462 Arc::new(Idefics3Processor::new(
2463 processor_config.unwrap_or_default(),
2464 preprocessor_config,
2465 max_edge,
2466 ))
2467 }
2468 fn supports_paged_attention(&self, _config: &str) -> bool {
2469 true
2470 }
2471 fn supports_prefix_cacher(&self, _config: &str) -> bool {
2472 true
2473 }
2474 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
2475 Arc::new(Idefics3Prefixer)
2476 }
2477 fn modalities(&self, _config: &str) -> Result<Modalities> {
2478 Ok(Modalities {
2479 input: vec![SupportedModality::Text, SupportedModality::Vision],
2480 output: vec![SupportedModality::Text],
2481 })
2482 }
2483}
2484
2485impl IsqModelLoader for Idefics3Loader {
2486 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
2487 Ok(vec![
2488 Regex::new(r"lm_head\.(weight|bias)$")?,
2489 Regex::new(r"model.text_model.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
2491 Regex::new(r"model.text_model.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
2492 Regex::new(r"model.text_model.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
2493 Regex::new(r"model.text_model.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
2494 Regex::new(r"model.text_model.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
2496 Regex::new(r"model.text_model.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
2497 Regex::new(r"model.text_model.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
2498 ])
2499 }
2500 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
2501 Ok(vec![
2502 Regex::new(r"lm_head\.(weight|bias)$")?,
2503 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
2505 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
2506 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
2507 Regex::new(r"model\.text_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
2508 Regex::new(r"model\.text_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
2510 Regex::new(r"model\.text_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
2511 Regex::new(r"model\.text_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
2512 ])
2529 }
2530}
2531
2532impl DeviceMappedModelLoader for Idefics3Loader {
2533 fn mapped_max_act_size_elems(
2534 &self,
2535 config: &str,
2536 params: &AutoDeviceMapParams,
2537 ) -> Result<usize> {
2538 let AutoDeviceMapParams::Multimodal {
2539 max_seq_len,
2540 max_batch_size,
2541 max_image_shape: _,
2542 max_num_images,
2543 } = params
2544 else {
2545 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
2546 };
2547
2548 let cfg: Idefics3Config = serde_json::from_str(config)?;
2549
2550 let num_patches = (cfg.vision_config.image_size / cfg.vision_config.patch_size).pow(2);
2551 let img_seq_len = (num_patches + 1) * max_num_images;
2552
2553 let max_text_attn = {
2554 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
2556 max_batch_size * cfg.text_config.num_attention_heads * max_seq_len * max_seq_len
2557 };
2558
2559 Ok(max_text_attn)
2560 }
2561
2562 fn non_mapped_max_act_size_elems(
2563 &self,
2564 config: &str,
2565 params: &AutoDeviceMapParams,
2566 ) -> Result<usize> {
2567 let AutoDeviceMapParams::Multimodal {
2568 max_seq_len: _,
2569 max_batch_size,
2570 max_image_shape: _,
2571 max_num_images,
2572 } = params
2573 else {
2574 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
2575 };
2576
2577 let cfg: Idefics3Config = serde_json::from_str(config)?;
2578
2579 let num_patches = (cfg.vision_config.image_size / cfg.vision_config.patch_size).pow(2);
2580 let img_seq_len = num_patches + 1;
2581
2582 let max_vision_attn = {
2583 let images_factor = 5;
2585
2586 (max_batch_size * images_factor * max_num_images)
2587 * cfg.vision_config.num_attention_heads
2588 * img_seq_len
2589 * img_seq_len
2590 };
2591
2592 Ok(max_vision_attn)
2593 }
2594
2595 fn non_mapped_size_in_bytes(
2596 &self,
2597 config: &str,
2598 dtype: DType,
2599 weight_pack_factor: usize,
2600 _matformer_config: Option<&MatformerSliceConfig>,
2601 ) -> Result<usize> {
2602 let cfg: Idefics3Config = serde_json::from_str(config)?;
2603 let text_elems = {
2604 let cfg = &cfg.text_config;
2605
2606 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
2607 let lm_head = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
2608 let norm = cfg.hidden_size;
2609 embed_tokens + lm_head + norm
2610 };
2611
2612 let connector_elems = {
2613 let in_dim = cfg.vision_config.hidden_size * cfg.scale_factor.pow(2);
2614 let out_dim = cfg.text_config.hidden_size;
2615
2616 in_dim * out_dim
2617 };
2618
2619 let vision_transformer = {
2620 let cfg = &cfg.vision_config;
2621
2622 let post_layernorm = cfg.hidden_size;
2623
2624 let conv_config = Conv2dConfig {
2625 stride: cfg.patch_size,
2626 ..Default::default()
2627 };
2628 let patch_embedding = cfg.num_channels * cfg.hidden_size / conv_config.groups
2629 * cfg.patch_size
2630 * cfg.patch_size;
2631
2632 let num_patches_per_side = cfg.image_size / cfg.patch_size;
2633 let num_patches = num_patches_per_side.pow(2);
2634 let position_embedding = num_patches * cfg.hidden_size;
2635
2636 let layer_elems = {
2637 let layer_norm_1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
2638 let layer_norm_2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
2639
2640 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
2641 let fc2 = cfg.intermediate_size * cfg.hidden_size + cfg.hidden_size;
2642
2643 let q_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2644 let k_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2645 let v_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2646 let o_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2647
2648 layer_norm_1 + layer_norm_2 + fc1 + fc2 + q_proj + k_proj + v_proj + o_proj
2649 };
2650
2651 post_layernorm
2652 + patch_embedding
2653 + position_embedding
2654 + layer_elems * cfg.num_hidden_layers
2655 };
2656
2657 let elems = text_elems + connector_elems + vision_transformer;
2658
2659 Ok(elems * dtype.size_in_bytes())
2660 }
2661
2662 fn layer_sizes_in_bytes(
2663 &self,
2664 config: &str,
2665 dtype: DType,
2666 weight_pack_factor: usize,
2667 _matformer_config: Option<&MatformerSliceConfig>,
2668 ) -> Result<Vec<usize>> {
2669 let cfg: Idefics3Config = serde_json::from_str(config)?;
2670 let cfg = cfg.text_config;
2671 let per_layer_elems = {
2672 let input_layernorm = cfg.hidden_size;
2673 let post_attention_layernorm = cfg.hidden_size;
2674
2675 let size_in = cfg.hidden_size;
2676 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
2677 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
2678 let q_proj = size_in * size_q / weight_pack_factor;
2679 let k_proj = size_in * size_kv / weight_pack_factor;
2680 let v_proj = size_in * size_kv / weight_pack_factor;
2681 let o_proj = size_q * size_in / weight_pack_factor;
2682
2683 let h_size = cfg.hidden_size;
2684 let i_size = cfg.intermediate_size;
2685 let gate_proj = h_size * i_size / weight_pack_factor;
2686 let up_proj = h_size * i_size / weight_pack_factor;
2687 let down_proj = i_size * h_size / weight_pack_factor;
2688
2689 input_layernorm
2690 + post_attention_layernorm
2691 + q_proj
2692 + k_proj
2693 + v_proj
2694 + o_proj
2695 + gate_proj
2696 + up_proj
2697 + down_proj
2698 };
2699 Ok(vec![
2700 per_layer_elems * dtype.size_in_bytes();
2701 cfg.num_hidden_layers
2702 ])
2703 }
2704
2705 fn num_layers(&self, config: &str) -> Result<usize> {
2706 let cfg: Idefics3Config = serde_json::from_str(config)?;
2707 Ok(cfg.text_config.num_hidden_layers)
2708 }
2709 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
2710 let cfg: Idefics3Config = serde_json::from_str(config)?;
2711 let cfg = &cfg.text_config;
2712
2713 let cfg = ModelConfigMetadata {
2714 max_seq_len: cfg.max_position_embeddings,
2715 num_layers: cfg.num_hidden_layers,
2716 hidden_size: cfg.hidden_size,
2717 num_kv_heads: cfg.num_key_value_heads,
2718 num_attn_heads: cfg.num_attention_heads,
2719 sliding_window: None,
2720 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
2721 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
2722 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
2723 };
2724
2725 Ok(Box::new(cfg))
2726 }
2727
2728 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
2729 Some(vec![NonMappedSubModel::Vision])
2730 }
2731}
2732
2733pub struct MiniCpmOLoader;
2739
2740pub struct MiniCpmOPrefixer;
2741
2742impl MultimodalPromptPrefixer for MiniCpmOPrefixer {
2743 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
2744 format!(
2745 "{}{prompt}",
2746 "(<image>./</image>)".repeat(image_indexes.len())
2747 )
2748 }
2749}
2750
2751impl MultimodalModelLoader for MiniCpmOLoader {
2752 fn load(
2753 &self,
2754 config: &str,
2755 vb: ShardedVarBuilder,
2756 normal_loading_metadata: NormalLoadingMetadata,
2757 attention_mechanism: AttentionImplementation,
2758 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
2759 let cfg: crate::vision_models::minicpmo::MiniCpmOConfig = serde_json::from_str(config)?;
2760 Ok(Box::new(MiniCpmOModel::new(
2761 &cfg,
2762 vb,
2763 self.is_gptx(config),
2764 normal_loading_metadata,
2765 attention_mechanism,
2766 )?))
2767 }
2768 fn is_gptx(&self, _config: &str) -> bool {
2769 true
2770 }
2771 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
2772 let cfg: crate::vision_models::minicpmo::MiniCpmOConfig = serde_json::from_str(config)?;
2773 Ok(Box::new(cfg))
2774 }
2775 fn get_processor(
2776 &self,
2777 _model_config: &str,
2778 processor_config: Option<ProcessorConfig>,
2779 preprocessor_config: PreProcessorConfig,
2780 max_edge: Option<u32>,
2781 ) -> Arc<dyn Processor + Send + Sync> {
2782 Arc::new(MiniCpmOProcessor::new(
2783 processor_config.unwrap_or_default(),
2784 preprocessor_config,
2785 max_edge,
2786 ))
2787 }
2788 fn supports_paged_attention(&self, _config: &str) -> bool {
2789 true
2790 }
2791 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
2792 Arc::new(MiniCpmOPrefixer)
2793 }
2794 fn modalities(&self, _config: &str) -> Result<Modalities> {
2795 Ok(Modalities {
2796 input: vec![SupportedModality::Text, SupportedModality::Vision],
2797 output: vec![SupportedModality::Text],
2798 })
2799 }
2800}
2801
2802impl IsqModelLoader for MiniCpmOLoader {
2803 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
2804 Ok(vec![
2805 Regex::new(r"llm.lm_head\.(weight|bias)$")?,
2806 Regex::new(r"llm.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
2808 Regex::new(r"llm.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
2809 Regex::new(r"llm.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
2810 Regex::new(r"llm.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
2811 Regex::new(r"llm.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
2813 Regex::new(r"llm.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
2814 Regex::new(r"llm.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
2815 ])
2816 }
2817 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
2818 self.isq_layer_regexes(config)
2819 }
2820}
2821
2822impl DeviceMappedModelLoader for MiniCpmOLoader {
2823 fn mapped_max_act_size_elems(
2824 &self,
2825 config: &str,
2826 params: &AutoDeviceMapParams,
2827 ) -> Result<usize> {
2828 let AutoDeviceMapParams::Multimodal {
2829 max_seq_len,
2830 max_batch_size,
2831 max_image_shape: _,
2832 max_num_images,
2833 } = params
2834 else {
2835 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
2836 };
2837
2838 let cfg: MiniCpmOConfig = serde_json::from_str(config)?;
2839
2840 let num_patches = (cfg.vision_config.image_size / cfg.vision_config.patch_size).pow(2);
2841 let img_seq_len = (num_patches + 1) * max_num_images;
2842
2843 let max_text_attn = {
2844 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
2846 max_batch_size * cfg.text_config.num_attention_heads * max_seq_len * max_seq_len
2847 };
2848
2849 Ok(max_text_attn)
2850 }
2851
2852 fn non_mapped_max_act_size_elems(
2853 &self,
2854 config: &str,
2855 params: &AutoDeviceMapParams,
2856 ) -> Result<usize> {
2857 let AutoDeviceMapParams::Multimodal {
2858 max_seq_len: _,
2859 max_batch_size,
2860 max_image_shape: _,
2861 max_num_images,
2862 } = params
2863 else {
2864 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
2865 };
2866
2867 let cfg: MiniCpmOConfig = serde_json::from_str(config)?;
2868
2869 let num_patches = (cfg.vision_config.image_size / cfg.vision_config.patch_size).pow(2);
2870 let img_seq_len = num_patches + 1;
2871
2872 let max_vision_attn = {
2873 let images_factor = 5;
2875
2876 (max_batch_size * images_factor * max_num_images)
2877 * cfg.vision_config.num_attention_heads
2878 * img_seq_len
2879 * img_seq_len
2880 };
2881
2882 Ok(max_vision_attn)
2883 }
2884
2885 fn non_mapped_size_in_bytes(
2886 &self,
2887 config: &str,
2888 dtype: DType,
2889 weight_pack_factor: usize,
2890 _matformer_config: Option<&MatformerSliceConfig>,
2891 ) -> Result<usize> {
2892 let cfg: MiniCpmOConfig = serde_json::from_str(config)?;
2893 let text_elems = {
2894 let cfg = &cfg.text_config;
2895
2896 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
2897 let lm_head = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
2898 let norm = cfg.hidden_size;
2899 embed_tokens + lm_head + norm
2900 };
2901
2902 let vision_transformer = {
2903 let cfg = &cfg.vision_config;
2904
2905 let post_layernorm = cfg.hidden_size;
2906
2907 let conv_config = Conv2dConfig {
2908 stride: cfg.patch_size,
2909 ..Default::default()
2910 };
2911 let patch_embedding = cfg.num_channels * cfg.hidden_size / conv_config.groups
2912 * cfg.patch_size
2913 * cfg.patch_size;
2914
2915 let num_patches_per_side = cfg.image_size / cfg.patch_size;
2916 let num_patches = num_patches_per_side.pow(2);
2917 let position_embedding = num_patches * cfg.hidden_size;
2918
2919 let layer_elems = {
2920 let layer_norm_1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
2921 let layer_norm_2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
2922
2923 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
2924 let fc2 = cfg.intermediate_size * cfg.hidden_size + cfg.hidden_size;
2925
2926 let q_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2927 let k_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2928 let v_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2929 let o_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
2930
2931 layer_norm_1 + layer_norm_2 + fc1 + fc2 + q_proj + k_proj + v_proj + o_proj
2932 };
2933
2934 post_layernorm
2935 + patch_embedding
2936 + position_embedding
2937 + layer_elems * cfg.num_hidden_layers
2938 };
2939
2940 let elems = text_elems + vision_transformer;
2941
2942 Ok(elems * dtype.size_in_bytes())
2943 }
2944
2945 fn layer_sizes_in_bytes(
2946 &self,
2947 config: &str,
2948 dtype: DType,
2949 weight_pack_factor: usize,
2950 _matformer_config: Option<&MatformerSliceConfig>,
2951 ) -> Result<Vec<usize>> {
2952 let cfg: MiniCpmOConfig = serde_json::from_str(config)?;
2953 let cfg = cfg.text_config;
2954 let per_layer_elems = {
2955 let input_layernorm = cfg.hidden_size;
2956 let post_attention_layernorm = cfg.hidden_size;
2957
2958 let size_in = cfg.hidden_size;
2959 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
2960 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
2961 let q_proj = size_in * size_q / weight_pack_factor;
2962 let k_proj = size_in * size_kv / weight_pack_factor;
2963 let v_proj = size_in * size_kv / weight_pack_factor;
2964 let o_proj = size_q * size_in / weight_pack_factor;
2965
2966 let h_size = cfg.hidden_size;
2967 let i_size = cfg.intermediate_size;
2968 let gate_proj = h_size * i_size / weight_pack_factor;
2969 let up_proj = h_size * i_size / weight_pack_factor;
2970 let down_proj = i_size * h_size / weight_pack_factor;
2971
2972 input_layernorm
2973 + post_attention_layernorm
2974 + q_proj
2975 + k_proj
2976 + v_proj
2977 + o_proj
2978 + gate_proj
2979 + up_proj
2980 + down_proj
2981 };
2982 Ok(vec![
2983 per_layer_elems * dtype.size_in_bytes();
2984 cfg.num_hidden_layers
2985 ])
2986 }
2987
2988 fn num_layers(&self, config: &str) -> Result<usize> {
2989 let cfg: MiniCpmOConfig = serde_json::from_str(config)?;
2990 Ok(cfg.text_config.num_hidden_layers)
2991 }
2992 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
2993 let cfg: MiniCpmOConfig = serde_json::from_str(config)?;
2994 let cfg = &cfg.text_config;
2995
2996 let cfg = ModelConfigMetadata {
2997 max_seq_len: cfg.max_position_embeddings,
2998 num_layers: cfg.num_hidden_layers,
2999 hidden_size: cfg.hidden_size,
3000 num_kv_heads: cfg.num_key_value_heads,
3001 num_attn_heads: cfg.num_attention_heads,
3002 sliding_window: None,
3003 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
3004 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
3005 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
3006 };
3007
3008 Ok(Box::new(cfg))
3009 }
3010}
3011
3012pub struct Phi4MMLoader;
3018
3019pub struct Phi4MMPrefixer;
3020
3021impl MultimodalPromptPrefixer for Phi4MMPrefixer {
3022 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
3023 format!(
3026 "{}{prompt}",
3027 image_indexes
3028 .into_iter()
3029 .map(|image_index| format!("<|image_{}|>", image_index + 1))
3030 .join("")
3031 )
3032 }
3033 fn prefix_audio(&self, audio_indexes: Vec<usize>, prompt: &str) -> String {
3034 format!(
3037 "{}{prompt}",
3038 audio_indexes
3039 .into_iter()
3040 .map(|audio_index| format!("<|audio_{}|>", audio_index + 1))
3041 .join("")
3042 )
3043 }
3044}
3045
3046impl MultimodalModelLoader for Phi4MMLoader {
3047 fn load(
3048 &self,
3049 config: &str,
3050 vb: ShardedVarBuilder,
3051 normal_loading_metadata: NormalLoadingMetadata,
3052 attention_mechanism: AttentionImplementation,
3053 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
3054 let cfg: crate::vision_models::phi4::Phi4MMConfig = serde_json::from_str(config)?;
3055 Ok(Box::new(Phi4MMModel::new(
3056 &cfg,
3057 vb,
3058 self.is_gptx(config),
3059 normal_loading_metadata,
3060 attention_mechanism,
3061 )?))
3062 }
3063 fn is_gptx(&self, _config: &str) -> bool {
3064 true
3065 }
3066 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
3067 let cfg: crate::vision_models::phi4::Phi4MMConfig = serde_json::from_str(config)?;
3068 Ok(Box::new(cfg))
3069 }
3070 fn get_processor(
3071 &self,
3072 _model_config: &str,
3073 processor_config: Option<ProcessorConfig>,
3074 preprocessor_config: PreProcessorConfig,
3075 _max_edge: Option<u32>,
3076 ) -> Arc<dyn Processor + Send + Sync> {
3077 Phi4MMProcessor::new_processor(processor_config, preprocessor_config)
3078 }
3079 fn supports_paged_attention(&self, _config: &str) -> bool {
3080 true
3081 }
3082 fn supports_prefix_cacher(&self, _config: &str) -> bool {
3083 true
3084 }
3085 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
3086 Arc::new(Phi4MMPrefixer)
3087 }
3088 fn modalities(&self, _config: &str) -> Result<Modalities> {
3089 Ok(Modalities {
3090 input: vec![
3091 SupportedModality::Text,
3092 SupportedModality::Vision,
3093 SupportedModality::Audio,
3094 ],
3095 output: vec![SupportedModality::Text],
3096 })
3097 }
3098}
3099
3100impl IsqModelLoader for Phi4MMLoader {
3101 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
3102 Ok(vec![
3103 Regex::new(r"lm_head\.(weight|bias)$")?,
3104 Regex::new(r"layers\.(\d+)\.self_attn\.qkv_proj\.(weight|bias)$")?,
3106 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
3107 Regex::new(r"layers\.(\d+)\.mlp\.gate_up_proj\.(weight|bias)$")?,
3109 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
3110 ])
3111 }
3112 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
3113 self.isq_layer_regexes(config)
3114 }
3115}
3116
3117impl DeviceMappedModelLoader for Phi4MMLoader {
3118 fn mapped_max_act_size_elems(
3119 &self,
3120 config: &str,
3121 params: &AutoDeviceMapParams,
3122 ) -> Result<usize> {
3123 let AutoDeviceMapParams::Multimodal {
3125 max_seq_len,
3126 max_batch_size,
3127 max_image_shape: _,
3128 max_num_images,
3129 } = params
3130 else {
3131 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
3132 };
3133
3134 let cfg: Phi4MMConfig = serde_json::from_str(config)?;
3135
3136 let vcfg = &PHI4_MM_VISION_CFG;
3137
3138 let num_patches = (vcfg.image_size / vcfg.patch_size).pow(2);
3139 let img_seq_len = (num_patches + 1) * max_num_images;
3140
3141 let max_text_attn = {
3142 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
3144 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
3145 };
3146
3147 Ok(max_text_attn)
3148 }
3149
3150 fn non_mapped_max_act_size_elems(
3151 &self,
3152 _config: &str,
3153 params: &AutoDeviceMapParams,
3154 ) -> Result<usize> {
3155 let AutoDeviceMapParams::Multimodal {
3156 max_seq_len: _,
3157 max_batch_size,
3158 max_image_shape,
3159 max_num_images,
3160 } = params
3161 else {
3162 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
3163 };
3164
3165 let vcfg = &PHI4_MM_VISION_CFG;
3166
3167 let num_patches = (vcfg.image_size / vcfg.patch_size).pow(2);
3168 let img_seq_len = num_patches + 1;
3169
3170 let max_batch_size = max_batch_size
3171 * (max_image_shape
3172 .0
3173 .div_ceil(phi4::inputs_processor::DYHD_BASE_RESOLUTION)
3174 * max_image_shape
3175 .1
3176 .div_ceil(phi4::inputs_processor::DYHD_BASE_RESOLUTION)
3177 + 1);
3178
3179 let max_vision_attn = (max_batch_size * max_num_images)
3180 * vcfg.num_attention_heads
3181 * img_seq_len
3182 * img_seq_len;
3183 let max_qkv = 3
3184 * (max_batch_size
3185 * vcfg.num_attention_heads
3186 * img_seq_len
3187 * (vcfg.hidden_size / vcfg.num_attention_heads));
3188
3189 Ok(max_vision_attn + max_qkv)
3190 }
3191
3192 fn non_mapped_size_in_bytes(
3193 &self,
3194 config: &str,
3195 dtype: DType,
3196 weight_pack_factor: usize,
3197 _matformer_config: Option<&MatformerSliceConfig>,
3198 ) -> Result<usize> {
3199 let cfg: Phi4MMConfig = serde_json::from_str(config)?;
3200 let elems = {
3201 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
3202 let lm_head = if !cfg.tie_word_embeddings || weight_pack_factor != 1 {
3204 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
3205 } else {
3206 0
3207 };
3208 let norm = cfg.hidden_size;
3209
3210 let image_embed = if let Some(img_embed) = &cfg.embd_layer.image_embd_layer {
3211 let projection_cls = img_embed
3212 .projection_cls
3213 .clone()
3214 .unwrap_or("linear".to_string());
3215 let with_learnable_separator = img_embed.with_learnable_separator.unwrap_or(false);
3216 let use_hd_transform = img_embed.use_hd_transform.unwrap_or(false);
3217 let image_dim_out = PHI4_MM_VISION_CFG.hidden_size;
3218
3219 let proj = match (projection_cls.as_str(), use_hd_transform) {
3220 ("linear", _) => image_dim_out * cfg.hidden_size + cfg.hidden_size,
3221 ("mlp", true) => {
3222 let a = (image_dim_out * 4) * cfg.hidden_size + cfg.hidden_size;
3223 let b = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3224 a + b
3225 }
3226 ("mlp", false) => {
3227 let a = image_dim_out * cfg.hidden_size + cfg.hidden_size;
3228 let b = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3229 a + b
3230 }
3231 _ => {
3232 anyhow::bail!("projection_cls=`{projection_cls}` not implemented.");
3233 }
3234 };
3235
3236 let (glb_gn, sub_gn) = if with_learnable_separator {
3237 let glb_gn = image_dim_out * 4;
3238 let sub_gn = image_dim_out * 4;
3239 (glb_gn, sub_gn)
3240 } else {
3241 (0, 0)
3242 };
3243
3244 let vision_transformer = {
3245 let cfg = &PHI4_MM_VISION_CFG;
3246
3247 let post_layernorm = cfg.hidden_size;
3248
3249 let conv_config = Conv2dConfig {
3250 stride: cfg.patch_size,
3251 ..Default::default()
3252 };
3253 let patch_embedding = cfg.num_channels * cfg.hidden_size / conv_config.groups
3254 * cfg.patch_size
3255 * cfg.patch_size;
3256
3257 let num_patches_per_side = cfg.image_size / cfg.patch_size;
3258 let num_patches = num_patches_per_side.pow(2);
3259 let position_embedding = num_patches * cfg.hidden_size;
3260
3261 let layer_elems = {
3262 let layer_norm_1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
3263 let layer_norm_2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
3264
3265 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
3266 let fc2 = cfg.intermediate_size * cfg.hidden_size + cfg.hidden_size;
3267
3268 let q_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3269 let k_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3270 let v_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3271 let o_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3272
3273 layer_norm_1 + layer_norm_2 + fc1 + fc2 + q_proj + k_proj + v_proj + o_proj
3274 };
3275
3276 post_layernorm
3277 + patch_embedding
3278 + position_embedding
3279 + layer_elems * cfg.num_hidden_layers
3280 };
3281
3282 proj + glb_gn + sub_gn + vision_transformer
3283 } else {
3284 0
3285 };
3286
3287 embed_tokens + lm_head + norm + image_embed
3288 };
3289
3290 Ok(elems * dtype.size_in_bytes())
3291 }
3292
3293 fn layer_sizes_in_bytes(
3294 &self,
3295 config: &str,
3296 dtype: DType,
3297 weight_pack_factor: usize,
3298 _matformer_config: Option<&MatformerSliceConfig>,
3299 ) -> Result<Vec<usize>> {
3300 let cfg: Phi4MMConfig = serde_json::from_str(config)?;
3301 let per_layer_elems = {
3302 let input_layernorm = cfg.hidden_size;
3303 let post_attention_layernorm = cfg.hidden_size;
3304
3305 let size_in = cfg.hidden_size;
3306 let head_dim = cfg.head_dim();
3307 let op_size =
3308 cfg.num_attention_heads * head_dim + 2 * cfg.num_key_value_heads() * head_dim;
3309 let qkv_proj = size_in * op_size / weight_pack_factor;
3310 let o_proj = (cfg.num_attention_heads * head_dim) * size_in / weight_pack_factor;
3311
3312 let h_size = cfg.hidden_size;
3313 let i_size = cfg.intermediate_size;
3314 let gate_up_proj = h_size * (2 * i_size) / weight_pack_factor;
3315 let down_proj = h_size * i_size / weight_pack_factor;
3316
3317 input_layernorm
3318 + post_attention_layernorm
3319 + qkv_proj
3320 + o_proj
3321 + gate_up_proj
3322 + down_proj
3323 };
3324 Ok(vec![
3325 per_layer_elems * dtype.size_in_bytes();
3326 cfg.num_hidden_layers
3327 ])
3328 }
3329
3330 fn num_layers(&self, config: &str) -> Result<usize> {
3331 let cfg: Phi4MMConfig = serde_json::from_str(config)?;
3332 Ok(cfg.num_hidden_layers)
3333 }
3334
3335 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
3336 let cfg: Phi4MMConfig = serde_json::from_str(config)?;
3337
3338 let cfg = ModelConfigMetadata {
3339 max_seq_len: cfg.max_position_embeddings,
3340 num_layers: cfg.num_hidden_layers,
3341 hidden_size: cfg.hidden_size,
3342 num_kv_heads: cfg.num_key_value_heads(),
3343 num_attn_heads: cfg.num_attention_heads,
3344 sliding_window: cfg.sliding_window,
3345 k_head_dim: cfg.head_dim(),
3346 v_head_dim: cfg.head_dim(),
3347 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
3348 };
3349
3350 Ok(Box::new(cfg))
3351 }
3352
3353 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
3354 Some(vec![NonMappedSubModel::Vision, NonMappedSubModel::Audio])
3355 }
3356}
3357
3358pub struct Qwen2_5VLLoader;
3364
3365pub struct Qwen2_5VLPrefixer;
3366
3367impl MultimodalPromptPrefixer for Qwen2_5VLPrefixer {
3368 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
3369 format!(
3370 "{}{prompt}",
3371 format!(
3372 "{}{}{}",
3373 Qwen2_5VLProcessor::VISION_START,
3374 Qwen2_5VLProcessor::IMAGE_PAD,
3375 Qwen2_5VLProcessor::VISION_END
3376 )
3377 .repeat(image_indexes.len())
3378 )
3379 }
3380}
3381
3382impl MultimodalModelLoader for Qwen2_5VLLoader {
3383 fn load(
3384 &self,
3385 config: &str,
3386 vb: ShardedVarBuilder,
3387 normal_loading_metadata: NormalLoadingMetadata,
3388 attention_mechanism: AttentionImplementation,
3389 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
3390 let cfg: Qwen2_5VLConfig = serde_json::from_str(config)?;
3391 Ok(Box::new(Qwen2_5VLModel::new(
3392 &cfg,
3393 vb,
3394 self.is_gptx(config),
3395 normal_loading_metadata,
3396 attention_mechanism,
3397 )?))
3398 }
3399 fn is_gptx(&self, _config: &str) -> bool {
3400 true
3401 }
3402 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
3403 let config: Qwen2_5VLConfig = serde_json::from_str(config)?;
3404 Ok(Box::new(config))
3405 }
3406 fn get_processor(
3407 &self,
3408 _model_config: &str,
3409 _processor_config: Option<ProcessorConfig>,
3410 _preprocessor_config: PreProcessorConfig,
3411 max_edge: Option<u32>,
3412 ) -> Arc<dyn Processor + Send + Sync> {
3413 Arc::new(Qwen2_5VLProcessor::new(max_edge))
3414 }
3415 fn supports_paged_attention(&self, _config: &str) -> bool {
3416 false
3417 }
3418 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
3419 Arc::new(Qwen2_5VLPrefixer)
3420 }
3421 fn modalities(&self, _config: &str) -> Result<Modalities> {
3422 Ok(Modalities {
3423 input: vec![SupportedModality::Text, SupportedModality::Vision],
3424 output: vec![SupportedModality::Text],
3425 })
3426 }
3427}
3428
3429impl IsqModelLoader for Qwen2_5VLLoader {
3430 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
3431 Ok(vec![
3432 Regex::new(r"lm_head\.(weight|bias)$")?,
3433 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
3435 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
3436 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
3437 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
3438 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
3440 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
3441 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
3442 ])
3443 }
3444 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
3445 self.isq_layer_regexes(config)
3446 }
3447}
3448
3449impl DeviceMappedModelLoader for Qwen2_5VLLoader {
3450 fn mapped_max_act_size_elems(
3451 &self,
3452 config: &str,
3453 params: &AutoDeviceMapParams,
3454 ) -> Result<usize> {
3455 let AutoDeviceMapParams::Multimodal {
3456 max_seq_len,
3457 max_batch_size,
3458 max_image_shape,
3459 max_num_images,
3460 } = params
3461 else {
3462 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
3463 };
3464
3465 let cfg: Qwen2_5VLConfig = serde_json::from_str(config)?;
3466
3467 let img_seq_len = {
3468 let cfg = &cfg.vision_config;
3469 let grid_t = max_num_images / cfg.temporal_patch_size;
3470 let grid_h = max_image_shape.0 / cfg.patch_size;
3471 let grid_w = max_image_shape.1 / cfg.patch_size;
3472 grid_t * grid_h * grid_w
3473 };
3474 let img_seq_len = img_seq_len * max_num_images;
3475
3476 let max_text_attn = {
3477 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
3479 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
3480 };
3481
3482 Ok(max_text_attn)
3483 }
3484
3485 fn non_mapped_max_act_size_elems(
3486 &self,
3487 config: &str,
3488 params: &AutoDeviceMapParams,
3489 ) -> Result<usize> {
3490 let AutoDeviceMapParams::Multimodal {
3491 max_seq_len: _,
3492 max_batch_size,
3493 max_image_shape,
3494 max_num_images,
3495 } = params
3496 else {
3497 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
3498 };
3499
3500 let cfg: Qwen2_5VLConfig = serde_json::from_str(config)?;
3501
3502 let img_seq_len = {
3503 let cfg = &cfg.vision_config;
3504 let grid_t = max_num_images / cfg.temporal_patch_size;
3505 let grid_h = max_image_shape.0 / cfg.patch_size;
3506 let grid_w = max_image_shape.1 / cfg.patch_size;
3507 grid_t * grid_h * grid_w
3508 };
3509
3510 let max_vision_attn = {
3511 let cfg = &cfg.vision_config;
3512 (max_batch_size * max_num_images) * cfg.num_heads * img_seq_len * img_seq_len
3513 };
3514
3515 Ok(max_vision_attn)
3516 }
3517
3518 fn non_mapped_size_in_bytes(
3519 &self,
3520 config: &str,
3521 dtype: DType,
3522 weight_pack_factor: usize,
3523 _matformer_config: Option<&MatformerSliceConfig>,
3524 ) -> Result<usize> {
3525 let cfg: Qwen2_5VLConfig = serde_json::from_str(config)?;
3526 let text_elems = {
3527 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
3528 let lm_head = if !cfg.tie_word_embeddings || weight_pack_factor != 1 {
3530 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
3531 } else {
3532 0
3533 };
3534 let norm = cfg.hidden_size;
3535 embed_tokens + lm_head + norm
3536 };
3537
3538 let patch_merger = {
3539 let cfg = &cfg.vision_config;
3540 let hidden_size = cfg.hidden_size * cfg.spatial_merge_size.pow(2);
3541
3542 let mlp0 = hidden_size * hidden_size + hidden_size;
3543 let mlp2 = hidden_size * cfg.hidden_size + cfg.hidden_size;
3544
3545 let ln_q = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
3546
3547 mlp0 + mlp2 + ln_q
3548 };
3549
3550 let patch_embed = {
3551 let cfg = &cfg.vision_config;
3552 let conv_cfg = Conv3dConfig {
3553 stride: cfg.patch_size,
3554 ..Default::default()
3555 };
3556 let kernel_sizes = [cfg.temporal_patch_size, cfg.patch_size, cfg.patch_size];
3557 cfg.in_chans * cfg.hidden_size / conv_cfg.groups
3558 * kernel_sizes[0]
3559 * kernel_sizes[1]
3560 * kernel_sizes[2]
3561 };
3562
3563 let encoder_layer = {
3564 let cfg = &cfg.vision_config;
3565 let norm1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
3566 let norm2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
3567
3568 #[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
3569 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
3570 let fc2 = cfg.hidden_size * cfg.intermediate_size + cfg.hidden_size;
3571
3572 let qkv = cfg.hidden_size * cfg.hidden_size * 3 + cfg.hidden_size * 3;
3573 let out = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3574
3575 norm1 + norm2 + fc1 + fc2 + qkv + out
3576 };
3577
3578 let elems =
3579 text_elems + patch_merger + patch_embed + encoder_layer * cfg.vision_config.depth;
3580
3581 Ok(elems * dtype.size_in_bytes())
3582 }
3583
3584 fn layer_sizes_in_bytes(
3585 &self,
3586 config: &str,
3587 dtype: DType,
3588 weight_pack_factor: usize,
3589 _matformer_config: Option<&MatformerSliceConfig>,
3590 ) -> Result<Vec<usize>> {
3591 let cfg: Qwen2_5VLConfig = serde_json::from_str(config)?;
3592 let per_layer_elems = {
3593 let input_layernorm = cfg.hidden_size;
3594 let post_attention_layernorm = cfg.hidden_size;
3595
3596 let size_in = cfg.hidden_size;
3597 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
3598 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
3599 let q_proj = size_in * size_q / weight_pack_factor + size_q;
3600 let k_proj = size_in * size_kv / weight_pack_factor + size_kv;
3601 let v_proj = size_in * size_kv / weight_pack_factor + size_kv;
3602 let o_proj = size_q * size_in / weight_pack_factor;
3603
3604 let h_size = cfg.hidden_size;
3605 let i_size = cfg.intermediate_size;
3606 let gate_proj = h_size * i_size / weight_pack_factor;
3607 let up_proj = h_size * i_size / weight_pack_factor;
3608 let down_proj = i_size * h_size / weight_pack_factor;
3609
3610 input_layernorm
3611 + post_attention_layernorm
3612 + q_proj
3613 + k_proj
3614 + v_proj
3615 + o_proj
3616 + gate_proj
3617 + up_proj
3618 + down_proj
3619 };
3620 Ok(vec![
3621 per_layer_elems * dtype.size_in_bytes();
3622 cfg.num_hidden_layers
3623 ])
3624 }
3625
3626 fn num_layers(&self, config: &str) -> Result<usize> {
3627 let cfg: Qwen2_5VLConfig = serde_json::from_str(config)?;
3628 Ok(cfg.num_hidden_layers)
3629 }
3630
3631 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
3632 let cfg: Qwen2_5VLConfig = serde_json::from_str(config)?;
3633
3634 let cfg = ModelConfigMetadata {
3635 max_seq_len: cfg.max_position_embeddings,
3636 num_layers: cfg.num_hidden_layers,
3637 hidden_size: cfg.hidden_size,
3638 num_kv_heads: cfg.num_key_value_heads,
3639 num_attn_heads: cfg.num_attention_heads,
3640 sliding_window: cfg.sliding_window,
3641 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
3642 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
3643 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
3644 };
3645
3646 Ok(Box::new(cfg))
3647 }
3648
3649 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
3650 Some(vec![NonMappedSubModel::Vision])
3651 }
3652}
3653
3654pub struct Gemma3Loader;
3660
3661pub struct Gemma3Prefixer;
3662
3663impl MultimodalPromptPrefixer for Gemma3Prefixer {
3664 fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
3665 prompt.to_string()
3666 }
3667}
3668
3669impl MultimodalModelLoader for Gemma3Loader {
3670 fn load(
3671 &self,
3672 config: &str,
3673 vb: ShardedVarBuilder,
3674 normal_loading_metadata: NormalLoadingMetadata,
3675 attention_mechanism: AttentionImplementation,
3676 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
3677 let cfg: Gemma3Config = serde_json::from_str(config)?;
3678 Ok(Box::new(Gemma3Model::new(
3679 &cfg,
3680 vb,
3681 self.is_gptx(config),
3682 normal_loading_metadata,
3683 attention_mechanism,
3684 )?))
3685 }
3686 fn is_gptx(&self, _config: &str) -> bool {
3687 true
3688 }
3689 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
3690 let config: Gemma3Config = serde_json::from_str(config)?;
3691 Ok(Box::new(config))
3692 }
3693 fn get_processor(
3694 &self,
3695 config: &str,
3696 processor_config: Option<ProcessorConfig>,
3697 _preprocessor_config: PreProcessorConfig,
3698 _max_edge: Option<u32>,
3699 ) -> Arc<dyn Processor + Send + Sync> {
3700 let config: Gemma3Config = serde_json::from_str(config).unwrap();
3701 Arc::new(Gemma3Processor::new(
3703 processor_config.unwrap_or_default(),
3704 matches!(config, Gemma3Config::WithVision { .. }),
3705 ))
3706 }
3707 fn supports_paged_attention(&self, _config: &str) -> bool {
3708 true
3709 }
3710 fn supports_prefix_cacher(&self, _config: &str) -> bool {
3711 true
3712 }
3713 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
3714 Arc::new(Gemma3Prefixer)
3715 }
3716 fn modalities(&self, _config: &str) -> Result<Modalities> {
3717 Ok(Modalities {
3718 input: vec![SupportedModality::Text, SupportedModality::Vision],
3719 output: vec![SupportedModality::Text],
3720 })
3721 }
3722}
3723
3724impl IsqModelLoader for Gemma3Loader {
3725 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
3726 Ok(vec![
3727 Regex::new(r"lm_head\.(weight|bias)$")?,
3728 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
3730 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
3731 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
3732 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
3733 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
3735 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
3736 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
3737 ])
3738 }
3739 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
3740 Ok(vec![
3741 Regex::new(r"lm_head\.(weight|bias)$")?,
3742 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
3744 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
3745 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
3746 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
3747 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
3749 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
3750 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
3751 ])
3752 }
3753}
3754
3755impl DeviceMappedModelLoader for Gemma3Loader {
3756 fn mapped_max_act_size_elems(
3757 &self,
3758 config: &str,
3759 params: &AutoDeviceMapParams,
3760 ) -> Result<usize> {
3761 let AutoDeviceMapParams::Multimodal {
3762 max_seq_len,
3763 max_batch_size,
3764 max_image_shape: _,
3765 max_num_images,
3766 } = params
3767 else {
3768 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
3769 };
3770
3771 let cfg: Gemma3Config = serde_json::from_str(config)?;
3772
3773 match cfg {
3774 Gemma3Config::Text(text_config) => Ok(max_batch_size
3775 * text_config.num_attention_heads
3776 * max_seq_len.min(&ATTENTION_CHUNK_SIZE).pow(2)),
3777 Gemma3Config::WithVision {
3778 text_config,
3779 vision_config,
3780 ..
3781 } => {
3782 let num_patches = (vision_config.image_size / vision_config.patch_size).pow(2);
3783 let img_seq_len = (num_patches + 1) * max_num_images;
3784
3785 let max_text_attn = {
3786 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
3788 max_batch_size * text_config.num_attention_heads * max_seq_len * max_seq_len
3789 };
3790 Ok(max_text_attn)
3791 }
3792 }
3793 }
3794
3795 fn non_mapped_max_act_size_elems(
3796 &self,
3797 config: &str,
3798 params: &AutoDeviceMapParams,
3799 ) -> Result<usize> {
3800 let AutoDeviceMapParams::Multimodal {
3801 max_seq_len: _,
3802 max_batch_size,
3803 max_image_shape: _,
3804 max_num_images,
3805 } = params
3806 else {
3807 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
3808 };
3809
3810 let cfg: Gemma3Config = serde_json::from_str(config)?;
3811
3812 match cfg {
3813 Gemma3Config::WithVision { vision_config, .. } => {
3814 let num_patches = (vision_config.image_size / vision_config.patch_size).pow(2);
3815 let img_seq_len = num_patches + 1;
3816
3817 let max_vision_attn = {
3818 (max_batch_size * max_num_images)
3819 * vision_config.num_attention_heads
3820 * img_seq_len
3821 * img_seq_len
3822 };
3823
3824 Ok(max_vision_attn)
3825 }
3826 Gemma3Config::Text(_) => Ok(0),
3827 }
3828 }
3829
3830 fn non_mapped_size_in_bytes(
3831 &self,
3832 config: &str,
3833 dtype: DType,
3834 weight_pack_factor: usize,
3835 _matformer_config: Option<&MatformerSliceConfig>,
3836 ) -> Result<usize> {
3837 let cfg: Gemma3Config = serde_json::from_str(config)?;
3838
3839 let text_elems = {
3840 let cfg = match &cfg {
3841 Gemma3Config::Text(cfg) => cfg,
3842 Gemma3Config::WithVision { text_config, .. } => text_config,
3843 };
3844 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
3845 let lm_head = if !cfg.tie_word_embeddings || weight_pack_factor != 1 {
3847 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
3848 } else {
3849 0
3850 };
3851 let norm = cfg.hidden_size;
3852 embed_tokens + lm_head + norm
3853 };
3854
3855 let vision_transformer = if let Gemma3Config::WithVision {
3856 vision_config: cfg, ..
3857 } = &cfg
3858 {
3859 let post_layernorm = cfg.hidden_size;
3860
3861 let conv_config = Conv2dConfig {
3862 stride: cfg.patch_size,
3863 ..Default::default()
3864 };
3865 let patch_embedding = cfg.num_channels * cfg.hidden_size / conv_config.groups
3866 * cfg.patch_size
3867 * cfg.patch_size;
3868
3869 let num_patches_per_side = cfg.image_size / cfg.patch_size;
3870 let num_patches = num_patches_per_side.pow(2);
3871 let position_embedding = num_patches * cfg.hidden_size;
3872
3873 let layer_elems = {
3874 let layer_norm_1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
3875 let layer_norm_2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
3876
3877 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
3878 let fc2 = cfg.intermediate_size * cfg.hidden_size + cfg.hidden_size;
3879
3880 let q_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3881 let k_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3882 let v_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3883 let o_proj = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
3884
3885 layer_norm_1 + layer_norm_2 + fc1 + fc2 + q_proj + k_proj + v_proj + o_proj
3886 };
3887
3888 post_layernorm
3889 + patch_embedding
3890 + position_embedding
3891 + layer_elems * cfg.num_hidden_layers
3892 } else {
3893 0
3894 };
3895
3896 let elems = text_elems + vision_transformer;
3897
3898 Ok(elems * dtype.size_in_bytes())
3899 }
3900
3901 fn layer_sizes_in_bytes(
3902 &self,
3903 config: &str,
3904 dtype: DType,
3905 weight_pack_factor: usize,
3906 _matformer_config: Option<&MatformerSliceConfig>,
3907 ) -> Result<Vec<usize>> {
3908 let cfg: Gemma3Config = serde_json::from_str(config)?;
3909
3910 let txt_cfg = match &cfg {
3911 Gemma3Config::Text(cfg) => cfg,
3912 Gemma3Config::WithVision { text_config, .. } => text_config,
3913 };
3914 let per_layer_elems = {
3915 let cfg = txt_cfg;
3916
3917 let input_layernorm = cfg.hidden_size;
3918 let post_attention_layernorm = cfg.hidden_size;
3919
3920 let size_in = cfg.hidden_size;
3921 let size_q = cfg.head_dim * cfg.num_attention_heads;
3922 let size_kv = cfg.head_dim * cfg.num_key_value_heads;
3923 let q_proj =
3924 size_in * size_q / weight_pack_factor + bias_if!(cfg.attention_bias, size_q);
3925 let k_proj =
3926 size_in * size_kv / weight_pack_factor + bias_if!(cfg.attention_bias, size_kv);
3927 let v_proj =
3928 size_in * size_kv / weight_pack_factor + bias_if!(cfg.attention_bias, size_kv);
3929 let o_proj =
3930 size_q * size_in / weight_pack_factor + bias_if!(cfg.attention_bias, size_in);
3931
3932 let h_size = cfg.hidden_size;
3933 let i_size = cfg.intermediate_size;
3934 let gate_proj = h_size * i_size / weight_pack_factor;
3935 let up_proj = h_size * i_size / weight_pack_factor;
3936 let down_proj = i_size * h_size / weight_pack_factor;
3937
3938 input_layernorm
3939 + post_attention_layernorm
3940 + q_proj
3941 + k_proj
3942 + v_proj
3943 + o_proj
3944 + gate_proj
3945 + up_proj
3946 + down_proj
3947 };
3948 Ok(vec![
3949 per_layer_elems * dtype.size_in_bytes();
3950 txt_cfg.num_hidden_layers
3951 ])
3952 }
3953
3954 fn num_layers(&self, config: &str) -> Result<usize> {
3955 let cfg: Gemma3Config = serde_json::from_str(config)?;
3956
3957 let txt_cfg = match &cfg {
3958 Gemma3Config::Text(cfg) => cfg,
3959 Gemma3Config::WithVision { text_config, .. } => text_config,
3960 };
3961
3962 Ok(txt_cfg.num_hidden_layers)
3963 }
3964
3965 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
3966 let cfg: Gemma3Config = serde_json::from_str(config)?;
3967
3968 let cfg = match &cfg {
3969 Gemma3Config::Text(cfg) => cfg,
3970 Gemma3Config::WithVision { text_config, .. } => text_config,
3971 };
3972
3973 let cfg = ModelConfigMetadata {
3974 max_seq_len: cfg.max_position_embeddings,
3975 num_layers: cfg.num_hidden_layers,
3976 hidden_size: cfg.hidden_size,
3977 num_kv_heads: cfg.num_key_value_heads,
3978 num_attn_heads: cfg.num_attention_heads,
3979 sliding_window: None, k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
3981 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
3982 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
3983 };
3984
3985 Ok(Box::new(cfg))
3986 }
3987
3988 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
3989 Some(vec![NonMappedSubModel::Vision])
3990 }
3991}
3992
3993pub struct Mistral3Loader;
3999
4000pub struct Mistral3Prefixer;
4001
4002impl MultimodalPromptPrefixer for Mistral3Prefixer {
4003 fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
4004 prompt.to_string()
4005 }
4006}
4007
4008impl MultimodalModelLoader for Mistral3Loader {
4009 fn load(
4010 &self,
4011 config: &str,
4012 vb: ShardedVarBuilder,
4013 normal_loading_metadata: NormalLoadingMetadata,
4014 attention_mechanism: AttentionImplementation,
4015 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
4016 let mut cfg: crate::vision_models::mistral3::Mistral3Config = serde_json::from_str(config)?;
4017 cfg.propagate_quantization_config();
4018 Ok(Box::new(Mistral3Model::new(
4019 &cfg,
4020 vb,
4021 self.is_gptx(config),
4022 normal_loading_metadata,
4023 attention_mechanism,
4024 )?))
4025 }
4026 fn is_gptx(&self, _config: &str) -> bool {
4027 true
4028 }
4029 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
4030 let cfg: crate::vision_models::mistral3::Mistral3Config = serde_json::from_str(config)?;
4031 Ok(Box::new(cfg))
4032 }
4033 fn get_processor(
4034 &self,
4035 _model_config: &str,
4036 processor_config: Option<ProcessorConfig>,
4037 _preprocessor_config: PreProcessorConfig,
4038 _max_edge: Option<u32>,
4039 ) -> Arc<dyn Processor + Send + Sync> {
4040 Arc::new(Mistral3Processor::new(processor_config.unwrap_or_default()))
4041 }
4042 fn supports_paged_attention(&self, _config: &str) -> bool {
4043 true
4044 }
4045 fn supports_prefix_cacher(&self, _config: &str) -> bool {
4046 true
4047 }
4048 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
4049 Arc::new(Mistral3Prefixer)
4050 }
4051 fn modalities(&self, _config: &str) -> Result<Modalities> {
4052 Ok(Modalities {
4053 input: vec![SupportedModality::Text, SupportedModality::Vision],
4054 output: vec![SupportedModality::Text],
4055 })
4056 }
4057}
4058
4059impl IsqModelLoader for Mistral3Loader {
4060 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
4061 Ok(vec![
4062 Regex::new(r"lm_head\.(weight|bias)$")?,
4063 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
4065 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
4066 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
4067 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
4068 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
4070 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
4071 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
4072 ])
4073 }
4074 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
4075 Ok(vec![
4076 Regex::new(r"lm_head\.(weight|bias)$")?,
4077 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
4079 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
4080 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
4081 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
4082 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
4084 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
4085 Regex::new(r"language_model\.model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
4086 ])
4087 }
4088}
4089
4090#[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
4091impl DeviceMappedModelLoader for Mistral3Loader {
4092 fn mapped_max_act_size_elems(
4093 &self,
4094 config: &str,
4095 params: &AutoDeviceMapParams,
4096 ) -> Result<usize> {
4097 let cfg: Mistral3Config = serde_json::from_str(config)?;
4098 let vcfg = &cfg.vision_config;
4099 let tcfg = &cfg.text_config;
4100
4101 let AutoDeviceMapParams::Multimodal {
4102 max_seq_len,
4103 max_batch_size,
4104 max_image_shape: (mut height, mut width),
4105 max_num_images,
4106 } = params
4107 else {
4108 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
4109 };
4110
4111 let img_seq_len = {
4112 let (max_height, max_width) = (1540, 1540);
4116 let ratio = (height as f64 / max_height as f64).max(width as f64 / max_width as f64);
4117 if ratio > 1. {
4118 height = (height as f64 / ratio).floor() as usize;
4119 width = (width as f64 / ratio).floor() as usize;
4120 }
4121
4122 let num_height_tokens = (height - 1) / vcfg.patch_size + 1;
4123 let num_width_tokens = (width - 1) / vcfg.patch_size + 1;
4124
4125 height = num_height_tokens * vcfg.patch_size;
4126 width = num_width_tokens * vcfg.patch_size;
4127
4128 let num_height_tokens = height / vcfg.patch_size;
4129 let num_width_tokens = width / vcfg.patch_size;
4130
4131 (num_width_tokens + 1) * num_height_tokens
4132 };
4133
4134 let max_seq_len = img_seq_len * max_num_images + *max_seq_len.min(&ATTENTION_CHUNK_SIZE);
4136 Ok(max_batch_size * tcfg.num_attention_heads * max_seq_len * max_seq_len)
4137 }
4138
4139 fn non_mapped_max_act_size_elems(
4140 &self,
4141 config: &str,
4142 params: &AutoDeviceMapParams,
4143 ) -> Result<usize> {
4144 let cfg: Mistral3Config = serde_json::from_str(config)?;
4145 let cfg = &cfg.vision_config;
4146
4147 let AutoDeviceMapParams::Multimodal {
4148 max_seq_len: _,
4149 max_batch_size,
4150 max_image_shape: (mut height, mut width),
4151 max_num_images,
4152 } = params
4153 else {
4154 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
4155 };
4156
4157 let img_seq_len = {
4158 let (max_height, max_width) = (1540, 1540);
4162 let ratio = (height as f64 / max_height as f64).max(width as f64 / max_width as f64);
4163 if ratio > 1. {
4164 height = (height as f64 / ratio).floor() as usize;
4165 width = (width as f64 / ratio).floor() as usize;
4166 }
4167
4168 let num_height_tokens = (height - 1) / cfg.patch_size + 1;
4169 let num_width_tokens = (width - 1) / cfg.patch_size + 1;
4170
4171 height = num_height_tokens * cfg.patch_size;
4172 width = num_width_tokens * cfg.patch_size;
4173
4174 let num_height_tokens = height / cfg.patch_size;
4175 let num_width_tokens = width / cfg.patch_size;
4176
4177 (num_width_tokens + 1) * num_height_tokens
4178 };
4179
4180 Ok((max_batch_size * max_num_images) * cfg.num_attention_heads * img_seq_len * img_seq_len)
4181 }
4182
4183 fn non_mapped_size_in_bytes(
4184 &self,
4185 config: &str,
4186 dtype: DType,
4187 weight_pack_factor: usize,
4188 _matformer_config: Option<&MatformerSliceConfig>,
4189 ) -> Result<usize> {
4190 let cfg: Mistral3Config = serde_json::from_str(config)?;
4191
4192 let text_elems = {
4193 let cfg = &cfg.text_config;
4194
4195 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
4196 let lm_head = if !cfg.tie_word_embeddings || weight_pack_factor != 1 {
4198 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
4199 } else {
4200 0
4201 };
4202 let norm = cfg.hidden_size;
4203 embed_tokens + lm_head + norm
4204 };
4205
4206 let vision_elems = {
4207 let cfg = &cfg.vision_config;
4208
4209 let patch_embed = {
4210 let conv_cfg = Conv2dConfig {
4211 stride: cfg.patch_size,
4212 ..Default::default()
4213 };
4214 cfg.num_channels * cfg.hidden_size / conv_cfg.groups
4215 * cfg.patch_size
4216 * cfg.patch_size
4217 * cfg.patch_size
4218 };
4219 let ln_pre = cfg.hidden_size;
4220 let vision_layer = {
4221 let attn_norm = cfg.hidden_size;
4222 let ffn_norm = cfg.hidden_size;
4223
4224 let gate = cfg.hidden_size * cfg.intermediate_size;
4225 let up = cfg.hidden_size * cfg.intermediate_size;
4226 let down = cfg.hidden_size * cfg.intermediate_size;
4227
4228 let q = cfg.hidden_size * cfg.hidden_size;
4229 let k = cfg.hidden_size * cfg.hidden_size;
4230 let v = cfg.hidden_size * cfg.hidden_size;
4231 let o = cfg.hidden_size * cfg.hidden_size;
4232
4233 attn_norm + ffn_norm + gate + up + down + q + k + v + o
4234 };
4235
4236 patch_embed + ln_pre + vision_layer * cfg.num_hidden_layers
4237 };
4238
4239 let elems = text_elems + vision_elems;
4240
4241 Ok(elems * dtype.size_in_bytes())
4242 }
4243
4244 fn layer_sizes_in_bytes(
4245 &self,
4246 config: &str,
4247 dtype: DType,
4248 weight_pack_factor: usize,
4249 _matformer_config: Option<&MatformerSliceConfig>,
4250 ) -> Result<Vec<usize>> {
4251 let cfg: Mistral3Config = serde_json::from_str(config)?;
4252 let cfg = &cfg.text_config;
4253
4254 let per_layer_elems = {
4255 let input_layernorm = cfg.hidden_size;
4256 let post_attention_layernorm = cfg.hidden_size;
4257
4258 let size_in = cfg.hidden_size;
4259 let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
4260 let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
4261 let q_proj = size_in * size_q / weight_pack_factor;
4262 let k_proj = size_in * size_kv / weight_pack_factor;
4263 let v_proj = size_in * size_kv / weight_pack_factor;
4264 let o_proj = size_q * size_in / weight_pack_factor;
4265
4266 let h_size = cfg.hidden_size;
4267 let i_size = cfg.intermediate_size;
4268 let gate_proj = h_size * i_size / weight_pack_factor;
4269 let up_proj = h_size * i_size / weight_pack_factor;
4270 let down_proj = i_size * h_size / weight_pack_factor;
4271
4272 input_layernorm
4273 + post_attention_layernorm
4274 + q_proj
4275 + k_proj
4276 + v_proj
4277 + o_proj
4278 + gate_proj
4279 + up_proj
4280 + down_proj
4281 };
4282 Ok(vec![
4283 per_layer_elems * dtype.size_in_bytes();
4284 cfg.num_hidden_layers
4285 ])
4286 }
4287
4288 fn num_layers(&self, config: &str) -> Result<usize> {
4289 let cfg: Mistral3Config = serde_json::from_str(config)?;
4290 let cfg = &cfg.text_config;
4291 Ok(cfg.num_hidden_layers)
4292 }
4293
4294 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
4295 let cfg: Mistral3Config = serde_json::from_str(config)?;
4296 let cfg = &cfg.text_config;
4297
4298 let cfg = ModelConfigMetadata {
4299 max_seq_len: cfg.max_position_embeddings,
4300 num_layers: cfg.num_hidden_layers,
4301 hidden_size: cfg.hidden_size,
4302 num_kv_heads: cfg.num_key_value_heads,
4303 num_attn_heads: cfg.num_attention_heads,
4304 sliding_window: cfg.sliding_window,
4305 k_head_dim: cfg.head_dim(),
4306 v_head_dim: cfg.head_dim(),
4307 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
4308 };
4309
4310 Ok(Box::new(cfg))
4311 }
4312
4313 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
4314 Some(vec![NonMappedSubModel::Vision])
4315 }
4316}
4317
4318pub struct VLlama4Loader;
4324
4325pub struct VLlama4Prefixer;
4326
4327impl MultimodalPromptPrefixer for VLlama4Prefixer {
4328 fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
4329 format!(
4330 "{}{prompt}",
4331 llama4::IMAGE_TOKEN.repeat(image_indexes.len())
4332 )
4333 }
4334}
4335
4336impl MultimodalModelLoader for VLlama4Loader {
4337 fn load(
4338 &self,
4339 config: &str,
4340 vb: ShardedVarBuilder,
4341 normal_loading_metadata: NormalLoadingMetadata,
4342 attention_mechanism: AttentionImplementation,
4343 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
4344 let mut cfg: crate::vision_models::llama4::Llama4Config = serde_json::from_str(config)?;
4345 cfg.propagate_quantization_config();
4346 Ok(Box::new(Llama4Model::new(
4347 &cfg,
4348 vb,
4349 self.is_gptx(config),
4350 normal_loading_metadata,
4351 attention_mechanism,
4352 )?))
4353 }
4354 fn is_gptx(&self, _config: &str) -> bool {
4355 false
4356 }
4357 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
4358 let mut cfg: crate::vision_models::llama4::Llama4Config = serde_json::from_str(config)?;
4359 cfg.propagate_quantization_config();
4360 Ok(Box::new(cfg))
4361 }
4362 fn get_processor(
4363 &self,
4364 _model_config: &str,
4365 processor_config: Option<ProcessorConfig>,
4366 _preprocessor_config: PreProcessorConfig,
4367 _max_edge: Option<u32>,
4368 ) -> Arc<dyn Processor + Send + Sync> {
4369 Arc::new(Llama4Processor::new(&processor_config.unwrap()))
4370 }
4371 fn supports_paged_attention(&self, _config: &str) -> bool {
4372 true
4373 }
4374 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
4375 Arc::new(VLlama4Prefixer)
4376 }
4377 fn modalities(&self, _config: &str) -> Result<Modalities> {
4378 Ok(Modalities {
4379 input: vec![SupportedModality::Text, SupportedModality::Vision],
4380 output: vec![SupportedModality::Text],
4381 })
4382 }
4383}
4384
4385impl IsqModelLoader for VLlama4Loader {
4386 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
4387 Ok(vec![
4388 Regex::new(r"lm_head\.(weight|bias)$")?,
4389 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
4391 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
4392 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
4393 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
4394 Regex::new(r"layers\.(\d+)\.feed_forward\.experts\.gate_up_proj\.(weight|bias)$")?,
4396 Regex::new(r"layers\.(\d+)\.feed_forward\.experts\.gate_proj\.(weight|bias)$")?,
4397 Regex::new(r"layers\.(\d+)\.feed_forward\.experts\.up_proj\.(weight|bias)$")?,
4398 Regex::new(r"layers\.(\d+)\.feed_forward\.experts\.down_proj\.(weight|bias)$")?,
4399 Regex::new(r"layers\.(\d+)\.feed_forward\.router\.(weight|bias)$")?,
4400 Regex::new(r"layers\.(\d+)\.feed_forward\.shared_expert\.(weight|bias)$")?,
4401 Regex::new(r"layers\.(\d+)\.feed_forward\.shared_expert\.(weight|bias)$")?,
4402 Regex::new(r"layers\.(\d+)\.feed_forward\.shared_expert\.(weight|bias)$")?,
4403 Regex::new(r"layers\.(\d+)\.feed_forward\.gate_proj\.(weight|bias)$")?,
4405 Regex::new(r"layers\.(\d+)\.feed_forward\.up_proj\.(weight|bias)$")?,
4406 Regex::new(r"layers\.(\d+)\.feed_forward\.down_proj\.(weight|bias)$")?,
4407 ])
4408 }
4409 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
4410 Ok(vec![
4411 Regex::new(r"lm_head\.(weight|bias)$")?,
4412 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
4414 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
4415 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
4416 Regex::new(r"language_model\.model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
4417 Regex::new(
4419 r"language_model\.model\.layers\.(\d+)\.feed_forward\.experts\.(\d+)\.gate_up_proj\.(weight|bias)$",
4420 )?,
4421 Regex::new(
4422 r"language_model\.model\.layers\.(\d+)\.feed_forward\.experts\.(\d+)\.gate_proj\.(weight|bias)$",
4423 )?,
4424 Regex::new(
4425 r"language_model\.model\.layers\.(\d+)\.feed_forward\.experts\.(\d+)\.up_proj\.(weight|bias)$",
4426 )?,
4427 Regex::new(
4428 r"language_model\.model\.layers\.(\d+)\.feed_forward\.experts\.(\d+)\.down_proj\.(weight|bias)$",
4429 )?,
4430 Regex::new(
4431 r"language_model\.model\.layers\.(\d+)\.feed_forward\.router\.(weight|bias)$",
4432 )?,
4433 Regex::new(
4434 r"language_model\.model\.layers\.(\d+)\.feed_forward\.shared_expert\.(weight|bias)$",
4435 )?,
4436 Regex::new(
4437 r"language_model\.model\.layers\.(\d+)\.feed_forward\.shared_expert\.(weight|bias)$",
4438 )?,
4439 Regex::new(
4440 r"language_model\.model\.layers\.(\d+)\.feed_forward\.shared_expert\.(weight|bias)$",
4441 )?,
4442 Regex::new(
4444 r"language_model\.model\.layers\.(\d+)\.feed_forward\.gate_proj\.(weight|bias)$",
4445 )?,
4446 Regex::new(
4447 r"language_model\.model\.layers\.(\d+)\.feed_forward\.up_proj\.(weight|bias)$",
4448 )?,
4449 Regex::new(
4450 r"language_model\.model\.layers\.(\d+)\.feed_forward\.down_proj\.(weight|bias)$",
4451 )?,
4452 ])
4453 }
4454}
4455
4456impl VLlama4Loader {
4457 #[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
4460 fn run_dummy_processing(
4461 &self,
4462 cfg: &Llama4Config,
4463 height: usize,
4464 width: usize,
4465 max_num_images: usize,
4466 max_batch_size: usize,
4467 ) -> Result<(usize, usize)> {
4468 let cfg = &cfg.vision_config;
4469
4470 let img_processor =
4471 Llama4ImageProcessor::new(Some(cfg.patch_size), Some(cfg.pixel_shuffle_ratio));
4472 let image = DynamicImage::new(width as u32, height as u32, ColorType::Rgb8);
4473 let res = img_processor.preprocess(
4474 vec![image; max_num_images],
4475 vec![],
4476 &PreProcessorConfig::default(),
4477 &Device::Cpu,
4478 (max_batch_size, max_num_images),
4479 )?;
4480
4481 let pixels_batch_size = res.pixel_values.dim(0)?;
4482 let pixels_max_batch_size = pixels_batch_size * max_batch_size;
4483
4484 let (image_h, image_w) = (
4485 res.pixel_values.dim(D::Minus2).unwrap(),
4486 res.pixel_values.dim(D::Minus1).unwrap(),
4487 );
4488 let num_patches_per_chunk = (image_h / img_processor.patch_size)
4489 * (image_w / img_processor.patch_size)
4490 / img_processor.downsample_ratio;
4491
4492 Ok((
4493 pixels_max_batch_size,
4494 num_patches_per_chunk * pixels_max_batch_size,
4495 ))
4496 }
4497}
4498
4499impl DeviceMappedModelLoader for VLlama4Loader {
4500 fn mapped_max_act_size_elems(
4501 &self,
4502 config: &str,
4503 params: &AutoDeviceMapParams,
4504 ) -> Result<usize> {
4505 let AutoDeviceMapParams::Multimodal {
4506 max_seq_len,
4507 max_batch_size,
4508 max_image_shape: (height, width),
4509 max_num_images,
4510 } = params
4511 else {
4512 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
4513 };
4514
4515 let cfg: Llama4Config = serde_json::from_str(config)?;
4516
4517 let (_pixels_batch_size, num_text_image_toks) =
4518 self.run_dummy_processing(&cfg, *height, *width, *max_num_images, *max_batch_size)?;
4519
4520 let max_seq_len = max_seq_len.min(&ATTENTION_CHUNK_SIZE) + num_text_image_toks;
4521
4522 Ok(max_batch_size * cfg.text_config.num_attention_heads * max_seq_len * max_seq_len)
4523 }
4524 fn non_mapped_max_act_size_elems(
4525 &self,
4526 config: &str,
4527 params: &AutoDeviceMapParams,
4528 ) -> Result<usize> {
4529 let AutoDeviceMapParams::Multimodal {
4530 max_seq_len: _,
4531 max_batch_size,
4532 max_image_shape: (height, width),
4533 max_num_images,
4534 } = params
4535 else {
4536 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
4537 };
4538
4539 let cfg: Llama4Config = serde_json::from_str(config)?;
4540
4541 let (pixels_batch_size, _num_text_image_toks) =
4542 self.run_dummy_processing(&cfg, *height, *width, *max_num_images, *max_batch_size)?;
4543 let max_seq_len = cfg.vision_config.num_patches();
4544
4545 Ok((max_batch_size * pixels_batch_size)
4546 * cfg.vision_config.num_attention_heads
4547 * max_seq_len
4548 * max_seq_len)
4549 }
4550
4551 fn non_mapped_size_in_bytes(
4552 &self,
4553 config: &str,
4554 dtype: DType,
4555 weight_pack_factor: usize,
4556 _matformer_config: Option<&MatformerSliceConfig>,
4557 ) -> Result<usize> {
4558 let cfg: Llama4Config = serde_json::from_str(config)?;
4559 let tcfg = &cfg.text_config;
4560
4561 let text_elems = {
4562 let embed_tokens = tcfg.hidden_size * tcfg.vocab_size / weight_pack_factor;
4563 let lm_head = if !tcfg.tie_word_embeddings {
4564 tcfg.hidden_size * tcfg.vocab_size
4565 } else {
4566 0
4567 };
4568 let norm = tcfg.hidden_size;
4569 embed_tokens + lm_head + norm
4570 };
4571
4572 let vision_elems = {
4573 let cfg = &cfg.vision_config;
4574
4575 let num_patches = cfg.num_patches();
4576
4577 let unfold_elems =
4578 (cfg.num_channels * cfg.patch_size * cfg.patch_size) * cfg.hidden_size;
4579 let class_embeddng_elems = cfg.hidden_size;
4580 let positional_embedding_vlm_elems = num_patches * cfg.hidden_size;
4581 let layernorm_pre_elems = cfg.hidden_size;
4582 let layernorm_post_elems = cfg.hidden_size;
4583
4584 let pixel_shuffle_elems = cfg.intermediate_size * cfg.projector_input_dim
4585 / weight_pack_factor
4586 + cfg.projector_input_dim * cfg.projector_output_dim / weight_pack_factor;
4587
4588 let encoder_layer = {
4589 let input_layernorm = cfg.hidden_size + cfg.hidden_size;
4590 let post_attention_layernorm = cfg.hidden_size + cfg.hidden_size;
4591
4592 let head_dim = cfg.hidden_size / cfg.num_attention_heads;
4593 let q_proj = cfg.hidden_size * cfg.num_attention_heads * head_dim
4594 / weight_pack_factor
4595 + cfg.num_attention_heads * head_dim;
4596 let k_proj = cfg.hidden_size * cfg.num_attention_heads * head_dim
4597 / weight_pack_factor
4598 + cfg.num_attention_heads * head_dim;
4599 let v_proj = cfg.hidden_size * cfg.num_attention_heads * head_dim
4600 / weight_pack_factor
4601 + cfg.num_attention_heads * head_dim;
4602 let o_proj = cfg.hidden_size * cfg.num_attention_heads * head_dim
4603 / weight_pack_factor
4604 + cfg.num_attention_heads * head_dim;
4605
4606 let fc1 = (cfg.hidden_size * cfg.intermediate_size) / weight_pack_factor
4607 + cfg.intermediate_size;
4608 let fc2 = (cfg.intermediate_size * cfg.hidden_size) / weight_pack_factor
4609 + cfg.hidden_size;
4610
4611 input_layernorm
4612 + post_attention_layernorm
4613 + q_proj
4614 + k_proj
4615 + v_proj
4616 + o_proj
4617 + fc1
4618 + fc2
4619 };
4620
4621 unfold_elems
4622 + class_embeddng_elems
4623 + positional_embedding_vlm_elems
4624 + layernorm_post_elems
4625 + layernorm_pre_elems
4626 + pixel_shuffle_elems
4627 + encoder_layer * cfg.num_hidden_layers
4628 };
4629
4630 let elems = text_elems + vision_elems;
4631
4632 Ok(elems * dtype.size_in_bytes())
4633 }
4634
4635 fn layer_sizes_in_bytes(
4636 &self,
4637 config: &str,
4638 dtype: DType,
4639 weight_pack_factor: usize,
4640 _matformer_config: Option<&MatformerSliceConfig>,
4641 ) -> Result<Vec<usize>> {
4642 let cfg: Llama4Config = serde_json::from_str(config)?;
4643 let tcfg = &cfg.text_config;
4644
4645 let mut per_layer_elems = Vec::new();
4646
4647 for layer_idx in 0..tcfg.num_hidden_layers {
4648 let input_layernorm = tcfg.hidden_size;
4649 let post_attention_layernorm = tcfg.hidden_size;
4650
4651 let size_in = tcfg.hidden_size;
4652 let size_q = (tcfg.hidden_size / tcfg.num_attention_heads) * tcfg.num_attention_heads;
4653 let size_kv = (tcfg.hidden_size / tcfg.num_attention_heads) * tcfg.num_key_value_heads;
4654 let q_proj = size_in * size_q / weight_pack_factor;
4655 let k_proj = size_in * size_kv / weight_pack_factor;
4656 let v_proj = size_in * size_kv / weight_pack_factor;
4657 let o_proj = size_q * size_in / weight_pack_factor;
4658
4659 let use_moe = tcfg.moe_layers().contains(&layer_idx);
4660 let moe_block = if use_moe {
4661 let h_size = tcfg.hidden_size;
4662 let i_size = tcfg.intermediate_size;
4663 let gate_proj = tcfg.num_local_experts * h_size * i_size / weight_pack_factor;
4664 let up_proj = tcfg.num_local_experts * h_size * i_size / weight_pack_factor;
4665 let down_proj = tcfg.num_local_experts * i_size * h_size / weight_pack_factor;
4666
4667 gate_proj + up_proj + down_proj
4668 } else {
4669 let h_size = tcfg.hidden_size;
4670 let i_size = tcfg.intermediate_size_mlp;
4671 let gate_proj = h_size * i_size / weight_pack_factor;
4672 let up_proj = h_size * i_size / weight_pack_factor;
4673 let down_proj = i_size * h_size / weight_pack_factor;
4674
4675 gate_proj + up_proj + down_proj
4676 };
4677
4678 per_layer_elems.push(
4679 input_layernorm
4680 + post_attention_layernorm
4681 + q_proj
4682 + k_proj
4683 + v_proj
4684 + o_proj
4685 + moe_block,
4686 );
4687 }
4688
4689 Ok(per_layer_elems
4690 .into_iter()
4691 .map(|x| x * dtype.size_in_bytes())
4692 .collect())
4693 }
4694
4695 fn num_layers(&self, config: &str) -> Result<usize> {
4696 let cfg: Llama4Config = serde_json::from_str(config)?;
4697 Ok(cfg.text_config.num_hidden_layers)
4698 }
4699
4700 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
4701 let cfg: Llama4Config = serde_json::from_str(config)?;
4702 let cfg = &cfg.text_config;
4703
4704 let cfg = ModelConfigMetadata {
4705 max_seq_len: cfg.max_position_embeddings,
4706 num_layers: cfg.num_hidden_layers,
4707 hidden_size: cfg.hidden_size,
4708 num_kv_heads: cfg.num_attention_heads,
4709 num_attn_heads: cfg.num_attention_heads,
4710 sliding_window: None,
4711 k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
4712 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
4713 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
4714 };
4715
4716 Ok(Box::new(cfg))
4717 }
4718
4719 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
4720 Some(vec![NonMappedSubModel::Vision])
4721 }
4722}
4723
4724pub struct Gemma3nLoader;
4730
4731#[allow(dead_code)]
4732pub struct Gemma3nPrefixer;
4733
4734impl MultimodalPromptPrefixer for Gemma3nPrefixer {
4735 fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
4736 prompt.to_string()
4737 }
4738}
4739
4740impl MultimodalModelLoader for Gemma3nLoader {
4741 fn load(
4742 &self,
4743 config: &str,
4744 vb: ShardedVarBuilder,
4745 normal_loading_metadata: NormalLoadingMetadata,
4746 attention_mechanism: AttentionImplementation,
4747 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
4748 let cfg: Gemma3nConfig = serde_json::from_str(config)?;
4749 Ok(Box::new(Gemma3nModel::new(
4750 &cfg,
4751 vb,
4752 self.is_gptx(config),
4753 normal_loading_metadata,
4754 attention_mechanism,
4755 )?))
4756 }
4757 fn is_gptx(&self, _config: &str) -> bool {
4758 true
4759 }
4760 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
4761 let config: Gemma3nConfig = serde_json::from_str(config)?;
4762 Ok(Box::new(config))
4763 }
4764 fn get_processor(
4765 &self,
4766 _config: &str,
4767 processor_config: Option<ProcessorConfig>,
4768 _preprocessor_config: PreProcessorConfig,
4769 _max_edge: Option<u32>,
4770 ) -> Arc<dyn Processor + Send + Sync> {
4771 Arc::new(Gemma3nProcessor::new(
4773 processor_config.unwrap_or_default(),
4774 true,
4775 ))
4776 }
4777 fn supports_paged_attention(&self, _config: &str) -> bool {
4778 false
4779 }
4780 fn supports_prefix_cacher(&self, _config: &str) -> bool {
4781 true
4782 }
4783 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
4784 Arc::new(Gemma3Prefixer)
4785 }
4786 fn modalities(&self, _config: &str) -> Result<Modalities> {
4787 Ok(Modalities {
4788 input: vec![
4789 SupportedModality::Text,
4790 SupportedModality::Vision,
4791 SupportedModality::Audio,
4792 ],
4793 output: vec![SupportedModality::Text],
4794 })
4795 }
4796}
4797
4798impl IsqModelLoader for Gemma3nLoader {
4799 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
4800 Ok(vec![
4801 Regex::new(r"lm_head\.(weight|bias)$")?,
4802 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
4804 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
4805 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
4806 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
4807 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
4809 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
4810 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
4811 Regex::new(r"conformer\.(\d+)\.attention\.attn\.q_proj\.(weight|bias)$")?,
4813 Regex::new(r"conformer\.(\d+)\.attention\.attn\.k_proj\.(weight|bias)$")?,
4814 Regex::new(r"conformer\.(\d+)\.attention\.attn\.v_proj\.(weight|bias)$")?,
4815 Regex::new(
4816 r"conformer\.(\d+)\.attention\.attn\.relative_position_embedding\.pos_proj\.(weight|bias)$",
4817 )?,
4818 Regex::new(r"conformer\.(\d+)\.attention\.post\.(weight|bias)$")?,
4819 Regex::new(r"conformer\.(\d+)\.ffw_layer_start\.ffw_layer_1\.(weight|bias)$")?,
4821 Regex::new(r"conformer\.(\d+)\.ffw_layer_start\.ffw_layer_2\.(weight|bias)$")?,
4822 Regex::new(r"conformer\.(\d+)\.ffw_layer_end\.ffw_layer_1\.(weight|bias)$")?,
4823 Regex::new(r"conformer\.(\d+)\.ffw_layer_end\.ffw_layer_2\.(weight|bias)$")?,
4824 Regex::new(r"conformer\.(\d+)\.lconv1d\.linear_start\.(weight|bias)$")?,
4826 Regex::new(r"conformer\.(\d+)\.lconv1d\.linear_end\.(weight|bias)$")?,
4827 Regex::new(r"subsample_conv_projection\.input_proj_linear\.(weight|bias)$")?,
4829 Regex::new(r"embed_vision\.embedding_projection\.(weight|bias)$")?,
4831 Regex::new(r"embed_audio\.embedding_projection\.(weight|bias)$")?,
4832 ])
4833 }
4834 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
4835 Ok(vec![
4836 Regex::new(r"lm_head\.(weight|bias)$")?,
4837 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
4839 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
4840 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
4841 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
4842 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
4844 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
4845 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
4846 Regex::new(r"model\.language_model\.per_layer_model_projection\.(weight|bias)$")?,
4848 Regex::new(r"model\.language_model\.altup_projections\.(\d+)\.(weight|bias)$")?,
4849 Regex::new(r"model\.language_model\.altup_unembed_projections\.(\d+)\.(weight|bias)$")?,
4850 Regex::new(
4852 r"model\.audio_tower\.conformer\.(\d+)\.attention\.attn\.q_proj\.(weight|bias)$",
4853 )?,
4854 Regex::new(
4855 r"model\.audio_tower\.conformer\.(\d+)\.attention\.attn\.k_proj\.(weight|bias)$",
4856 )?,
4857 Regex::new(
4858 r"model\.audio_tower\.conformer\.(\d+)\.attention\.attn\.v_proj\.(weight|bias)$",
4859 )?,
4860 Regex::new(
4861 r"model\.audio_tower\.conformer\.(\d+)\.attention\.attn\.relative_position_embedding\.pos_proj\.(weight|bias)$",
4862 )?,
4863 Regex::new(r"model\.audio_tower\.conformer\.(\d+)\.attention\.post\.(weight|bias)$")?,
4864 Regex::new(
4866 r"model\.audio_tower\.conformer\.(\d+)\.ffw_layer_start\.ffw_layer_1\.(weight|bias)$",
4867 )?,
4868 Regex::new(
4869 r"model\.audio_tower\.conformer\.(\d+)\.ffw_layer_start\.ffw_layer_2\.(weight|bias)$",
4870 )?,
4871 Regex::new(
4872 r"model\.audio_tower\.conformer\.(\d+)\.ffw_layer_end\.ffw_layer_1\.(weight|bias)$",
4873 )?,
4874 Regex::new(
4875 r"model\.audio_tower\.conformer\.(\d+)\.ffw_layer_end\.ffw_layer_2\.(weight|bias)$",
4876 )?,
4877 Regex::new(
4879 r"model\.audio_tower\.conformer\.(\d+)\.lconv1d\.linear_start\.(weight|bias)$",
4880 )?,
4881 Regex::new(
4882 r"model\.audio_tower\.conformer\.(\d+)\.lconv1d\.linear_end\.(weight|bias)$",
4883 )?,
4884 Regex::new(
4886 r"model\.audio_tower\.subsample_conv_projection\.input_proj_linear\.(weight|bias)$",
4887 )?,
4888 Regex::new(r"model\.embed_vision\.embedding_projection\.(weight|bias)$")?,
4890 Regex::new(r"model\.embed_audio\.embedding_projection\.(weight|bias)$")?,
4891 ])
4892 }
4893}
4894
4895impl DeviceMappedModelLoader for Gemma3nLoader {
4896 fn mapped_max_act_size_elems(
4897 &self,
4898 config: &str,
4899 params: &AutoDeviceMapParams,
4900 ) -> Result<usize> {
4901 let AutoDeviceMapParams::Multimodal {
4902 max_seq_len,
4903 max_batch_size,
4904 max_image_shape: _,
4905 max_num_images,
4906 } = params
4907 else {
4908 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
4909 };
4910
4911 let cfg: Gemma3nConfig = serde_json::from_str(config)?;
4912 let text_cfg = &cfg.text_config;
4913
4914 let mut total_seq_len = *max_seq_len.min(&ATTENTION_CHUNK_SIZE);
4918
4919 {
4921 let msfa_spatial_size = 16; let vision_tokens_per_image = msfa_spatial_size * msfa_spatial_size; total_seq_len += vision_tokens_per_image * max_num_images;
4926 }
4927
4928 {
4930 let audio_tokens = cfg.audio_soft_tokens_per_image;
4933 total_seq_len += audio_tokens;
4934 }
4935
4936 let max_text_attn =
4938 max_batch_size * text_cfg.num_attention_heads * total_seq_len * total_seq_len;
4939
4940 Ok(max_text_attn)
4941 }
4942
4943 fn non_mapped_max_act_size_elems(
4944 &self,
4945 config: &str,
4946 params: &AutoDeviceMapParams,
4947 ) -> Result<usize> {
4948 let AutoDeviceMapParams::Multimodal {
4949 max_seq_len: _,
4950 max_batch_size,
4951 max_image_shape: _,
4952 max_num_images,
4953 } = params
4954 else {
4955 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
4956 };
4957
4958 let cfg: Gemma3nConfig = serde_json::from_str(config)?;
4959
4960 let mut max_activation = 0;
4962
4963 {
4965 let vision_tower_act = {
4975 let num_heads = 16; let spatial_size = 24; let seq_len = spatial_size * spatial_size;
4981
4982 max_batch_size * max_num_images * num_heads * seq_len * seq_len
4984 };
4985
4986 let vision_embed_act = {
4988 let msfa_channels = 2048; let spatial_size = 16; let vision_features =
4992 max_batch_size * max_num_images * msfa_channels * spatial_size * spatial_size;
4993
4994 let projected = max_batch_size
4996 * max_num_images
4997 * spatial_size
4998 * spatial_size
4999 * cfg.text_config.hidden_size;
5000
5001 vision_features.max(projected)
5002 };
5003
5004 max_activation = max_activation.max(vision_tower_act).max(vision_embed_act);
5005 }
5006
5007 {
5009 let audio_cfg = &cfg.audio_config;
5010
5011 let max_audio_frames = 1280;
5016
5017 let subsample_factor: usize = audio_cfg
5018 .sscp_conv_stride_size
5019 .iter()
5020 .map(|stride| stride[0]) .product();
5022 let audio_seq_after_subsample = max_audio_frames / subsample_factor;
5023
5024 let audio_encoder_act = {
5026 let intermediate_size = audio_cfg.hidden_size * 4; max_batch_size * audio_seq_after_subsample * intermediate_size
5031 };
5032
5033 let audio_attn_act = {
5035 let chunk_size = audio_cfg.conf_attention_chunk_size;
5037 let context_size = chunk_size + audio_cfg.conf_attention_context_left - 1
5038 + audio_cfg.conf_attention_context_right;
5039
5040 let num_chunks = audio_seq_after_subsample.div_ceil(chunk_size);
5042
5043 max_batch_size
5044 * audio_cfg.conf_num_attention_heads
5045 * num_chunks
5046 * chunk_size
5047 * context_size
5048 };
5049
5050 max_activation = max_activation.max(audio_encoder_act).max(audio_attn_act);
5051 }
5052
5053 Ok(max_activation)
5054 }
5055
5056 fn non_mapped_size_in_bytes(
5057 &self,
5058 config: &str,
5059 dtype: DType,
5060 weight_pack_factor: usize,
5061 matformer_config: Option<&MatformerSliceConfig>,
5062 ) -> Result<usize> {
5063 let cfg: Gemma3nConfig = serde_json::from_str(config)?;
5064
5065 let text_cfg = if let Some(matformer_cfg) = matformer_config {
5067 use crate::device_map::DummyDeviceMapper;
5068 use crate::vision_models::gemma3n::text::handle_matformer_slicing;
5069
5070 let dummy_mapper = DummyDeviceMapper {
5071 nm_device: Device::Cpu,
5072 };
5073 let (adjusted_cfg, _, _, _, _) = handle_matformer_slicing(
5074 &cfg.text_config,
5075 &Some(matformer_cfg.clone()),
5076 &dummy_mapper,
5077 )?;
5078 adjusted_cfg
5079 } else {
5080 cfg.text_config.clone()
5081 };
5082
5083 let text_cfg = &text_cfg;
5084
5085 let text_elems = {
5087 let embed_tokens = text_cfg.hidden_size * text_cfg.vocab_size;
5089 let embed_tokens_per_layer = text_cfg.num_hidden_layers
5090 * text_cfg.hidden_size_per_layer_input
5091 * text_cfg.vocab_size_per_layer_input;
5092
5093 let lm_head = if !text_cfg.tie_word_embeddings || weight_pack_factor != 1 {
5095 text_cfg.hidden_size * text_cfg.vocab_size / weight_pack_factor
5096 } else {
5097 0
5098 };
5099
5100 let norm = text_cfg.hidden_size;
5102
5103 let altup_projections =
5105 (text_cfg.altup_num_inputs - 1) * text_cfg.hidden_size * text_cfg.hidden_size
5106 / weight_pack_factor;
5107 let altup_unembed_projections =
5108 (text_cfg.altup_num_inputs - 1) * text_cfg.hidden_size * text_cfg.hidden_size
5109 / weight_pack_factor;
5110
5111 let per_layer_model_projection = text_cfg.num_hidden_layers
5113 * text_cfg.hidden_size
5114 * text_cfg.hidden_size_per_layer_input
5115 / weight_pack_factor;
5116 let per_layer_projection_norm = text_cfg.hidden_size;
5117
5118 embed_tokens
5119 + embed_tokens_per_layer
5120 + lm_head
5121 + norm
5122 + altup_projections
5123 + altup_unembed_projections
5124 + per_layer_model_projection
5125 + per_layer_projection_norm
5126 };
5127
5128 let vision_elems = {
5130 let multimodal_cfg = &cfg.vision_config;
5131 let vision_tower_elems = {
5135 use crate::vision_models::gemma3n::vision::{
5136 gemma3n_mobilenet_def, make_divisible, BlockType, INPUT_CHANNELS,
5137 MSFA_EXPANSION_RATIO, MSFA_IN_CHANNELS, MSFA_OUT_CHANNELS, STEM_KERNEL_SIZE,
5138 STEM_OUT_CHANNELS,
5139 };
5140
5141 let stem_conv =
5143 INPUT_CHANNELS * STEM_OUT_CHANNELS * STEM_KERNEL_SIZE * STEM_KERNEL_SIZE;
5144 let stem_norm = STEM_OUT_CHANNELS; let mut in_chs = STEM_OUT_CHANNELS;
5148 let mut total_elems = stem_conv + stem_norm;
5149
5150 let block_defs = gemma3n_mobilenet_def();
5152
5153 for stage_blocks in block_defs.iter() {
5154 for block_type in stage_blocks.iter() {
5155 match block_type {
5156 BlockType::EdgeResidual {
5157 out_channels,
5158 kernel_size,
5159 stride: _,
5160 expand_ratio,
5161 ..
5162 } => {
5163 #[allow(clippy::cast_precision_loss)]
5164 let mid_chs = make_divisible(in_chs as f64 * expand_ratio, 8);
5165 total_elems += in_chs * mid_chs * kernel_size * kernel_size; total_elems += mid_chs; total_elems += mid_chs * out_channels; total_elems += out_channels; in_chs = *out_channels;
5171 }
5172 BlockType::UniversalInvertedResidual {
5173 out_channels,
5174 start_kernel_size,
5175 mid_kernel_size,
5176 stride: _,
5177 expand_ratio,
5178 ..
5179 } => {
5180 #[allow(clippy::cast_precision_loss)]
5181 let mid_chs = make_divisible(in_chs as f64 * expand_ratio, 8);
5182 if *expand_ratio != 1.0 {
5184 total_elems += in_chs * mid_chs; total_elems += mid_chs; }
5187 if *start_kernel_size > 0 {
5188 total_elems += mid_chs * start_kernel_size * start_kernel_size; total_elems += mid_chs; }
5191 if *mid_kernel_size > 0 {
5192 total_elems += mid_chs * mid_kernel_size * mid_kernel_size; total_elems += mid_chs; }
5195 total_elems += mid_chs * out_channels; total_elems += out_channels; total_elems += out_channels; in_chs = *out_channels;
5199 }
5200 BlockType::MultiQueryAttention {
5201 num_heads,
5202 kv_dim,
5203 kv_stride: _,
5204 ..
5205 } => {
5206 let dw_kernel_size = 3; total_elems += in_chs; total_elems += in_chs * num_heads * kv_dim; total_elems += in_chs * kv_dim; total_elems += in_chs * dw_kernel_size * dw_kernel_size; total_elems += *kv_dim; total_elems += 1; total_elems += *kv_dim; total_elems += num_heads * kv_dim * in_chs; total_elems += in_chs; }
5218 }
5219 }
5220 }
5221
5222 let msfa_in = MSFA_IN_CHANNELS.iter().sum::<usize>();
5224 let msfa_out = MSFA_OUT_CHANNELS;
5225 #[allow(clippy::cast_precision_loss)]
5226 let msfa_mid = make_divisible(msfa_in as f64 * MSFA_EXPANSION_RATIO, 8);
5227
5228 total_elems += msfa_in * msfa_mid; total_elems += msfa_mid; total_elems += msfa_mid * msfa_out; total_elems += msfa_out; total_elems += msfa_out; total_elems
5236 };
5237
5238 let embed_vision_elems = {
5240 let embedding = multimodal_cfg.vocab_size * multimodal_cfg.hidden_size;
5242
5243 let hard_norm = multimodal_cfg.hidden_size;
5245 let soft_norm = multimodal_cfg.hidden_size;
5246
5247 let projection =
5249 multimodal_cfg.hidden_size * text_cfg.hidden_size / weight_pack_factor;
5250
5251 let post_norm = text_cfg.hidden_size;
5253
5254 embedding + hard_norm + soft_norm + projection + post_norm
5255 };
5256
5257 vision_tower_elems + embed_vision_elems
5258 };
5259
5260 let audio_elems = {
5262 let audio_cfg = &cfg.audio_config;
5263
5264 let subsample_conv_projection_elems = {
5266 let mut conv_elems = 0;
5268
5269 let in_ch_0 = 1;
5271 let out_ch_0 = audio_cfg.sscp_conv_channel_size[0];
5272 let kernel_0 = &audio_cfg.sscp_conv_kernel_size[0];
5273 conv_elems += in_ch_0 * out_ch_0 * kernel_0[0] * kernel_0[1];
5274
5275 let in_ch_1 = out_ch_0;
5277 let out_ch_1 = audio_cfg.sscp_conv_channel_size[1];
5278 let kernel_1 = &audio_cfg.sscp_conv_kernel_size[1];
5279 conv_elems += in_ch_1 * out_ch_1 * kernel_1[0] * kernel_1[1];
5280
5281 let norm_0 = out_ch_0; let norm_1 = out_ch_1; let mut f_out = audio_cfg.input_feat_size;
5287 for i in 0..2 {
5288 let kernel_w = audio_cfg.sscp_conv_kernel_size[i][1];
5289 let stride_w = audio_cfg.sscp_conv_stride_size[i][1];
5290 let pad_left = 1;
5291 let pad_right = 1;
5292 f_out = (f_out + pad_left + pad_right + stride_w - kernel_w) / stride_w;
5293 }
5294 let input_proj_in_features = out_ch_1 * f_out;
5295 let input_proj_linear =
5296 input_proj_in_features * audio_cfg.hidden_size / weight_pack_factor;
5297
5298 conv_elems + norm_0 + norm_1 + input_proj_linear
5299 };
5300
5301 let conformer_elems = {
5303 let mut total = 0;
5304
5305 for _ in 0..audio_cfg.conf_num_hidden_layers {
5306 let attention_elems = {
5308 let pre_attn_norm = audio_cfg.hidden_size;
5310 let post_norm = audio_cfg.hidden_size;
5311
5312 let q_proj =
5314 audio_cfg.hidden_size * audio_cfg.hidden_size / weight_pack_factor;
5315 let k_proj =
5316 audio_cfg.hidden_size * audio_cfg.hidden_size / weight_pack_factor;
5317 let v_proj =
5318 audio_cfg.hidden_size * audio_cfg.hidden_size / weight_pack_factor;
5319 let post =
5320 audio_cfg.hidden_size * audio_cfg.hidden_size / weight_pack_factor;
5321
5322 let pos_proj =
5324 audio_cfg.hidden_size * audio_cfg.hidden_size / weight_pack_factor;
5325 let per_dim_scale =
5326 audio_cfg.hidden_size / audio_cfg.conf_num_attention_heads; let inv_timescales = audio_cfg.hidden_size / 2; let pos_indices = audio_cfg.conf_attention_context_left
5329 + audio_cfg.conf_attention_context_right
5330 + 1;
5331
5332 let chunk_size = audio_cfg.conf_attention_chunk_size;
5334 let context_size = chunk_size + audio_cfg.conf_attention_context_left - 1
5335 + audio_cfg.conf_attention_context_right;
5336 let local_causal_valid_mask = chunk_size * context_size; let invalid_logits_tensor = 1; pre_attn_norm
5340 + post_norm
5341 + q_proj
5342 + k_proj
5343 + v_proj
5344 + post
5345 + pos_proj
5346 + per_dim_scale
5347 + inv_timescales
5348 + pos_indices
5349 + local_causal_valid_mask
5350 + invalid_logits_tensor
5351 };
5352
5353 let ffw_elems = {
5355 let intermediate_size = audio_cfg.hidden_size * 4;
5361
5362 let ffw_start = {
5363 let pre_norm = audio_cfg.hidden_size;
5364 let layer_1 =
5365 audio_cfg.hidden_size * intermediate_size / weight_pack_factor;
5366 let layer_2 =
5367 intermediate_size * audio_cfg.hidden_size / weight_pack_factor;
5368 let post_norm = audio_cfg.hidden_size;
5369 pre_norm + layer_1 + layer_2 + post_norm
5370 };
5371
5372 let ffw_end = ffw_start; ffw_start + ffw_end
5375 };
5376
5377 let lconv1d_elems = {
5379 let pre_layer_norm = audio_cfg.hidden_size;
5381 let conv_norm = audio_cfg.hidden_size;
5382
5383 let linear_start = audio_cfg.hidden_size * (audio_cfg.hidden_size * 2)
5385 / weight_pack_factor;
5386 let linear_end =
5387 audio_cfg.hidden_size * audio_cfg.hidden_size / weight_pack_factor;
5388
5389 let depthwise = audio_cfg.hidden_size * audio_cfg.conf_conv_kernel_size;
5391
5392 pre_layer_norm + conv_norm + linear_start + linear_end + depthwise
5393 };
5394
5395 let block_norm = audio_cfg.hidden_size;
5397
5398 total += attention_elems + ffw_elems + lconv1d_elems + block_norm;
5399 }
5400
5401 total
5402 };
5403
5404 let embed_audio_elems = {
5406 let embedding = audio_cfg.vocab_size * audio_cfg.hidden_size;
5408
5409 let hard_embedding_norm = audio_cfg.hidden_size; let soft_embedding_norm = audio_cfg.hidden_size; let embedding_post_projection_norm = text_cfg.hidden_size; let embedding_projection =
5416 audio_cfg.hidden_size * text_cfg.hidden_size / weight_pack_factor;
5417
5418 embedding
5419 + hard_embedding_norm
5420 + soft_embedding_norm
5421 + embedding_post_projection_norm
5422 + embedding_projection
5423 };
5424
5425 subsample_conv_projection_elems + conformer_elems + embed_audio_elems
5426 };
5427
5428 let vision_dtype = if dtype == DType::F16 {
5429 DType::F32
5431 } else {
5432 dtype
5433 };
5434
5435 let total_elems = text_elems * dtype.size_in_bytes()
5436 + vision_elems * vision_dtype.size_in_bytes()
5437 + audio_elems * dtype.size_in_bytes();
5438
5439 Ok(total_elems)
5440 }
5441
5442 fn layer_sizes_in_bytes(
5443 &self,
5444 config: &str,
5445 dtype: DType,
5446 weight_pack_factor: usize,
5447 matformer_config: Option<&MatformerSliceConfig>,
5448 ) -> Result<Vec<usize>> {
5449 let cfg: Gemma3nConfig = serde_json::from_str(config)?;
5450
5451 let (text_cfg, _layer_rename_map, _layers_skipped) = if let Some(matformer_cfg) =
5453 matformer_config
5454 {
5455 use crate::device_map::DummyDeviceMapper;
5456 use crate::vision_models::gemma3n::text::handle_matformer_slicing;
5457
5458 let dummy_mapper = DummyDeviceMapper {
5459 nm_device: Device::Cpu,
5460 };
5461 let (adjusted_cfg, _, _, layer_rename_map, layers_skipped) = handle_matformer_slicing(
5462 &cfg.text_config,
5463 &Some(matformer_cfg.clone()),
5464 &dummy_mapper,
5465 )?;
5466 (adjusted_cfg, layer_rename_map, layers_skipped)
5467 } else {
5468 (cfg.text_config.clone(), None, None)
5469 };
5470
5471 let text_cfg = &text_cfg;
5472
5473 let mut layer_sizes = Vec::new();
5475
5476 for layer_idx in 0..text_cfg.num_hidden_layers {
5480 let per_layer_elems = {
5481 let input_layernorm = text_cfg.hidden_size;
5483 let post_attention_layernorm = text_cfg.hidden_size;
5484 let pre_feedforward_layernorm = text_cfg.hidden_size;
5485 let post_feedforward_layernorm = text_cfg.hidden_size;
5486 let post_per_layer_input_norm = text_cfg.hidden_size;
5487
5488 let size_in = text_cfg.hidden_size;
5490 let size_q = text_cfg.num_attention_heads * text_cfg.head_dim;
5491 let size_kv = text_cfg.num_key_value_heads * text_cfg.head_dim;
5492
5493 let q_proj = size_in * size_q / weight_pack_factor;
5494 let k_proj = size_in * size_kv / weight_pack_factor;
5495 let v_proj = size_in * size_kv / weight_pack_factor;
5496 let o_proj = size_q * size_in / weight_pack_factor;
5497
5498 let q_norm = text_cfg.head_dim;
5500 let k_norm = text_cfg.head_dim;
5501 let v_norm = text_cfg.head_dim; let intermediate_size = match &text_cfg.intermediate_size {
5505 IntermediateSize::Single(size) => *size,
5506 IntermediateSize::PerLayer(sizes) => sizes[layer_idx],
5507 IntermediateSize::Matformer(sizes, _) => sizes[layer_idx],
5508 };
5509 let gate_proj = text_cfg.hidden_size * intermediate_size / weight_pack_factor;
5510 let up_proj = text_cfg.hidden_size * intermediate_size / weight_pack_factor;
5511 let down_proj = intermediate_size * text_cfg.hidden_size / weight_pack_factor;
5512
5513 let altup_elems = {
5515 let correct_output_scale = text_cfg.hidden_size;
5516 let correction_coefs = text_cfg.altup_num_inputs * text_cfg.altup_num_inputs;
5517 let prediction_coefs =
5518 text_cfg.altup_num_inputs * text_cfg.altup_num_inputs.pow(2);
5519 let modality_router = text_cfg.hidden_size * text_cfg.altup_num_inputs;
5520 let router_norm = text_cfg.hidden_size;
5521
5522 correct_output_scale
5523 + correction_coefs
5524 + prediction_coefs
5525 + modality_router
5526 + router_norm
5527 };
5528
5529 let laurel_elems = {
5531 let left = text_cfg.hidden_size * text_cfg.laurel_rank;
5532 let right = text_cfg.laurel_rank * text_cfg.hidden_size;
5533 let post_norm = text_cfg.hidden_size;
5534
5535 left + right + post_norm
5536 };
5537
5538 let per_layer_input_gate =
5540 text_cfg.hidden_size * text_cfg.hidden_size_per_layer_input;
5541 let per_layer_projection =
5542 text_cfg.hidden_size_per_layer_input * text_cfg.hidden_size;
5543
5544 input_layernorm
5545 + post_attention_layernorm
5546 + pre_feedforward_layernorm
5547 + post_feedforward_layernorm
5548 + post_per_layer_input_norm
5549 + q_proj
5550 + k_proj
5551 + v_proj
5552 + o_proj
5553 + q_norm
5554 + k_norm
5555 + v_norm
5556 + gate_proj
5557 + up_proj
5558 + down_proj
5559 + altup_elems
5560 + laurel_elems
5561 + per_layer_input_gate
5562 + per_layer_projection
5563 };
5564
5565 layer_sizes.push(per_layer_elems * dtype.size_in_bytes());
5566 }
5567
5568 Ok(layer_sizes)
5569 }
5570
5571 fn num_layers(&self, config: &str) -> Result<usize> {
5572 let cfg: Gemma3nConfig = serde_json::from_str(config)?;
5573 Ok(cfg.text_config.num_hidden_layers)
5574 }
5575
5576 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
5577 let cfg: Gemma3nConfig = serde_json::from_str(config)?;
5578 let cfg = cfg.text_config;
5579
5580 let cfg = ModelConfigMetadata {
5581 max_seq_len: cfg.max_position_embeddings,
5582 num_layers: cfg.num_hidden_layers,
5583 hidden_size: cfg.hidden_size,
5584 num_kv_heads: cfg.num_key_value_heads,
5585 num_attn_heads: cfg.num_attention_heads,
5586 sliding_window: None, k_head_dim: cfg.hidden_size / cfg.num_attention_heads,
5588 v_head_dim: cfg.hidden_size / cfg.num_attention_heads,
5589 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
5590 };
5591
5592 Ok(Box::new(cfg))
5593 }
5594
5595 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
5596 Some(vec![NonMappedSubModel::Vision, NonMappedSubModel::Audio])
5597 }
5598}
5599
5600pub struct Qwen3VLLoader;
5606
5607pub struct Qwen3VLPrefixer;
5608
5609impl MultimodalPromptPrefixer for Qwen3VLPrefixer {
5610 }
5613
5614impl MultimodalModelLoader for Qwen3VLLoader {
5615 fn load(
5616 &self,
5617 config: &str,
5618 vb: ShardedVarBuilder,
5619 normal_loading_metadata: NormalLoadingMetadata,
5620 attention_mechanism: AttentionImplementation,
5621 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
5622 let cfg: Qwen3VLConfig = serde_json::from_str(config)?;
5623 Ok(Box::new(Qwen3VLModel::new(
5624 &cfg,
5625 vb,
5626 self.is_gptx(config),
5627 normal_loading_metadata,
5628 attention_mechanism,
5629 )?))
5630 }
5631 fn is_gptx(&self, _config: &str) -> bool {
5632 true
5633 }
5634 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
5635 let config: Qwen3VLConfig = serde_json::from_str(config)?;
5636 Ok(Box::new(config))
5637 }
5638 fn get_processor(
5639 &self,
5640 _model_config: &str,
5641 _processor_config: Option<ProcessorConfig>,
5642 _preprocessor_config: PreProcessorConfig,
5643 max_edge: Option<u32>,
5644 ) -> Arc<dyn Processor + Send + Sync> {
5645 Arc::new(Qwen3VLProcessor::new(max_edge))
5646 }
5647 fn supports_paged_attention(&self, _config: &str) -> bool {
5648 true
5649 }
5650 fn supports_prefix_cacher(&self, _config: &str) -> bool {
5651 true
5652 }
5653 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
5654 Arc::new(Qwen3VLPrefixer)
5655 }
5656 fn modalities(&self, _config: &str) -> Result<Modalities> {
5657 Ok(Modalities {
5658 input: vec![SupportedModality::Text, SupportedModality::Vision],
5659 output: vec![SupportedModality::Text],
5660 })
5661 }
5662}
5663
5664impl IsqModelLoader for Qwen3VLLoader {
5665 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
5666 Ok(vec![
5667 Regex::new(r"lm_head\.(weight|bias)$")?,
5668 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
5670 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
5671 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
5672 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
5673 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
5675 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
5676 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
5677 ])
5678 }
5679 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
5680 self.isq_layer_regexes(config)
5681 }
5682}
5683
5684impl DeviceMappedModelLoader for Qwen3VLLoader {
5685 fn mapped_max_act_size_elems(
5686 &self,
5687 config: &str,
5688 params: &AutoDeviceMapParams,
5689 ) -> Result<usize> {
5690 let AutoDeviceMapParams::Multimodal {
5691 max_seq_len,
5692 max_batch_size,
5693 max_image_shape,
5694 max_num_images,
5695 } = params
5696 else {
5697 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
5698 };
5699
5700 let cfg: Qwen3VLConfig = serde_json::from_str(config)?;
5701
5702 let img_seq_len = {
5704 let cfg = &cfg.vision_config;
5705 let grid_t = 1;
5707 let grid_h = (max_image_shape.0 / cfg.patch_size) / cfg.spatial_merge_size;
5709 let grid_w = (max_image_shape.1 / cfg.patch_size) / cfg.spatial_merge_size;
5710 grid_t * grid_h * grid_w * max_num_images
5711 };
5712
5713 let max_text_attn = {
5714 let cfg = &cfg.text_config;
5715 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
5717 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
5718 };
5719
5720 Ok(max_text_attn)
5721 }
5722
5723 fn non_mapped_max_act_size_elems(
5724 &self,
5725 config: &str,
5726 params: &AutoDeviceMapParams,
5727 ) -> Result<usize> {
5728 let AutoDeviceMapParams::Multimodal {
5729 max_seq_len: _,
5730 max_batch_size,
5731 max_image_shape,
5732 max_num_images,
5733 } = params
5734 else {
5735 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
5736 };
5737
5738 let cfg: Qwen3VLConfig = serde_json::from_str(config)?;
5739
5740 let img_seq_len = {
5742 let cfg = &cfg.vision_config;
5743 let grid_t = 1;
5745 let grid_h = max_image_shape.0 / cfg.patch_size;
5746 let grid_w = max_image_shape.1 / cfg.patch_size;
5747 grid_t * grid_h * grid_w
5748 };
5749
5750 let max_vision_attn = {
5751 let cfg = &cfg.vision_config;
5752 (max_batch_size * max_num_images) * cfg.num_heads * img_seq_len * img_seq_len
5753 };
5754
5755 Ok(max_vision_attn)
5756 }
5757
5758 fn non_mapped_size_in_bytes(
5759 &self,
5760 config: &str,
5761 dtype: DType,
5762 weight_pack_factor: usize,
5763 _matformer_config: Option<&MatformerSliceConfig>,
5764 ) -> Result<usize> {
5765 let cfg: Qwen3VLConfig = serde_json::from_str(config)?;
5766 let tie = cfg.tie_word_embeddings;
5767 let text_elems = {
5768 let cfg = &cfg.text_config;
5769 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
5770 let lm_head = if !tie || weight_pack_factor != 1 {
5772 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
5773 } else {
5774 0
5775 };
5776 let norm = cfg.hidden_size;
5777 embed_tokens + lm_head + norm
5778 };
5779
5780 let (patch_merger, deepstack_mergers) = {
5781 let cfg = &cfg.vision_config;
5782 let hidden_size = cfg.hidden_size * cfg.spatial_merge_size.pow(2);
5783
5784 let mlp0 = hidden_size * hidden_size + hidden_size;
5785 let mlp2 = hidden_size * cfg.out_hidden_size + cfg.out_hidden_size;
5786
5787 let ln_q = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
5789 let merger = mlp0 + mlp2 + ln_q;
5790
5791 let ds_ln = hidden_size + bias_if!(true, hidden_size);
5793 let ds_merger = mlp0 + mlp2 + ds_ln;
5794 let deepstack = cfg.deepstack_visual_indexes.len() * ds_merger;
5795
5796 (merger, deepstack)
5797 };
5798
5799 let patch_embed = {
5800 let cfg = &cfg.vision_config;
5801 let conv_cfg = Conv3dConfig {
5802 stride: cfg.patch_size,
5803 ..Default::default()
5804 };
5805 let kernel_sizes = [cfg.temporal_patch_size, cfg.patch_size, cfg.patch_size];
5806 let weight = cfg.in_chans * cfg.hidden_size / conv_cfg.groups
5807 * kernel_sizes[0]
5808 * kernel_sizes[1]
5809 * kernel_sizes[2];
5810 let bias = cfg.hidden_size;
5811 weight + bias
5812 };
5813
5814 let pos_embed = {
5815 let cfg = &cfg.vision_config;
5816 cfg.num_position_embeddings * cfg.hidden_size
5817 };
5818
5819 let encoder_layer = {
5820 let cfg = &cfg.vision_config;
5821 let norm1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
5822 let norm2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
5823
5824 #[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
5825 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
5826 let fc2 = cfg.hidden_size * cfg.intermediate_size + cfg.hidden_size;
5827
5828 let qkv = cfg.hidden_size * cfg.hidden_size * 3 + cfg.hidden_size * 3;
5829 let out = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
5830
5831 norm1 + norm2 + fc1 + fc2 + qkv + out
5832 };
5833
5834 let elems = text_elems
5835 + patch_merger
5836 + deepstack_mergers
5837 + patch_embed
5838 + pos_embed
5839 + encoder_layer * cfg.vision_config.depth;
5840
5841 Ok(elems * dtype.size_in_bytes())
5842 }
5843
5844 fn layer_sizes_in_bytes(
5845 &self,
5846 config: &str,
5847 dtype: DType,
5848 weight_pack_factor: usize,
5849 _matformer_config: Option<&MatformerSliceConfig>,
5850 ) -> Result<Vec<usize>> {
5851 let cfg: Qwen3VLConfig = serde_json::from_str(config)?;
5852 let per_layer_elems = {
5853 let cfg = &cfg.text_config;
5854 let input_layernorm = cfg.hidden_size;
5855 let post_attention_layernorm = cfg.hidden_size;
5856
5857 let size_in = cfg.hidden_size;
5858 let size_q = cfg.head_dim * cfg.num_attention_heads;
5859 let size_kv = cfg.head_dim * cfg.num_key_value_heads;
5860 let q_proj = size_in * size_q / weight_pack_factor;
5861 let k_proj = size_in * size_kv / weight_pack_factor;
5862 let v_proj = size_in * size_kv / weight_pack_factor;
5863 let o_proj = size_q * size_in / weight_pack_factor;
5864
5865 let q_norm = cfg.head_dim;
5866 let k_norm = cfg.head_dim;
5867
5868 let h_size = cfg.hidden_size;
5869 let i_size = cfg.intermediate_size;
5870 let gate_proj = h_size * i_size / weight_pack_factor;
5871 let up_proj = h_size * i_size / weight_pack_factor;
5872 let down_proj = i_size * h_size / weight_pack_factor;
5873
5874 input_layernorm
5875 + post_attention_layernorm
5876 + q_proj
5877 + k_proj
5878 + v_proj
5879 + o_proj
5880 + q_norm
5881 + k_norm
5882 + gate_proj
5883 + up_proj
5884 + down_proj
5885 };
5886 Ok(vec![
5887 per_layer_elems * dtype.size_in_bytes();
5888 cfg.text_config.num_hidden_layers
5889 ])
5890 }
5891
5892 fn num_layers(&self, config: &str) -> Result<usize> {
5893 let cfg: Qwen3VLConfig = serde_json::from_str(config)?;
5894 let cfg = &cfg.text_config;
5895 Ok(cfg.num_hidden_layers)
5896 }
5897
5898 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
5899 let cfg: Qwen3VLConfig = serde_json::from_str(config)?;
5900 let cfg = &cfg.text_config;
5901
5902 let cfg = ModelConfigMetadata {
5903 max_seq_len: cfg.max_position_embeddings,
5904 num_layers: cfg.num_hidden_layers,
5905 hidden_size: cfg.hidden_size,
5906 num_kv_heads: cfg.num_key_value_heads,
5907 num_attn_heads: cfg.num_attention_heads,
5908 sliding_window: cfg.sliding_window,
5909 k_head_dim: cfg.head_dim,
5910 v_head_dim: cfg.head_dim,
5911 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
5912 };
5913
5914 Ok(Box::new(cfg))
5915 }
5916
5917 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
5918 Some(vec![NonMappedSubModel::Vision])
5919 }
5920}
5921
5922pub struct Qwen3VLMoELoader;
5928
5929pub struct Qwen3VLMoEPrefixer;
5930
5931impl MultimodalPromptPrefixer for Qwen3VLMoEPrefixer {
5932 }
5935
5936impl MultimodalModelLoader for Qwen3VLMoELoader {
5937 fn load(
5938 &self,
5939 config: &str,
5940 vb: ShardedVarBuilder,
5941 normal_loading_metadata: NormalLoadingMetadata,
5942 attention_mechanism: AttentionImplementation,
5943 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
5944 let cfg: Qwen3VLMoEConfig = serde_json::from_str(config)?;
5945 Ok(Box::new(Qwen3VLMoEModel::new(
5946 &cfg,
5947 vb,
5948 self.is_gptx(config),
5949 normal_loading_metadata,
5950 attention_mechanism,
5951 )?))
5952 }
5953 fn is_gptx(&self, _config: &str) -> bool {
5954 true
5955 }
5956 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
5957 let config: Qwen3VLMoEConfig = serde_json::from_str(config)?;
5958 Ok(Box::new(config))
5959 }
5960 fn get_processor(
5961 &self,
5962 _model_config: &str,
5963 _processor_config: Option<ProcessorConfig>,
5964 _preprocessor_config: PreProcessorConfig,
5965 max_edge: Option<u32>,
5966 ) -> Arc<dyn Processor + Send + Sync> {
5967 Arc::new(Qwen3VLMoEProcessor::new(max_edge))
5968 }
5969 fn supports_paged_attention(&self, _config: &str) -> bool {
5970 true
5971 }
5972 fn supports_prefix_cacher(&self, _config: &str) -> bool {
5973 true
5974 }
5975 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
5976 Arc::new(Qwen3VLMoEPrefixer)
5977 }
5978 fn modalities(&self, _config: &str) -> Result<Modalities> {
5979 Ok(Modalities {
5980 input: vec![SupportedModality::Text, SupportedModality::Vision],
5981 output: vec![SupportedModality::Text],
5982 })
5983 }
5984}
5985
5986impl IsqModelLoader for Qwen3VLMoELoader {
5987 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
5988 Ok(vec![
5989 Regex::new(r"lm_head\.(weight|bias)$")?,
5990 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
5992 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
5993 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
5994 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
5995 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
5997 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
5998 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
5999 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate\.(weight|bias)$")?,
6001 Regex::new(
6003 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.gate_proj\.(weight|bias)$",
6004 )?,
6005 Regex::new(
6006 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.up_proj\.(weight|bias)$",
6007 )?,
6008 Regex::new(
6009 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.down_proj\.(weight|bias)$",
6010 )?,
6011 ])
6012 }
6013 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
6014 self.isq_layer_regexes(config)
6015 }
6016 fn isq_layer_regexes_moqe(&self, _config: &str) -> Result<Vec<Regex>> {
6017 Ok(vec![
6018 Regex::new(r"lm_head\.(weight|bias)$")?,
6019 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
6021 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
6022 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
6023 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate\.(weight|bias)$")?,
6025 Regex::new(
6027 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.gate_proj\.(weight|bias)$",
6028 )?,
6029 Regex::new(
6030 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.up_proj\.(weight|bias)$",
6031 )?,
6032 Regex::new(
6033 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.down_proj\.(weight|bias)$",
6034 )?,
6035 ])
6036 }
6037 fn immediate_isq_predicates_moqe(&self, config: &str) -> Result<Vec<Regex>> {
6038 self.isq_layer_regexes_moqe(config)
6039 }
6040}
6041
6042impl DeviceMappedModelLoader for Qwen3VLMoELoader {
6043 fn mapped_max_act_size_elems(
6044 &self,
6045 config: &str,
6046 params: &AutoDeviceMapParams,
6047 ) -> Result<usize> {
6048 let AutoDeviceMapParams::Multimodal {
6049 max_seq_len,
6050 max_batch_size,
6051 max_image_shape,
6052 max_num_images,
6053 } = params
6054 else {
6055 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
6056 };
6057
6058 let cfg: Qwen3VLMoEConfig = serde_json::from_str(config)?;
6059
6060 let img_seq_len = {
6062 let cfg = &cfg.vision_config;
6063 let grid_t = 1;
6065 let grid_h = (max_image_shape.0 / cfg.patch_size) / cfg.spatial_merge_size;
6067 let grid_w = (max_image_shape.1 / cfg.patch_size) / cfg.spatial_merge_size;
6068 grid_t * grid_h * grid_w * max_num_images
6069 };
6070
6071 let max_text_attn = {
6072 let cfg = &cfg.text_config;
6073 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
6075 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
6076 };
6077
6078 Ok(max_text_attn)
6079 }
6080
6081 fn non_mapped_max_act_size_elems(
6082 &self,
6083 config: &str,
6084 params: &AutoDeviceMapParams,
6085 ) -> Result<usize> {
6086 let AutoDeviceMapParams::Multimodal {
6087 max_seq_len: _,
6088 max_batch_size,
6089 max_image_shape,
6090 max_num_images,
6091 } = params
6092 else {
6093 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
6094 };
6095
6096 let cfg: Qwen3VLMoEConfig = serde_json::from_str(config)?;
6097
6098 let img_seq_len = {
6100 let cfg = &cfg.vision_config;
6101 let grid_t = 1;
6103 let grid_h = max_image_shape.0 / cfg.patch_size;
6104 let grid_w = max_image_shape.1 / cfg.patch_size;
6105 grid_t * grid_h * grid_w
6106 };
6107
6108 let max_vision_attn = {
6109 let cfg = &cfg.vision_config;
6110 (max_batch_size * max_num_images) * cfg.num_heads * img_seq_len * img_seq_len
6111 };
6112
6113 Ok(max_vision_attn)
6114 }
6115
6116 fn non_mapped_size_in_bytes(
6117 &self,
6118 config: &str,
6119 dtype: DType,
6120 weight_pack_factor: usize,
6121 _matformer_config: Option<&MatformerSliceConfig>,
6122 ) -> Result<usize> {
6123 let cfg: Qwen3VLMoEConfig = serde_json::from_str(config)?;
6124 let tie = cfg.tie_word_embeddings;
6125 let text_elems = {
6126 let cfg = &cfg.text_config;
6127 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
6128 let lm_head = if !tie || weight_pack_factor != 1 {
6130 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
6131 } else {
6132 0
6133 };
6134 let norm = cfg.hidden_size;
6135 embed_tokens + lm_head + norm
6136 };
6137
6138 let (patch_merger, deepstack_mergers) = {
6139 let cfg = &cfg.vision_config;
6140 let hidden_size = cfg.hidden_size * cfg.spatial_merge_size.pow(2);
6141
6142 let mlp0 = hidden_size * hidden_size + hidden_size;
6143 let mlp2 = hidden_size * cfg.out_hidden_size + cfg.out_hidden_size;
6144
6145 let ln_q = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6147 let merger = mlp0 + mlp2 + ln_q;
6148
6149 let ds_ln = hidden_size + bias_if!(true, hidden_size);
6151 let ds_merger = mlp0 + mlp2 + ds_ln;
6152 let deepstack = cfg.deepstack_visual_indexes.len() * ds_merger;
6153
6154 (merger, deepstack)
6155 };
6156
6157 let patch_embed = {
6158 let cfg = &cfg.vision_config;
6159 let conv_cfg = Conv3dConfig {
6160 stride: cfg.patch_size,
6161 ..Default::default()
6162 };
6163 let kernel_sizes = [cfg.temporal_patch_size, cfg.patch_size, cfg.patch_size];
6164 let weight = cfg.in_chans * cfg.hidden_size / conv_cfg.groups
6165 * kernel_sizes[0]
6166 * kernel_sizes[1]
6167 * kernel_sizes[2];
6168 let bias = cfg.hidden_size;
6169 weight + bias
6170 };
6171
6172 let pos_embed = {
6173 let cfg = &cfg.vision_config;
6174 cfg.num_position_embeddings * cfg.hidden_size
6175 };
6176
6177 let encoder_layer = {
6178 let cfg = &cfg.vision_config;
6179 let norm1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6180 let norm2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6181
6182 #[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
6183 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
6184 let fc2 = cfg.hidden_size * cfg.intermediate_size + cfg.hidden_size;
6185
6186 let qkv = cfg.hidden_size * cfg.hidden_size * 3 + cfg.hidden_size * 3;
6187 let out = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
6188
6189 norm1 + norm2 + fc1 + fc2 + qkv + out
6190 };
6191
6192 let elems = text_elems
6193 + patch_merger
6194 + deepstack_mergers
6195 + patch_embed
6196 + pos_embed
6197 + encoder_layer * cfg.vision_config.depth;
6198
6199 Ok(elems * dtype.size_in_bytes())
6200 }
6201
6202 fn layer_sizes_in_bytes(
6203 &self,
6204 config: &str,
6205 dtype: DType,
6206 weight_pack_factor: usize,
6207 _matformer_config: Option<&MatformerSliceConfig>,
6208 ) -> Result<Vec<usize>> {
6209 let cfg: Qwen3VLMoEConfig = serde_json::from_str(config)?;
6210 let text_cfg = &cfg.text_config;
6211
6212 let mut layer_sizes = Vec::with_capacity(text_cfg.num_hidden_layers);
6213
6214 for layer_idx in 0..text_cfg.num_hidden_layers {
6215 let input_layernorm = text_cfg.hidden_size;
6216 let post_attention_layernorm = text_cfg.hidden_size;
6217
6218 let size_in = text_cfg.hidden_size;
6219 let size_q = text_cfg.head_dim * text_cfg.num_attention_heads;
6220 let size_kv = text_cfg.head_dim * text_cfg.num_key_value_heads;
6221 let q_proj = size_in * size_q / weight_pack_factor;
6222 let k_proj = size_in * size_kv / weight_pack_factor;
6223 let v_proj = size_in * size_kv / weight_pack_factor;
6224 let o_proj = size_q * size_in / weight_pack_factor;
6225
6226 let q_norm = text_cfg.head_dim;
6227 let k_norm = text_cfg.head_dim;
6228
6229 let is_moe = !text_cfg.mlp_only_layers.contains(&layer_idx)
6231 && (text_cfg.num_experts > 0
6232 && (layer_idx + 1) % text_cfg.decoder_sparse_step == 0);
6233
6234 let mlp_elems = if is_moe {
6235 let gate = text_cfg.hidden_size * text_cfg.num_experts;
6237 let per_expert = {
6238 let h_size = text_cfg.hidden_size;
6239 let i_size = text_cfg.moe_intermediate_size;
6240 let gate_proj = h_size * i_size / weight_pack_factor;
6241 let up_proj = h_size * i_size / weight_pack_factor;
6242 let down_proj = i_size * h_size / weight_pack_factor;
6243 gate_proj + up_proj + down_proj
6244 };
6245 gate + per_expert * text_cfg.num_experts
6246 } else {
6247 let h_size = text_cfg.hidden_size;
6249 let i_size = text_cfg.intermediate_size;
6250 let gate_proj = h_size * i_size / weight_pack_factor;
6251 let up_proj = h_size * i_size / weight_pack_factor;
6252 let down_proj = i_size * h_size / weight_pack_factor;
6253 gate_proj + up_proj + down_proj
6254 };
6255
6256 let per_layer_elems = input_layernorm
6257 + post_attention_layernorm
6258 + q_proj
6259 + k_proj
6260 + v_proj
6261 + o_proj
6262 + q_norm
6263 + k_norm
6264 + mlp_elems;
6265
6266 layer_sizes.push(per_layer_elems * dtype.size_in_bytes());
6267 }
6268
6269 Ok(layer_sizes)
6270 }
6271
6272 fn num_layers(&self, config: &str) -> Result<usize> {
6273 let cfg: Qwen3VLMoEConfig = serde_json::from_str(config)?;
6274 let cfg = &cfg.text_config;
6275 Ok(cfg.num_hidden_layers)
6276 }
6277
6278 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
6279 let cfg: Qwen3VLMoEConfig = serde_json::from_str(config)?;
6280 let cfg = &cfg.text_config;
6281
6282 let cfg = ModelConfigMetadata {
6283 max_seq_len: cfg.max_position_embeddings,
6284 num_layers: cfg.num_hidden_layers,
6285 hidden_size: cfg.hidden_size,
6286 num_kv_heads: cfg.num_key_value_heads,
6287 num_attn_heads: cfg.num_attention_heads,
6288 sliding_window: cfg.sliding_window,
6289 k_head_dim: cfg.head_dim,
6290 v_head_dim: cfg.head_dim,
6291 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
6292 };
6293
6294 Ok(Box::new(cfg))
6295 }
6296
6297 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
6298 Some(vec![NonMappedSubModel::Vision])
6299 }
6300}
6301
6302pub struct Qwen3_5Loader;
6308
6309pub struct Qwen3_5Prefixer;
6310
6311impl MultimodalPromptPrefixer for Qwen3_5Prefixer {
6312 }
6315
6316impl MultimodalModelLoader for Qwen3_5Loader {
6317 fn load(
6318 &self,
6319 config: &str,
6320 vb: ShardedVarBuilder,
6321 normal_loading_metadata: NormalLoadingMetadata,
6322 attention_mechanism: AttentionImplementation,
6323 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
6324 let cfg: Qwen3_5Config = serde_json::from_str(config)?;
6325 Ok(Box::new(Qwen3_5Model::new(
6326 &cfg,
6327 vb,
6328 self.is_gptx(config),
6329 normal_loading_metadata,
6330 attention_mechanism,
6331 )?))
6332 }
6333 fn is_gptx(&self, _config: &str) -> bool {
6334 true
6335 }
6336 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
6337 let config: Qwen3_5Config = serde_json::from_str(config)?;
6338 Ok(Box::new(config))
6339 }
6340 fn get_processor(
6341 &self,
6342 _model_config: &str,
6343 _processor_config: Option<ProcessorConfig>,
6344 _preprocessor_config: PreProcessorConfig,
6345 max_edge: Option<u32>,
6346 ) -> Arc<dyn Processor + Send + Sync> {
6347 Arc::new(Qwen3_5Processor::new(max_edge))
6348 }
6349 fn supports_paged_attention(&self, _config: &str) -> bool {
6350 true
6351 }
6352 fn supports_prefix_cacher(&self, _config: &str) -> bool {
6353 true
6354 }
6355 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
6356 Arc::new(Qwen3_5Prefixer)
6357 }
6358 fn modalities(&self, _config: &str) -> Result<Modalities> {
6359 Ok(Modalities {
6360 input: vec![SupportedModality::Text, SupportedModality::Vision],
6361 output: vec![SupportedModality::Text],
6362 })
6363 }
6364}
6365
6366impl IsqModelLoader for Qwen3_5Loader {
6367 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
6368 Ok(vec![
6369 Regex::new(r"lm_head\.(weight|bias)$")?,
6370 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
6372 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
6373 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
6374 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
6375 Regex::new(
6377 r"model\.language_model\.layers\.(\d+)\.linear_attn\.out_proj\.(weight|bias)$",
6378 )?,
6379 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
6381 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
6382 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
6383 ])
6384 }
6385 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
6386 self.isq_layer_regexes(config)
6387 }
6388}
6389
6390impl DeviceMappedModelLoader for Qwen3_5Loader {
6391 fn mapped_max_act_size_elems(
6392 &self,
6393 config: &str,
6394 params: &AutoDeviceMapParams,
6395 ) -> Result<usize> {
6396 let AutoDeviceMapParams::Multimodal {
6397 max_seq_len,
6398 max_batch_size,
6399 max_image_shape,
6400 max_num_images,
6401 } = params
6402 else {
6403 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
6404 };
6405
6406 let cfg: Qwen3_5Config = serde_json::from_str(config)?;
6407
6408 let img_seq_len = {
6409 let cfg = &cfg.vision_config;
6410 let grid_t = 1;
6411 let grid_h = (max_image_shape.0 / cfg.patch_size) / cfg.spatial_merge_size;
6412 let grid_w = (max_image_shape.1 / cfg.patch_size) / cfg.spatial_merge_size;
6413 grid_t * grid_h * grid_w * max_num_images
6414 };
6415
6416 let max_text_attn = {
6417 let cfg = &cfg.text_config;
6418 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
6419 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
6420 };
6421
6422 Ok(max_text_attn)
6423 }
6424
6425 fn non_mapped_max_act_size_elems(
6426 &self,
6427 config: &str,
6428 params: &AutoDeviceMapParams,
6429 ) -> Result<usize> {
6430 let AutoDeviceMapParams::Multimodal {
6431 max_seq_len: _,
6432 max_batch_size,
6433 max_image_shape,
6434 max_num_images,
6435 } = params
6436 else {
6437 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
6438 };
6439
6440 let cfg: Qwen3_5Config = serde_json::from_str(config)?;
6441
6442 let img_seq_len = {
6443 let cfg = &cfg.vision_config;
6444 let grid_t = 1;
6445 let grid_h = max_image_shape.0 / cfg.patch_size;
6446 let grid_w = max_image_shape.1 / cfg.patch_size;
6447 grid_t * grid_h * grid_w
6448 };
6449
6450 let max_vision_attn = {
6451 let cfg = &cfg.vision_config;
6452 (max_batch_size * max_num_images) * cfg.num_heads * img_seq_len * img_seq_len
6453 };
6454
6455 Ok(max_vision_attn)
6456 }
6457
6458 fn non_mapped_size_in_bytes(
6459 &self,
6460 config: &str,
6461 dtype: DType,
6462 weight_pack_factor: usize,
6463 _matformer_config: Option<&MatformerSliceConfig>,
6464 ) -> Result<usize> {
6465 let cfg: Qwen3_5Config = serde_json::from_str(config)?;
6466 let tie = cfg.tie_word_embeddings;
6467 let text_elems = {
6468 let cfg = &cfg.text_config;
6469 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
6470 let lm_head = if !tie || weight_pack_factor != 1 {
6471 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
6472 } else {
6473 0
6474 };
6475 let norm = cfg.hidden_size;
6476 embed_tokens + lm_head + norm
6477 };
6478
6479 let (patch_merger, deepstack_mergers) = {
6480 let cfg = &cfg.vision_config;
6481 let hidden_size = cfg.hidden_size * cfg.spatial_merge_size.pow(2);
6482
6483 let mlp0 = hidden_size * hidden_size + hidden_size;
6484 let mlp2 = hidden_size * cfg.out_hidden_size + cfg.out_hidden_size;
6485
6486 let ln_q = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6487 let merger = mlp0 + mlp2 + ln_q;
6488
6489 let ds_ln = hidden_size + bias_if!(true, hidden_size);
6490 let ds_merger = mlp0 + mlp2 + ds_ln;
6491 let deepstack = cfg.deepstack_visual_indexes.len() * ds_merger;
6492
6493 (merger, deepstack)
6494 };
6495
6496 let patch_embed = {
6497 let cfg = &cfg.vision_config;
6498 let conv_cfg = Conv3dConfig {
6499 stride: cfg.patch_size,
6500 ..Default::default()
6501 };
6502 let kernel_sizes = [cfg.temporal_patch_size, cfg.patch_size, cfg.patch_size];
6503 let weight = cfg.in_chans * cfg.hidden_size / conv_cfg.groups
6504 * kernel_sizes[0]
6505 * kernel_sizes[1]
6506 * kernel_sizes[2];
6507 let bias = cfg.hidden_size;
6508 weight + bias
6509 };
6510
6511 let pos_embed = {
6512 let cfg = &cfg.vision_config;
6513 cfg.num_position_embeddings * cfg.hidden_size
6514 };
6515
6516 let encoder_layer = {
6517 let cfg = &cfg.vision_config;
6518 let norm1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6519 let norm2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6520
6521 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
6522 let fc2 = cfg.hidden_size * cfg.intermediate_size + cfg.hidden_size;
6523
6524 let qkv = cfg.hidden_size * cfg.hidden_size * 3 + cfg.hidden_size * 3;
6525 let out = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
6526
6527 norm1 + norm2 + fc1 + fc2 + qkv + out
6528 };
6529
6530 let elems = text_elems
6531 + patch_merger
6532 + deepstack_mergers
6533 + patch_embed
6534 + pos_embed
6535 + encoder_layer * cfg.vision_config.depth;
6536
6537 Ok(elems * dtype.size_in_bytes())
6538 }
6539
6540 fn layer_sizes_in_bytes(
6541 &self,
6542 config: &str,
6543 dtype: DType,
6544 weight_pack_factor: usize,
6545 _matformer_config: Option<&MatformerSliceConfig>,
6546 ) -> Result<Vec<usize>> {
6547 let cfg: Qwen3_5Config = serde_json::from_str(config)?;
6548 let text_cfg = &cfg.text_config;
6549 let layer_types = text_cfg.layer_types();
6550
6551 let mut layer_sizes = Vec::with_capacity(text_cfg.num_hidden_layers);
6552
6553 for layer_type in &layer_types {
6554 let input_layernorm = text_cfg.hidden_size;
6555 let post_attention_layernorm = text_cfg.hidden_size;
6556
6557 let attn_elems = match layer_type {
6558 crate::vision_models::qwen3_5::config::LayerType::FullAttention => {
6559 let size_in = text_cfg.hidden_size;
6560 let size_q = text_cfg.head_dim * text_cfg.num_attention_heads;
6561 let size_kv = text_cfg.head_dim * text_cfg.num_key_value_heads;
6562 let q_proj = size_in * size_q * 2 / weight_pack_factor;
6563 let k_proj = size_in * size_kv / weight_pack_factor;
6564 let v_proj = size_in * size_kv / weight_pack_factor;
6565 let o_proj = size_q * size_in / weight_pack_factor;
6566 let q_norm = text_cfg.head_dim;
6567 let k_norm = text_cfg.head_dim;
6568 q_proj + k_proj + v_proj + o_proj + q_norm + k_norm
6569 }
6570 crate::vision_models::qwen3_5::config::LayerType::LinearAttention => {
6571 let hidden = text_cfg.hidden_size;
6572 let key_dim = text_cfg.linear_key_dim();
6573 let value_dim = text_cfg.linear_value_dim();
6574 let conv_dim = text_cfg.linear_conv_dim();
6575 let in_proj_qkvz = hidden * (key_dim * 2 + value_dim * 2);
6577 let in_proj_ba = hidden * (text_cfg.linear_num_value_heads * 2);
6579 let out_proj = value_dim * hidden / weight_pack_factor;
6580 let conv1d = conv_dim * text_cfg.linear_conv_kernel_dim;
6581 let dt_bias = text_cfg.linear_num_value_heads;
6582 let a_log = text_cfg.linear_num_value_heads;
6583 let norm = text_cfg.linear_value_head_dim;
6585 in_proj_qkvz + in_proj_ba + out_proj + conv1d + dt_bias + a_log + norm
6586 }
6587 };
6588
6589 let mlp_elems = {
6591 let h_size = text_cfg.hidden_size;
6592 let i_size = text_cfg.intermediate_size;
6593 let gate_proj = h_size * i_size / weight_pack_factor;
6594 let up_proj = h_size * i_size / weight_pack_factor;
6595 let down_proj = i_size * h_size / weight_pack_factor;
6596 gate_proj + up_proj + down_proj
6597 };
6598
6599 let per_layer_elems =
6600 input_layernorm + post_attention_layernorm + attn_elems + mlp_elems;
6601
6602 layer_sizes.push(per_layer_elems * dtype.size_in_bytes());
6603 }
6604
6605 Ok(layer_sizes)
6606 }
6607
6608 fn num_layers(&self, config: &str) -> Result<usize> {
6609 let cfg: Qwen3_5Config = serde_json::from_str(config)?;
6610 Ok(cfg.text_config.num_hidden_layers)
6611 }
6612
6613 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
6614 let cfg: Qwen3_5Config = serde_json::from_str(config)?;
6615 let cfg = &cfg.text_config;
6616
6617 let cfg = ModelConfigMetadata {
6618 max_seq_len: cfg.max_position_embeddings,
6619 num_layers: cfg.num_hidden_layers,
6620 hidden_size: cfg.hidden_size,
6621 num_kv_heads: cfg.num_key_value_heads,
6622 num_attn_heads: cfg.num_attention_heads,
6623 sliding_window: None,
6624 k_head_dim: cfg.head_dim,
6625 v_head_dim: cfg.head_dim,
6626 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
6627 };
6628
6629 Ok(Box::new(cfg))
6630 }
6631
6632 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
6633 Some(vec![NonMappedSubModel::Vision])
6634 }
6635}
6636
6637pub struct Qwen3_5MoeLoader;
6643
6644pub struct Qwen3_5MoePrefixer;
6645
6646impl MultimodalPromptPrefixer for Qwen3_5MoePrefixer {
6647 }
6650
6651impl MultimodalModelLoader for Qwen3_5MoeLoader {
6652 fn load(
6653 &self,
6654 config: &str,
6655 vb: ShardedVarBuilder,
6656 normal_loading_metadata: NormalLoadingMetadata,
6657 attention_mechanism: AttentionImplementation,
6658 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
6659 let cfg: Qwen3_5MoeConfig = serde_json::from_str(config)?;
6660 Ok(Box::new(Qwen3_5MoeModel::new(
6661 &cfg,
6662 vb,
6663 self.is_gptx(config),
6664 normal_loading_metadata,
6665 attention_mechanism,
6666 )?))
6667 }
6668 fn is_gptx(&self, _config: &str) -> bool {
6669 true
6670 }
6671 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
6672 let config: Qwen3_5MoeConfig = serde_json::from_str(config)?;
6673 Ok(Box::new(config))
6674 }
6675 fn get_processor(
6676 &self,
6677 _model_config: &str,
6678 _processor_config: Option<ProcessorConfig>,
6679 _preprocessor_config: PreProcessorConfig,
6680 max_edge: Option<u32>,
6681 ) -> Arc<dyn Processor + Send + Sync> {
6682 Arc::new(Qwen3_5MoeProcessor::new(max_edge))
6683 }
6684 fn supports_paged_attention(&self, _config: &str) -> bool {
6685 true
6686 }
6687 fn supports_prefix_cacher(&self, _config: &str) -> bool {
6688 true
6689 }
6690 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
6691 Arc::new(Qwen3_5MoePrefixer)
6692 }
6693 fn modalities(&self, _config: &str) -> Result<Modalities> {
6694 Ok(Modalities {
6695 input: vec![SupportedModality::Text, SupportedModality::Vision],
6696 output: vec![SupportedModality::Text],
6697 })
6698 }
6699}
6700
6701impl IsqModelLoader for Qwen3_5MoeLoader {
6702 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
6703 Ok(vec![
6704 Regex::new(r"lm_head\.(weight|bias)$")?,
6705 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
6707 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
6708 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
6709 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
6710 Regex::new(
6712 r"model\.language_model\.layers\.(\d+)\.linear_attn\.out_proj\.(weight|bias)$",
6713 )?,
6714 Regex::new(
6716 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.gate_proj\.(weight|bias)$",
6717 )?,
6718 Regex::new(
6719 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.up_proj\.(weight|bias)$",
6720 )?,
6721 Regex::new(
6722 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.down_proj\.(weight|bias)$",
6723 )?,
6724 Regex::new(
6726 r"model\.language_model\.layers\.(\d+)\.mlp\.shared_expert\.gate_proj\.(weight|bias)$",
6727 )?,
6728 Regex::new(
6729 r"model\.language_model\.layers\.(\d+)\.mlp\.shared_expert\.up_proj\.(weight|bias)$",
6730 )?,
6731 Regex::new(
6732 r"model\.language_model\.layers\.(\d+)\.mlp\.shared_expert\.down_proj\.(weight|bias)$",
6733 )?,
6734 ])
6735 }
6736 fn immediate_isq_predicates(&self, config: &str) -> Result<Vec<Regex>> {
6737 self.isq_layer_regexes(config)
6738 }
6739 fn isq_layer_regexes_moqe(&self, _config: &str) -> Result<Vec<Regex>> {
6740 Ok(vec![
6741 Regex::new(r"lm_head\.(weight|bias)$")?,
6742 Regex::new(
6744 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.gate_proj\.(weight|bias)$",
6745 )?,
6746 Regex::new(
6747 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.up_proj\.(weight|bias)$",
6748 )?,
6749 Regex::new(
6750 r"model\.language_model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.down_proj\.(weight|bias)$",
6751 )?,
6752 Regex::new(
6754 r"model\.language_model\.layers\.(\d+)\.mlp\.shared_expert\.gate_proj\.(weight|bias)$",
6755 )?,
6756 Regex::new(
6757 r"model\.language_model\.layers\.(\d+)\.mlp\.shared_expert\.up_proj\.(weight|bias)$",
6758 )?,
6759 Regex::new(
6760 r"model\.language_model\.layers\.(\d+)\.mlp\.shared_expert\.down_proj\.(weight|bias)$",
6761 )?,
6762 ])
6763 }
6764 fn immediate_isq_predicates_moqe(&self, config: &str) -> Result<Vec<Regex>> {
6765 self.isq_layer_regexes_moqe(config)
6766 }
6767}
6768
6769impl DeviceMappedModelLoader for Qwen3_5MoeLoader {
6770 fn mapped_max_act_size_elems(
6771 &self,
6772 config: &str,
6773 params: &AutoDeviceMapParams,
6774 ) -> Result<usize> {
6775 let AutoDeviceMapParams::Multimodal {
6776 max_seq_len,
6777 max_batch_size,
6778 max_image_shape,
6779 max_num_images,
6780 } = params
6781 else {
6782 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
6783 };
6784
6785 let cfg: Qwen3_5MoeConfig = serde_json::from_str(config)?;
6786
6787 let img_seq_len = {
6788 let cfg = &cfg.vision_config;
6789 let grid_t = 1;
6790 let grid_h = (max_image_shape.0 / cfg.patch_size) / cfg.spatial_merge_size;
6791 let grid_w = (max_image_shape.1 / cfg.patch_size) / cfg.spatial_merge_size;
6792 grid_t * grid_h * grid_w * max_num_images
6793 };
6794
6795 let max_text_attn = {
6796 let cfg = &cfg.text_config;
6797 let max_seq_len = img_seq_len + max_seq_len.min(&ATTENTION_CHUNK_SIZE);
6798 max_batch_size * cfg.num_attention_heads * max_seq_len * max_seq_len
6799 };
6800
6801 Ok(max_text_attn)
6802 }
6803
6804 fn non_mapped_max_act_size_elems(
6805 &self,
6806 config: &str,
6807 params: &AutoDeviceMapParams,
6808 ) -> Result<usize> {
6809 let AutoDeviceMapParams::Multimodal {
6810 max_seq_len: _,
6811 max_batch_size,
6812 max_image_shape,
6813 max_num_images,
6814 } = params
6815 else {
6816 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
6817 };
6818
6819 let cfg: Qwen3_5MoeConfig = serde_json::from_str(config)?;
6820
6821 let img_seq_len = {
6822 let cfg = &cfg.vision_config;
6823 let grid_t = 1;
6824 let grid_h = max_image_shape.0 / cfg.patch_size;
6825 let grid_w = max_image_shape.1 / cfg.patch_size;
6826 grid_t * grid_h * grid_w
6827 };
6828
6829 let max_vision_attn = {
6830 let cfg = &cfg.vision_config;
6831 (max_batch_size * max_num_images) * cfg.num_heads * img_seq_len * img_seq_len
6832 };
6833
6834 Ok(max_vision_attn)
6835 }
6836
6837 fn non_mapped_size_in_bytes(
6838 &self,
6839 config: &str,
6840 dtype: DType,
6841 weight_pack_factor: usize,
6842 _matformer_config: Option<&MatformerSliceConfig>,
6843 ) -> Result<usize> {
6844 let cfg: Qwen3_5MoeConfig = serde_json::from_str(config)?;
6845 let tie = cfg.tie_word_embeddings;
6846 let text_elems = {
6847 let cfg = &cfg.text_config;
6848 let embed_tokens = cfg.hidden_size * cfg.vocab_size / weight_pack_factor;
6849 let lm_head = if !tie || weight_pack_factor != 1 {
6850 cfg.hidden_size * cfg.vocab_size / weight_pack_factor
6851 } else {
6852 0
6853 };
6854 let norm = cfg.hidden_size;
6855 embed_tokens + lm_head + norm
6856 };
6857
6858 let (patch_merger, deepstack_mergers) = {
6859 let cfg = &cfg.vision_config;
6860 let hidden_size = cfg.hidden_size * cfg.spatial_merge_size.pow(2);
6861
6862 let mlp0 = hidden_size * hidden_size + hidden_size;
6863 let mlp2 = hidden_size * cfg.out_hidden_size + cfg.out_hidden_size;
6864
6865 let ln_q = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6866 let merger = mlp0 + mlp2 + ln_q;
6867
6868 let ds_ln = hidden_size + bias_if!(true, hidden_size);
6869 let ds_merger = mlp0 + mlp2 + ds_ln;
6870 let deepstack = cfg.deepstack_visual_indexes.len() * ds_merger;
6871
6872 (merger, deepstack)
6873 };
6874
6875 let patch_embed = {
6876 let cfg = &cfg.vision_config;
6877 let conv_cfg = Conv3dConfig {
6878 stride: cfg.patch_size,
6879 ..Default::default()
6880 };
6881 let kernel_sizes = [cfg.temporal_patch_size, cfg.patch_size, cfg.patch_size];
6882 let weight = cfg.in_chans * cfg.hidden_size / conv_cfg.groups
6883 * kernel_sizes[0]
6884 * kernel_sizes[1]
6885 * kernel_sizes[2];
6886 let bias = cfg.hidden_size;
6887 weight + bias
6888 };
6889
6890 let pos_embed = {
6891 let cfg = &cfg.vision_config;
6892 cfg.num_position_embeddings * cfg.hidden_size
6893 };
6894
6895 let encoder_layer = {
6896 let cfg = &cfg.vision_config;
6897 let norm1 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6898 let norm2 = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6899
6900 let fc1 = cfg.hidden_size * cfg.intermediate_size + cfg.intermediate_size;
6901 let fc2 = cfg.hidden_size * cfg.intermediate_size + cfg.hidden_size;
6902
6903 let qkv = cfg.hidden_size * cfg.hidden_size * 3 + cfg.hidden_size * 3;
6904 let out = cfg.hidden_size * cfg.hidden_size + cfg.hidden_size;
6905
6906 norm1 + norm2 + fc1 + fc2 + qkv + out
6907 };
6908
6909 let elems = text_elems
6910 + patch_merger
6911 + deepstack_mergers
6912 + patch_embed
6913 + pos_embed
6914 + encoder_layer * cfg.vision_config.depth;
6915
6916 Ok(elems * dtype.size_in_bytes())
6917 }
6918
6919 fn layer_sizes_in_bytes(
6920 &self,
6921 config: &str,
6922 dtype: DType,
6923 weight_pack_factor: usize,
6924 _matformer_config: Option<&MatformerSliceConfig>,
6925 ) -> Result<Vec<usize>> {
6926 let cfg: Qwen3_5MoeConfig = serde_json::from_str(config)?;
6927 let text_cfg = &cfg.text_config;
6928 let layer_types = text_cfg.layer_types();
6929
6930 let mut layer_sizes = Vec::with_capacity(text_cfg.num_hidden_layers);
6931
6932 for layer_type in &layer_types {
6933 let input_layernorm = text_cfg.hidden_size;
6934 let post_attention_layernorm = text_cfg.hidden_size;
6935
6936 let attn_elems = match layer_type {
6937 crate::vision_models::qwen3_5_moe::config::LayerType::FullAttention => {
6938 let size_in = text_cfg.hidden_size;
6939 let size_q = text_cfg.head_dim * text_cfg.num_attention_heads;
6940 let size_kv = text_cfg.head_dim * text_cfg.num_key_value_heads;
6941 let q_proj = size_in * size_q * 2 / weight_pack_factor;
6942 let k_proj = size_in * size_kv / weight_pack_factor;
6943 let v_proj = size_in * size_kv / weight_pack_factor;
6944 let o_proj = size_q * size_in / weight_pack_factor;
6945 let q_norm = text_cfg.head_dim;
6946 let k_norm = text_cfg.head_dim;
6947 q_proj + k_proj + v_proj + o_proj + q_norm + k_norm
6948 }
6949 crate::vision_models::qwen3_5_moe::config::LayerType::LinearAttention => {
6950 let hidden = text_cfg.hidden_size;
6951 let key_dim = text_cfg.linear_key_dim();
6952 let value_dim = text_cfg.linear_value_dim();
6953 let conv_dim = text_cfg.linear_conv_dim();
6954 let in_proj_qkvz = hidden * (key_dim * 2 + value_dim * 2);
6956 let in_proj_ba = hidden * (text_cfg.linear_num_value_heads * 2);
6958 let out_proj = value_dim * hidden / weight_pack_factor;
6960 let conv1d = conv_dim * text_cfg.linear_conv_kernel_dim;
6962 let dt_bias = text_cfg.linear_num_value_heads;
6964 let a_log = text_cfg.linear_num_value_heads;
6965 let norm = text_cfg.linear_value_head_dim;
6967 in_proj_qkvz + in_proj_ba + out_proj + conv1d + dt_bias + a_log + norm
6968 }
6969 };
6970
6971 let moe_elems = {
6973 let gate = text_cfg.hidden_size * text_cfg.num_experts;
6974 let per_expert = {
6975 let h_size = text_cfg.hidden_size;
6976 let i_size = text_cfg.moe_intermediate_size;
6977 let gate_proj = h_size * i_size / weight_pack_factor;
6978 let up_proj = h_size * i_size / weight_pack_factor;
6979 let down_proj = i_size * h_size / weight_pack_factor;
6980 gate_proj + up_proj + down_proj
6981 };
6982 let shared_expert = {
6983 let h_size = text_cfg.hidden_size;
6984 let i_size = text_cfg.shared_expert_intermediate_size;
6985 let gate_proj = h_size * i_size / weight_pack_factor;
6986 let up_proj = h_size * i_size / weight_pack_factor;
6987 let down_proj = i_size * h_size / weight_pack_factor;
6988 gate_proj + up_proj + down_proj
6989 };
6990 let shared_expert_gate = text_cfg.hidden_size;
6991 gate + per_expert * text_cfg.num_experts + shared_expert + shared_expert_gate
6992 };
6993
6994 let per_layer_elems =
6995 input_layernorm + post_attention_layernorm + attn_elems + moe_elems;
6996
6997 layer_sizes.push(per_layer_elems * dtype.size_in_bytes());
6998 }
6999
7000 Ok(layer_sizes)
7001 }
7002
7003 fn num_layers(&self, config: &str) -> Result<usize> {
7004 let cfg: Qwen3_5MoeConfig = serde_json::from_str(config)?;
7005 Ok(cfg.text_config.num_hidden_layers)
7006 }
7007
7008 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
7009 let cfg: Qwen3_5MoeConfig = serde_json::from_str(config)?;
7010 let cfg = &cfg.text_config;
7011
7012 let cfg = ModelConfigMetadata {
7013 max_seq_len: cfg.max_position_embeddings,
7014 num_layers: cfg.num_hidden_layers,
7015 hidden_size: cfg.hidden_size,
7016 num_kv_heads: cfg.num_key_value_heads,
7017 num_attn_heads: cfg.num_attention_heads,
7018 sliding_window: None,
7019 k_head_dim: cfg.head_dim,
7020 v_head_dim: cfg.head_dim,
7021 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
7022 };
7023
7024 Ok(Box::new(cfg))
7025 }
7026
7027 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
7028 Some(vec![NonMappedSubModel::Vision])
7029 }
7030}
7031
7032pub struct VoxtralLoader;
7038
7039pub struct VoxtralPrefixer;
7040
7041impl MultimodalPromptPrefixer for VoxtralPrefixer {
7042 fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
7043 prompt.to_string()
7044 }
7045}
7046
7047impl MultimodalModelLoader for VoxtralLoader {
7048 fn load(
7049 &self,
7050 config: &str,
7051 vb: ShardedVarBuilder,
7052 normal_loading_metadata: NormalLoadingMetadata,
7053 attention_mechanism: AttentionImplementation,
7054 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
7055 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7056 Ok(Box::new(VoxtralModel::new(
7057 &cfg,
7058 vb,
7059 self.is_gptx(config),
7060 normal_loading_metadata,
7061 attention_mechanism,
7062 )?))
7063 }
7064 fn is_gptx(&self, _config: &str) -> bool {
7065 true
7066 }
7067 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
7068 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7069 Ok(Box::new(cfg))
7070 }
7071 fn get_processor(
7072 &self,
7073 model_config: &str,
7074 _processor_config: Option<ProcessorConfig>,
7075 _preprocessor_config: PreProcessorConfig,
7076 _max_edge: Option<u32>,
7077 ) -> Arc<dyn Processor + Send + Sync> {
7078 let cfg: VoxtralConfig =
7079 serde_json::from_str(model_config).expect("Failed to parse VoxtralConfig");
7080 Arc::new(VoxtralProcessor::new(&cfg))
7081 }
7082 fn supports_paged_attention(&self, _config: &str) -> bool {
7083 false
7084 }
7085 fn supports_prefix_cacher(&self, _config: &str) -> bool {
7086 false
7087 }
7088 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
7089 Arc::new(VoxtralPrefixer)
7090 }
7091 fn modalities(&self, _config: &str) -> Result<Modalities> {
7092 Ok(Modalities {
7093 input: vec![SupportedModality::Text, SupportedModality::Audio],
7094 output: vec![SupportedModality::Text],
7095 })
7096 }
7097 fn default_chat_template(&self, _config: &str) -> Option<String> {
7098 Some("{{ bos_token }}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% if message['role'] == 'user' %}{{ '[INST] ' + message['content'] + ' [/INST]' }}{% elif message['role'] == 'assistant' %}{{ message['content'] + eos_token + ' ' }}{% else %}{{ raise_exception('Only user and assistant roles are supported!') }}{% endif %}{% endfor %}".to_string())
7100 }
7101 fn default_bos_eos(&self, _config: &str) -> Option<(String, String)> {
7102 Some(("<s>".to_string(), "</s>".to_string()))
7104 }
7105}
7106
7107impl IsqModelLoader for VoxtralLoader {
7108 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
7109 Ok(vec![
7110 Regex::new(r"lm_head\.(weight|bias)$")?,
7112 Regex::new(r"layers\.(\d+)\.attention\.wq\.(weight|bias)$")?,
7114 Regex::new(r"layers\.(\d+)\.attention\.wk\.(weight|bias)$")?,
7115 Regex::new(r"layers\.(\d+)\.attention\.wv\.(weight|bias)$")?,
7116 Regex::new(r"layers\.(\d+)\.attention\.wo\.(weight|bias)$")?,
7117 Regex::new(r"layers\.(\d+)\.feed_forward\.w1\.(weight|bias)$")?,
7119 Regex::new(r"layers\.(\d+)\.feed_forward\.w3\.(weight|bias)$")?,
7120 Regex::new(r"layers\.(\d+)\.feed_forward\.w2\.(weight|bias)$")?,
7121 ])
7122 }
7123 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
7124 Ok(vec![
7125 Regex::new(r"tok_embeddings\.(weight|bias)$")?,
7126 Regex::new(r"layers\.(\d+)\.attention\.wq\.(weight|bias)$")?,
7128 Regex::new(r"layers\.(\d+)\.attention\.wk\.(weight|bias)$")?,
7129 Regex::new(r"layers\.(\d+)\.attention\.wv\.(weight|bias)$")?,
7130 Regex::new(r"layers\.(\d+)\.attention\.wo\.(weight|bias)$")?,
7131 Regex::new(r"layers\.(\d+)\.feed_forward\.w1\.(weight|bias)$")?,
7133 Regex::new(r"layers\.(\d+)\.feed_forward\.w3\.(weight|bias)$")?,
7134 Regex::new(r"layers\.(\d+)\.feed_forward\.w2\.(weight|bias)$")?,
7135 ])
7136 }
7137}
7138
7139#[allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
7140impl DeviceMappedModelLoader for VoxtralLoader {
7141 fn mapped_max_act_size_elems(
7142 &self,
7143 config: &str,
7144 params: &AutoDeviceMapParams,
7145 ) -> Result<usize> {
7146 let AutoDeviceMapParams::Multimodal {
7147 max_seq_len,
7148 max_batch_size,
7149 ..
7150 } = params
7151 else {
7152 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
7153 };
7154
7155 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7156
7157 let max_audio_tokens = 375;
7160 let total_seq = max_audio_tokens + *max_seq_len.min(&ATTENTION_CHUNK_SIZE);
7161 Ok(max_batch_size * cfg.n_heads * total_seq * total_seq)
7162 }
7163
7164 fn non_mapped_max_act_size_elems(
7165 &self,
7166 config: &str,
7167 params: &AutoDeviceMapParams,
7168 ) -> Result<usize> {
7169 let AutoDeviceMapParams::Multimodal { max_batch_size, .. } = params else {
7170 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
7171 };
7172
7173 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7174 let enc = &cfg.multimodal.whisper_model_args.encoder_args;
7175 let max_enc_seq = 3000usize;
7178 Ok(max_batch_size * enc.n_heads * max_enc_seq * max_enc_seq)
7179 }
7180
7181 fn non_mapped_size_in_bytes(
7182 &self,
7183 config: &str,
7184 dtype: DType,
7185 _weight_pack_factor: usize,
7186 _matformer_config: Option<&MatformerSliceConfig>,
7187 ) -> Result<usize> {
7188 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7189 let enc = &cfg.multimodal.whisper_model_args.encoder_args;
7190 let ds = &cfg.multimodal.whisper_model_args.downsample_args;
7191
7192 let elem = dtype.size_in_bytes();
7193
7194 let conv1 = enc.dim * enc.audio_encoding_args.num_mel_bins * 3 + enc.dim; let conv2 = enc.dim * enc.dim * 3 + enc.dim;
7197
7198 let enc_attn_per_layer = 4 * enc.dim * enc.dim; let enc_mlp_per_layer = 3 * enc.dim * enc.hidden_dim; let enc_norm_per_layer = 2 * enc.dim; let enc_layers =
7203 enc.n_layers * (enc_attn_per_layer + enc_mlp_per_layer + enc_norm_per_layer);
7204 let enc_final_norm = enc.dim;
7205
7206 let adapter_in_features = enc.dim * ds.downsample_factor;
7208 let adapter = adapter_in_features * cfg.dim + cfg.dim + cfg.dim * cfg.dim + cfg.dim;
7209
7210 let total_encoder = conv1 + conv2 + enc_layers + enc_final_norm + adapter;
7211
7212 let embeddings = cfg.vocab_size * cfg.dim;
7214
7215 Ok((total_encoder + embeddings) * elem)
7216 }
7217
7218 fn layer_sizes_in_bytes(
7219 &self,
7220 config: &str,
7221 dtype: DType,
7222 weight_pack_factor: usize,
7223 _matformer_config: Option<&MatformerSliceConfig>,
7224 ) -> Result<Vec<usize>> {
7225 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7226 let elem = dtype.size_in_bytes();
7227
7228 let attn = (cfg.dim * cfg.n_heads * cfg.head_dim
7229 + cfg.dim * cfg.n_kv_heads * cfg.head_dim
7230 + cfg.dim * cfg.n_kv_heads * cfg.head_dim
7231 + cfg.n_heads * cfg.head_dim * cfg.dim)
7232 / weight_pack_factor;
7233 let mlp = (cfg.dim * cfg.hidden_dim + cfg.hidden_dim * cfg.dim + cfg.dim * cfg.hidden_dim)
7234 / weight_pack_factor;
7235 let norms = 2 * cfg.dim; let per_layer = (attn + mlp + norms) * elem;
7238
7239 Ok(vec![per_layer; cfg.n_layers])
7240 }
7241
7242 fn num_layers(&self, config: &str) -> Result<usize> {
7243 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7244 Ok(cfg.n_layers)
7245 }
7246
7247 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
7248 let cfg: VoxtralConfig = serde_json::from_str(config)?;
7249
7250 let cfg = ModelConfigMetadata {
7251 max_seq_len: cfg.model_max_length,
7252 num_layers: cfg.n_layers,
7253 hidden_size: cfg.dim,
7254 num_kv_heads: cfg.n_kv_heads,
7255 num_attn_heads: cfg.n_heads,
7256 sliding_window: cfg.sliding_window,
7257 k_head_dim: cfg.head_dim,
7258 v_head_dim: cfg.head_dim,
7259 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
7260 };
7261
7262 Ok(Box::new(cfg))
7263 }
7264}
7265
7266pub struct Gemma4Loader;
7269
7270#[allow(dead_code)]
7271pub struct Gemma4Prefixer;
7272
7273impl MultimodalPromptPrefixer for Gemma4Prefixer {
7274 fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
7275 prompt.to_string()
7276 }
7277 fn prefix_video(&self, _video_indexes: Vec<usize>, prompt: &str) -> String {
7278 prompt.to_string()
7279 }
7280}
7281
7282impl MultimodalModelLoader for Gemma4Loader {
7283 fn load(
7284 &self,
7285 config: &str,
7286 vb: ShardedVarBuilder,
7287 normal_loading_metadata: NormalLoadingMetadata,
7288 attention_mechanism: AttentionImplementation,
7289 ) -> Result<Box<dyn MultimodalModel + Send + Sync>> {
7290 let cfg: Gemma4Config = serde_json::from_str(config)?;
7291 Ok(Box::new(Gemma4Model::new(
7292 &cfg,
7293 vb,
7294 self.is_gptx(config),
7295 normal_loading_metadata,
7296 attention_mechanism,
7297 )?))
7298 }
7299 fn is_gptx(&self, _config: &str) -> bool {
7300 true
7301 }
7302 fn get_config_repr(&self, config: &str) -> Result<Box<dyn Debug>> {
7303 let config: Gemma4Config = serde_json::from_str(config)?;
7304 Ok(Box::new(config))
7305 }
7306 fn get_processor(
7307 &self,
7308 config: &str,
7309 processor_config: Option<ProcessorConfig>,
7310 _preprocessor_config: PreProcessorConfig,
7311 _max_edge: Option<u32>,
7312 ) -> Arc<dyn Processor + Send + Sync> {
7313 let cfg: Gemma4Config = serde_json::from_str(config).expect("Failed to parse Gemma4Config");
7314 Arc::new(Gemma4Processor::new(
7315 processor_config.unwrap_or_default(),
7316 cfg.vision_config.patch_size,
7317 cfg.vision_config.pooling_kernel_size,
7318 cfg.vision_config.default_output_length,
7319 true,
7320 cfg.audio_config.is_some(),
7321 ))
7322 }
7323 fn supports_paged_attention(&self, _config: &str) -> bool {
7324 true
7325 }
7326 fn supports_prefix_cacher(&self, _config: &str) -> bool {
7327 true
7328 }
7329 fn prefixer(&self, _config: &str) -> Arc<dyn MultimodalPromptPrefixer> {
7330 Arc::new(Gemma4Prefixer)
7331 }
7332 fn modalities(&self, config: &str) -> Result<Modalities> {
7333 let cfg: Gemma4Config = serde_json::from_str(config)?;
7334 let mut input = vec![
7335 SupportedModality::Text,
7336 SupportedModality::Vision,
7337 SupportedModality::Video,
7338 ];
7339 if cfg.audio_config.is_some() {
7340 input.push(SupportedModality::Audio);
7341 }
7342 Ok(Modalities {
7343 input,
7344 output: vec![SupportedModality::Text],
7345 })
7346 }
7347}
7348
7349impl IsqModelLoader for Gemma4Loader {
7350 fn isq_layer_regexes(&self, _config: &str) -> Result<Vec<Regex>> {
7351 Ok(vec![
7353 Regex::new(r"lm_head\.(weight|bias)$")?,
7354 Regex::new(r"layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
7355 Regex::new(r"layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
7356 Regex::new(r"layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
7357 Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
7358 Regex::new(r"layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
7359 Regex::new(r"layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
7360 Regex::new(r"layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
7361 Regex::new(r"layers\.(\d+)\.moe\.gate_up_proj\.weight$")?,
7362 Regex::new(r"layers\.(\d+)\.moe\.down_proj\.weight$")?,
7363 Regex::new(r"layers\.(\d+)\.experts\.gate_up_proj\.weight$")?,
7364 Regex::new(r"layers\.(\d+)\.experts\.down_proj\.weight$")?,
7365 Regex::new(r"per_layer_model_projection\.(weight|bias)$")?,
7366 Regex::new(r"layers\.(\d+)\.per_layer_input_gate\.(weight|bias)$")?,
7367 Regex::new(r"layers\.(\d+)\.per_layer_projection\.(weight|bias)$")?,
7368 ])
7369 }
7370 fn immediate_isq_predicates(&self, _config: &str) -> Result<Vec<Regex>> {
7371 Ok(vec![
7372 Regex::new(r"lm_head\.(weight|bias)$")?,
7373 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$")?,
7374 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$")?,
7375 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$")?,
7376 Regex::new(r"model\.language_model\.layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
7377 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate_proj\.(weight|bias)$")?,
7378 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.up_proj\.(weight|bias)$")?,
7379 Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.down_proj\.(weight|bias)$")?,
7380 Regex::new(r"model\.language_model\.layers\.(\d+)\.moe\.gate_up_proj\.weight$")?,
7381 Regex::new(r"model\.language_model\.layers\.(\d+)\.moe\.down_proj\.weight$")?,
7382 Regex::new(r"model\.language_model\.layers\.(\d+)\.experts\.gate_up_proj\.weight$")?,
7383 Regex::new(r"model\.language_model\.layers\.(\d+)\.experts\.down_proj\.weight$")?,
7384 Regex::new(r"model\.language_model\.per_layer_model_projection\.(weight|bias)$")?,
7385 Regex::new(
7386 r"model\.language_model\.layers\.(\d+)\.per_layer_input_gate\.(weight|bias)$",
7387 )?,
7388 Regex::new(
7389 r"model\.language_model\.layers\.(\d+)\.per_layer_projection\.(weight|bias)$",
7390 )?,
7391 ])
7392 }
7393}
7394
7395impl DeviceMappedModelLoader for Gemma4Loader {
7396 fn mapped_max_act_size_elems(
7397 &self,
7398 config: &str,
7399 params: &AutoDeviceMapParams,
7400 ) -> Result<usize> {
7401 let AutoDeviceMapParams::Multimodal {
7402 max_seq_len,
7403 max_batch_size,
7404 max_image_shape: _,
7405 max_num_images,
7406 } = params
7407 else {
7408 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
7409 };
7410
7411 let cfg: Gemma4Config = serde_json::from_str(config)?;
7412 let tc = &cfg.text_config;
7413
7414 let vision_tokens_per_image = cfg.vision_soft_tokens_per_image.unwrap_or(280);
7415 let audio_tokens = if cfg.audio_config.is_some() { 750 } else { 0 };
7416 let total_seq_len = *max_seq_len + vision_tokens_per_image * max_num_images + audio_tokens;
7417 let max_text_attn = max_batch_size * tc.num_attention_heads * total_seq_len * total_seq_len;
7418
7419 Ok(max_text_attn)
7420 }
7421
7422 fn non_mapped_max_act_size_elems(
7423 &self,
7424 config: &str,
7425 params: &AutoDeviceMapParams,
7426 ) -> Result<usize> {
7427 let AutoDeviceMapParams::Multimodal {
7428 max_seq_len: _,
7429 max_batch_size,
7430 max_image_shape: _,
7431 max_num_images,
7432 } = params
7433 else {
7434 anyhow::bail!("Expected multimodal AutoDeviceMapParams for this model!")
7435 };
7436
7437 let cfg: Gemma4Config = serde_json::from_str(config)?;
7438 let vc = &cfg.vision_config;
7439
7440 let max_patches =
7441 vc.default_output_length * vc.pooling_kernel_size * vc.pooling_kernel_size;
7442 let max_vision_attn =
7443 max_batch_size * max_num_images * vc.num_attention_heads * max_patches * max_patches;
7444 let max_vision_hidden = max_batch_size
7445 * max_num_images
7446 * max_patches
7447 * vc.hidden_size.max(vc.intermediate_size);
7448
7449 let max_audio_activation = cfg.audio_config.as_ref().map_or(0, |audio_cfg| {
7450 let subsample_factor: usize = audio_cfg
7451 .sscp_conv_stride_size
7452 .iter()
7453 .map(|stride| stride[0])
7454 .product();
7455 let max_audio_frames = 750 * subsample_factor.max(1);
7456 let audio_seq_after_subsample = max_audio_frames / subsample_factor.max(1);
7457
7458 let audio_encoder_act = audio_seq_after_subsample * (audio_cfg.hidden_size * 4);
7459 let chunk_size = audio_cfg.conf_attention_chunk_size;
7460 let context_size = chunk_size + audio_cfg.conf_attention_context_left - 1
7461 + audio_cfg.conf_attention_context_right;
7462 let num_chunks = audio_seq_after_subsample.div_ceil(chunk_size);
7463 let audio_attn_act =
7464 audio_cfg.conf_num_attention_heads * num_chunks * chunk_size * context_size;
7465
7466 max_batch_size * audio_encoder_act.max(audio_attn_act)
7467 });
7468
7469 Ok(max_vision_attn
7470 .max(max_vision_hidden)
7471 .max(max_audio_activation))
7472 }
7473
7474 fn non_mapped_size_in_bytes(
7475 &self,
7476 config: &str,
7477 dtype: DType,
7478 weight_pack_factor: usize,
7479 _matformer_config: Option<&MatformerSliceConfig>,
7480 ) -> Result<usize> {
7481 let cfg: Gemma4Config = serde_json::from_str(config)?;
7482 let tc = &cfg.text_config;
7483 let vc = &cfg.vision_config;
7484
7485 let text_elems = {
7486 let embed_tokens = tc.hidden_size * tc.vocab_size;
7487 let lm_head = if !tc.tie_word_embeddings || weight_pack_factor != 1 {
7488 tc.hidden_size * tc.vocab_size / weight_pack_factor
7489 } else {
7490 0
7491 };
7492 let norm = tc.hidden_size;
7493
7494 let ple_dim = tc.hidden_size_per_layer_input.unwrap_or(0);
7495 let ple_vocab = tc.vocab_size_per_layer_input.unwrap_or(tc.vocab_size);
7496 let embed_tokens_per_layer = if ple_dim > 0 {
7497 ple_vocab * tc.num_hidden_layers * ple_dim
7498 } else {
7499 0
7500 };
7501 let per_layer_model_projection = if ple_dim > 0 {
7502 tc.hidden_size * tc.num_hidden_layers * ple_dim / weight_pack_factor
7503 } else {
7504 0
7505 };
7506 let per_layer_projection_norm = ple_dim;
7507
7508 embed_tokens
7509 + lm_head
7510 + norm
7511 + embed_tokens_per_layer
7512 + per_layer_model_projection
7513 + per_layer_projection_norm
7514 };
7515
7516 let vision_layer_elems = {
7517 let quantized = vc.hidden_size * vc.num_attention_heads * vc.head_dim
7518 + 3 * (vc.hidden_size * vc.num_key_value_heads * vc.head_dim)
7519 + 2 * (vc.hidden_size * vc.intermediate_size)
7520 + vc.intermediate_size * vc.hidden_size;
7521 let norms = 2 * vc.head_dim + 4 * vc.hidden_size;
7522 quantized / weight_pack_factor + norms
7523 };
7524 let vision_elems = {
7525 let patch_embed = vc.patch_size * vc.patch_size * 3 * vc.hidden_size;
7526 let position_embedding_table = 2 * vc.position_embedding_size * vc.hidden_size;
7527 let patch_embedder = patch_embed / weight_pack_factor + position_embedding_table;
7528 let encoder = vc.num_hidden_layers * vision_layer_elems;
7529 let embed_vision = vc.hidden_size * tc.hidden_size / weight_pack_factor;
7530
7531 patch_embedder + encoder + embed_vision
7532 };
7533
7534 let audio_elems = cfg.audio_config.as_ref().map_or(0, |audio_cfg| {
7535 let mut f_out = audio_cfg.input_feat_size;
7536 for i in 0..2 {
7537 let kernel_w = audio_cfg.sscp_conv_kernel_size[i][1];
7538 let stride_w = audio_cfg.sscp_conv_stride_size[i][1];
7539 let pad_left = 1;
7540 let pad_right = 1;
7541 f_out = (f_out + pad_left + pad_right + stride_w - kernel_w) / stride_w;
7542 }
7543
7544 let subsample_conv_projection = {
7545 let conv_0 = audio_cfg.sscp_conv_channel_size[0]
7546 * audio_cfg.sscp_conv_kernel_size[0][0]
7547 * audio_cfg.sscp_conv_kernel_size[0][1];
7548 let conv_1 = audio_cfg.sscp_conv_channel_size[0]
7549 * audio_cfg.sscp_conv_channel_size[1]
7550 * audio_cfg.sscp_conv_kernel_size[1][0]
7551 * audio_cfg.sscp_conv_kernel_size[1][1];
7552 let norms =
7553 audio_cfg.sscp_conv_channel_size[0] + audio_cfg.sscp_conv_channel_size[1];
7554 let input_proj =
7555 audio_cfg.sscp_conv_channel_size[1] * f_out * audio_cfg.hidden_size
7556 / weight_pack_factor;
7557 conv_0 + conv_1 + norms + input_proj
7558 };
7559
7560 let conformer_block = {
7561 let attention = 5 * (audio_cfg.hidden_size * audio_cfg.hidden_size)
7562 / weight_pack_factor
7563 + 2 * audio_cfg.hidden_size
7564 + audio_cfg.hidden_size / audio_cfg.conf_num_attention_heads
7565 + audio_cfg.hidden_size / 2
7566 + (audio_cfg.conf_attention_context_left
7567 + audio_cfg.conf_attention_context_right
7568 + 1)
7569 + (audio_cfg.conf_attention_chunk_size
7570 * (audio_cfg.conf_attention_chunk_size
7571 + audio_cfg.conf_attention_context_left
7572 - 1
7573 + audio_cfg.conf_attention_context_right))
7574 + 1;
7575 let ffw = 2
7576 * (2 * audio_cfg.hidden_size
7577 + 2 * (audio_cfg.hidden_size * (audio_cfg.hidden_size * 4))
7578 / weight_pack_factor);
7579 let conv = 2 * audio_cfg.hidden_size
7580 + audio_cfg.hidden_size * (audio_cfg.hidden_size * 2) / weight_pack_factor
7581 + audio_cfg.hidden_size * audio_cfg.hidden_size / weight_pack_factor
7582 + audio_cfg.hidden_size * audio_cfg.conf_conv_kernel_size;
7583 attention + ffw + conv + audio_cfg.hidden_size
7584 };
7585
7586 let output_proj = audio_cfg.output_proj_dims.map_or(0, |output_dim| {
7587 audio_cfg.hidden_size * output_dim / weight_pack_factor + output_dim
7588 });
7589 let audio_embed_hidden = audio_cfg.output_proj_dims.unwrap_or(audio_cfg.hidden_size);
7590 let embed_audio = audio_embed_hidden * tc.hidden_size / weight_pack_factor;
7591
7592 subsample_conv_projection
7593 + audio_cfg.conf_num_hidden_layers * conformer_block
7594 + output_proj
7595 + embed_audio
7596 });
7597
7598 let vision_dtype = if dtype == DType::F16 {
7599 DType::F32
7600 } else {
7601 dtype
7602 };
7603
7604 Ok(text_elems * dtype.size_in_bytes()
7605 + vision_elems * vision_dtype.size_in_bytes()
7606 + audio_elems * dtype.size_in_bytes())
7607 }
7608
7609 fn layer_sizes_in_bytes(
7610 &self,
7611 config: &str,
7612 dtype: DType,
7613 weight_pack_factor: usize,
7614 _matformer_config: Option<&MatformerSliceConfig>,
7615 ) -> Result<Vec<usize>> {
7616 let cfg: Gemma4Config = serde_json::from_str(config)?;
7617 let tc = &cfg.text_config;
7618 let sizes: Vec<usize> = (0..tc.num_hidden_layers)
7619 .map(|layer_idx| {
7620 let is_sliding = {
7621 let is_last = layer_idx == tc.num_hidden_layers - 1;
7622 !is_last && (layer_idx + 1) % tc.sliding_window_pattern != 0
7623 };
7624 let hd = if is_sliding {
7625 tc.head_dim
7626 } else {
7627 tc.global_head_dim
7628 };
7629 let nkv = if is_sliding {
7630 tc.num_key_value_heads
7631 } else {
7632 tc.num_global_key_value_heads
7633 .unwrap_or(tc.num_key_value_heads)
7634 };
7635 let use_k_eq_v = tc.attention_k_eq_v && !is_sliding;
7636
7637 let mut attn = tc.hidden_size * tc.num_attention_heads * hd
7638 + tc.hidden_size * nkv * hd
7639 + tc.num_attention_heads * hd * tc.hidden_size;
7640 if !use_k_eq_v {
7641 attn += tc.hidden_size * nkv * hd;
7642 }
7643 attn += 2 * hd;
7644
7645 let mlp = 3 * tc.hidden_size * tc.intermediate_size;
7646
7647 let moe = if tc.enable_moe_block {
7648 let ne = tc.num_experts.unwrap_or(0);
7649 let ei = tc.expert_intermediate_size.unwrap_or(0);
7650 ne * tc.hidden_size * ei * 2
7651 + ne * ei * tc.hidden_size
7652 + ne
7653 + ne * tc.hidden_size
7654 + tc.hidden_size
7655 + 3 * tc.hidden_size
7656 } else {
7657 0
7658 };
7659
7660 let ple = if tc.hidden_size_per_layer_input.unwrap_or(0) > 0 {
7661 let pd = tc.hidden_size_per_layer_input.unwrap();
7662 tc.hidden_size * pd + pd * tc.hidden_size + tc.hidden_size
7663 } else {
7664 0
7665 };
7666
7667 let norms = 4 * tc.hidden_size + 1;
7668
7669 (attn + mlp + moe + ple + norms) * dtype.size_in_bytes() / weight_pack_factor
7670 })
7671 .collect();
7672 Ok(sizes)
7673 }
7674
7675 fn num_layers(&self, config: &str) -> Result<usize> {
7676 let cfg: Gemma4Config = serde_json::from_str(config)?;
7677 Ok(cfg.text_config.num_hidden_layers)
7678 }
7679
7680 fn non_mapped_sub_models(&self) -> Option<Vec<NonMappedSubModel>> {
7681 Some(vec![NonMappedSubModel::Vision, NonMappedSubModel::Audio])
7682 }
7683
7684 fn model_config(&self, config: &str) -> Result<Box<dyn ModelConfigLike>> {
7685 let cfg: Gemma4Config = serde_json::from_str(config)?;
7686 let tc = &cfg.text_config;
7687
7688 let cfg = ModelConfigMetadata {
7689 max_seq_len: tc.max_position_embeddings,
7690 num_layers: tc.num_hidden_layers,
7691 hidden_size: tc.hidden_size,
7692 num_kv_heads: tc.num_key_value_heads,
7693 num_attn_heads: tc.num_attention_heads,
7694 sliding_window: Some(tc.sliding_window),
7695 k_head_dim: tc.global_head_dim,
7696 v_head_dim: tc.global_head_dim,
7697 kv_cache_layout: crate::paged_attention::KvCacheLayout::Standard,
7698 };
7699
7700 Ok(Box::new(cfg))
7701 }
7702}