Skip to main content

mistralrs_core/pipeline/loaders/
multimodal_loaders.rs

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    // pixel_values and pixel_attention_mask only specified for prompt seqs
82    #[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>, // pixel attention mask, or image sizes, or anything else
91        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    /// For a prompt without images. Requires batch size of 1!
103    fn default_model_specific_args(&self, input_ids: &Tensor) -> Box<dyn Any>;
104    /// Return encoder cache hit/miss counters (hits, misses) if this model has an encoder cache.
105    fn encoder_cache_counters(&self) -> Option<(Arc<AtomicUsize>, Arc<AtomicUsize>)> {
106        None
107    }
108    /// Reset model-specific state (e.g. cached audio embeddings) between requests.
109    /// Called when the pipeline's non-granular state is reset.
110    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        // Default is false, specific model must override.
133        false
134    }
135    fn modalities(&self, config: &str) -> Result<Modalities>;
136    fn prefixer(&self, config: &str) -> Arc<dyn MultimodalPromptPrefixer>;
137    /// Return a default chat template (Jinja string) for models that don't ship a
138    /// `tokenizer_config.json` or `chat_template.jinja`. Returns `None` by default.
139    /// The `config` parameter is the raw model config JSON, used by `AutoMultimodalLoader`
140    /// to delegate to the correct concrete loader.
141    fn default_chat_template(&self, _config: &str) -> Option<String> {
142        None
143    }
144    /// Return default (bos_token, eos_token) strings for models that don't ship a
145    /// `tokenizer_config.json`. Used to populate the chat template context and
146    /// EOS token detection. Returns `None` by default.
147    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)]
181/// The architecture to load the multimodal model as.
182pub 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
225// https://github.com/huggingface/transformers/blob/cff06aac6fad28019930be03f5d467055bf62177/src/transformers/models/auto/modeling_auto.py#L448
226impl 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    /// Voxtral params.json uses a `multimodal` key instead of `architectures`.
319    #[serde(default)]
320    multimodal: Option<serde_json::Value>,
321}
322
323/// Automatically selects a MultimodalModelLoader implementation based on the JSON `architectures` field.
324pub 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        // Voxtral: params.json has `multimodal` but no `architectures`
331        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        // Delegate to the concrete loader
346        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
563// ======================== Phi 3 loader
564
565/// [`MultimodalLoader`] for a Phi 3 Vision model.
566///
567/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
568pub struct Phi3VLoader;
569
570pub struct Phi3VPrefixer;
571
572impl MultimodalPromptPrefixer for Phi3VPrefixer {
573    fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
574        // Image indexing starts at 0.
575        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            // Attention
640            Regex::new(r"layers\.(\d+)\.self_attn\.qkv_proj\.(weight|bias)$")?,
641            Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
642            // MLP
643            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        // NOTE: we ignore max_num_images although it can only be one...
659        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            // This model injects the vision information directly into the input embeddings
678            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        // NOTE: we ignore max_num_images although it can only be one...
691        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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
726            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
845// ======================== Idefics 2 loader
846
847/// [`MultimodalLoader`] for an Idefics 2 Vision model.
848///
849/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
850pub struct Idefics2Loader;
851
852pub struct Idefics2Prefixer;
853
854impl MultimodalPromptPrefixer for Idefics2Prefixer {
855    fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
856        // Chat template does it
857        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            // Attention
920            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            // MLP
925            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            // Attention
934            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            // MLP
939            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            // This model injects the vision information directly into the input embeddings
969            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            // do_image_splitting = true
998            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
1196// ======================== LLaVANext Loader
1197
1198/// [`MultimodalLoader`] for an LLaVANext Vision model.
1199///
1200/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
1201pub 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            // Attention
1266            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            // MLP
1271            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            // Attention
1280            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            // MLP
1285            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            // This model injects the vision information directly into the input embeddings
1321            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
1466// ======================== LLaVA Loader
1467
1468/// [`MultimodalLoader`] for an LLaVA Vision model.
1469///
1470/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
1471pub 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            // Attention
1536            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            // MLP
1541            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            // Attention
1550            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            // MLP
1555            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            // This model injects the vision information directly into the input embeddings
1587            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
1728// ======================== MLlama Loader
1729
1730/// [`MultimodalLoader`] for an Llama Vision model.
1731///
1732/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
1733pub 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                // Attention text
1803                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                // MLP text
1816                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            // Vision attention (transformer)
1829            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            // Vision attention (global transforemr)
1842            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            // MLP vision
1855            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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
1948            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                // No isq for cross attention
2045                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
2114// ======================== Qwen2VL Loader
2115
2116/// [`MultimodalLoader`] for an Qwen2-VL model.
2117///
2118/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
2119pub 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            // Attention
2190            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            // MLP
2195            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        // For images, grid_t=1. After spatial merging, grid_h and grid_w are reduced.
2224        let img_seq_len = {
2225            let cfg = &cfg.vision_config;
2226            // grid_t is 1 for images (temporal dimension is for video only)
2227            let grid_t = 1;
2228            // After patch embedding and spatial merge, the effective grid dimensions are reduced
2229            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            // This model injects the vision information directly into the input embeddings
2236            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        // For the vision encoder, before spatial merging
2261        let img_seq_len = {
2262            let cfg = &cfg.vision_config;
2263            // grid_t is 1 for images
2264            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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
2289            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
2415// ======================== Idefics 3 loader
2416
2417/// [`MultimodalLoader`] for an Idefics 3 Vision model.
2418///
2419/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
2420pub struct Idefics3Loader;
2421
2422pub struct Idefics3Prefixer;
2423
2424impl MultimodalPromptPrefixer for Idefics3Prefixer {
2425    fn prefix_image(&self, _image_indexes: Vec<usize>, prompt: &str) -> String {
2426        // Chat template does it
2427        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            // Attention
2490            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            // MLP
2495            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            // Attention
2504            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            // MLP
2509            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            // // Attention (vision)
2513            // Regex::new(
2514            //     r"model\.vision_model\.encoder\.layers\.(\d+)\.self_attn\.q_proj\.(weight|bias)$",
2515            // )?,
2516            // Regex::new(
2517            //     r"model\.vision_model\.encoder\.layers\.(\d+)\.self_attn\.k_proj\.(weight|bias)$",
2518            // )?,
2519            // Regex::new(
2520            //     r"model\.vision_model\.encoder\.layers\.(\d+)\.self_attn\.v_proj\.(weight|bias)$",
2521            // )?,
2522            // Regex::new(
2523            //     r"model\.vision_model\.encoder\.layers\.(\d+)\.self_attn\.out_proj\.(weight|bias)$",
2524            // )?,
2525            // MLP (vision)
2526            // Regex::new(r"model\.vision_model\.encoder\.layers\.(\d+)\.mlp\.fc1\.(weight|bias)$")?,
2527            // Regex::new(r"model\.vision_model\.encoder\.layers\.(\d+)\.mlp\.fc2\.(weight|bias)$")?,
2528        ])
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            // This model injects the vision information directly into the input embeddings
2555            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            // do_image_splitting = true
2584            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
2733// ======================== MiniCpm-O loader
2734
2735/// [`MultimodalLoader`] for an MiniCpm-O model.
2736///
2737/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
2738pub 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            // Attention
2807            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            // MLP
2812            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            // This model injects the vision information directly into the input embeddings
2845            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            // do_image_splitting = true
2874            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
3012// ======================== Phi 4MM loader
3013
3014/// [`MultimodalLoader`] for a Phi 4MM Vision model.
3015///
3016/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
3017pub struct Phi4MMLoader;
3018
3019pub struct Phi4MMPrefixer;
3020
3021impl MultimodalPromptPrefixer for Phi4MMPrefixer {
3022    fn prefix_image(&self, image_indexes: Vec<usize>, prompt: &str) -> String {
3023        // Image indexing starts at 0.
3024
3025        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        // Image indexing starts at 0.
3035
3036        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            // Attention
3105            Regex::new(r"layers\.(\d+)\.self_attn\.qkv_proj\.(weight|bias)$")?,
3106            Regex::new(r"layers\.(\d+)\.self_attn\.o_proj\.(weight|bias)$")?,
3107            // MLP
3108            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        // NOTE: we ignore max_num_images although it can only be one...
3124        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            // This model injects the vision information directly into the input embeddings
3143            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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
3203            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
3358// ======================== Qwen2_5VL Loader
3359
3360/// [`MultimodalLoader`] for an Qwen2_5VL model.
3361///
3362/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
3363pub 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            // Attention
3434            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            // MLP
3439            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            // This model injects the vision information directly into the input embeddings
3478            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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
3529            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
3654// ======================== Gemma 3 Loader
3655
3656/// [`MultimodalLoader`] for an Gemma 3 model.
3657///
3658/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
3659pub 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        // Handle the Gemma 3 1b case here
3702        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            // Attention
3729            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            // MLP
3734            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            // Attention
3743            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            // MLP
3748            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                    // This model injects the vision information directly into the input embeddings
3787                    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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
3846            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, // None to be more forgiving, some do not
3980            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
3993// ======================== Mistral 3 Loader
3994
3995/// [`MultimodalLoader`] for an Mistral 3 model.
3996///
3997/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
3998pub 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            // Attention
4064            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            // MLP
4069            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            // Attention
4078            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            // MLP
4083            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            // Reshaping algorithm
4113
4114            // https://huggingface.co/mistralai/Mistral-Small-3.1-24B-Instruct-2503/blob/main/preprocessor_config.json#L29
4115            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        // This model injects the vision information directly into the input embeddings
4135        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            // Reshaping algorithm
4159
4160            // https://huggingface.co/mistralai/Mistral-Small-3.1-24B-Instruct-2503/blob/main/preprocessor_config.json#L29
4161            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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
4197            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
4318// ======================== Llama 4 Loader
4319
4320/// [`MultimodalLoader`] for an Llama Vision model.
4321///
4322/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
4323pub 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            // Attention
4390            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            // FF MoE
4395            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            // FF MLP
4404            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            // Attention
4413            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            // FF MoE
4418            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            // FF MLP
4443            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    /// This incorporates the max batch size!
4458    /// Returns (pixels max batch size, num text image tokens)
4459    #[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
4724// ======================== Gemma 3n Loader
4725
4726/// [`MultimodalLoader`] for an Gemma 3n model.
4727///
4728/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
4729pub 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        // Handle the Gemma 3 1b case here
4772        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            // Language model attention
4803            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            // Language model MLP
4808            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            // Audio conformer attention layers
4812            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            // Audio conformer FFW layers
4820            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            // Audio conformer conv1d layers
4825            Regex::new(r"conformer\.(\d+)\.lconv1d\.linear_start\.(weight|bias)$")?,
4826            Regex::new(r"conformer\.(\d+)\.lconv1d\.linear_end\.(weight|bias)$")?,
4827            // Audio subsample projection
4828            Regex::new(r"subsample_conv_projection\.input_proj_linear\.(weight|bias)$")?,
4829            // Multimodal embedders
4830            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            // Language model attention
4838            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            // Language model MLP
4843            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            // Projections
4847            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            // Audio conformer attention layers
4851            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            // Audio conformer FFW layers
4865            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            // Audio conformer conv1d layers
4878            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            // Audio subsample projection
4885            Regex::new(
4886                r"model\.audio_tower\.subsample_conv_projection\.input_proj_linear\.(weight|bias)$",
4887            )?,
4888            // Multimodal embedders
4889            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        // Gemma3n is an "inject into the prompt" model, similar to Gemma3
4915        // We need to account for vision and audio tokens in the sequence length
4916
4917        let mut total_seq_len = *max_seq_len.min(&ATTENTION_CHUNK_SIZE);
4918
4919        // Add vision tokens
4920        {
4921            // Vision tokens are injected into the prompt
4922            // MSFA outputs fixed 16x16 features regardless of input size
4923            let msfa_spatial_size = 16; // Fixed from vision.rs line 1115
4924            let vision_tokens_per_image = msfa_spatial_size * msfa_spatial_size; // 256 tokens
4925            total_seq_len += vision_tokens_per_image * max_num_images;
4926        }
4927
4928        // Add audio tokens
4929        {
4930            // Audio tokens are injected into the prompt
4931            // From config field audio_soft_tokens_per_image (typically 188)
4932            let audio_tokens = cfg.audio_soft_tokens_per_image;
4933            total_seq_len += audio_tokens;
4934        }
4935
4936        // Calculate max attention size for text model with all injected tokens
4937        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        // Calculate max activation sizes for each modality
4961        let mut max_activation = 0;
4962
4963        // Vision activation size
4964        {
4965            // Vision is Gemma3n's MobileNetV5 architecture with Multi-Query Attention
4966            // The peak activation is in the Multi-Query Attention layers
4967
4968            // From the architecture: stages 3 and 4 have MMQA blocks
4969            // Input images are 768x768 (from inputs_processor.rs)
4970            // Stage 3: 640 channels at 48x48 (768/16 downsampling), MMQA with num_heads=12, kv_dim=64
4971            // Stage 4: 1280 channels at 24x24 (768/32 downsampling), MMQA with num_heads=16, kv_dim=96
4972            // MSFA output: 2048 channels at fixed 16x16
4973
4974            let vision_tower_act = {
4975                // Peak is during MMQA attention computation in stage 4
4976                // Stage 4 has higher memory usage than Stage 3 due to more heads (16 vs 12)
4977                // From vision.rs: Stage 4 has num_heads=16, kv_dim=96, kv_stride=1
4978                let num_heads = 16; // Stage 4 configuration
4979                let spatial_size = 24; // 768 / 32 = 24 (input 768x768, stage 4 has 32x downsampling)
4980                let seq_len = spatial_size * spatial_size;
4981
4982                // Attention scores: [B * num_images, num_heads, seq_len, seq_len]
4983                max_batch_size * max_num_images * num_heads * seq_len * seq_len
4984            };
4985
4986            // Vision embedder activations
4987            let vision_embed_act = {
4988                // MSFA output: 2048 channels at fixed 16x16 spatial (from vision.rs line 1115)
4989                let msfa_channels = 2048; // MSFA_OUT_CHANNELS from vision.rs
4990                let spatial_size = 16; // Fixed output resolution from MSFA
4991                let vision_features =
4992                    max_batch_size * max_num_images * msfa_channels * spatial_size * spatial_size;
4993
4994                // After embedding projection to text hidden size
4995                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        // Audio activation size
5008        {
5009            let audio_cfg = &cfg.audio_config;
5010
5011            // Calculate max audio sequence length based on config
5012            // Audio uses conformer with subsampling and reduction
5013
5014            // A rough estimate of max_audio_frames
5015            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]) // Time dimension stride
5021                .product();
5022            let audio_seq_after_subsample = max_audio_frames / subsample_factor;
5023
5024            // Audio encoder activations
5025            let audio_encoder_act = {
5026                // Conformer FFW layers have expansion factor from config
5027                let intermediate_size = audio_cfg.hidden_size * 4; // FFW expansion factor
5028
5029                // Peak is in the FFW layers before reduction
5030                max_batch_size * audio_seq_after_subsample * intermediate_size
5031            };
5032
5033            // Audio attention activations
5034            let audio_attn_act = {
5035                // Attention uses chunked processing with specific context sizes
5036                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                // Peak is attention scores: [B, num_heads, num_chunks, chunk_size, context_size]
5041                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        // Apply matformer slicing if configured
5066        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        // Text components that are not device-mapped
5086        let text_elems = {
5087            // Embeddings
5088            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            // LM head (if not tied)
5094            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            // Final layer norm
5101            let norm = text_cfg.hidden_size;
5102
5103            // AltUp projections (not device-mapped)
5104            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            // Per-layer model projection
5112            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        // Vision components
5129        let vision_elems = {
5130            let multimodal_cfg = &cfg.vision_config;
5131            // Vision tower - calculated from actual Gemma3n architecture
5132            // NOTE: Vision tower uses only Conv2d layers, NOT Arc<dyn QuantMethod>,
5133            // so NONE of these should be divided by weight_pack_factor
5134            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                // Stem: ConvNormAct (Conv2d + RMSNorm)
5142                let stem_conv =
5143                    INPUT_CHANNELS * STEM_OUT_CHANNELS * STEM_KERNEL_SIZE * STEM_KERNEL_SIZE;
5144                let stem_norm = STEM_OUT_CHANNELS; // RMSNorm weight
5145
5146                // Track input channels through the network
5147                let mut in_chs = STEM_OUT_CHANNELS;
5148                let mut total_elems = stem_conv + stem_norm;
5149
5150                // Process all stages from gemma3n_mobilenet_def
5151                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                                // EdgeResidual: all Conv2d layers, not quantizable
5166                                total_elems += in_chs * mid_chs * kernel_size * kernel_size; // conv_exp (Conv2d)
5167                                total_elems += mid_chs; // bn1 weight
5168                                total_elems += mid_chs * out_channels; // conv_pwl (Conv2d)
5169                                total_elems += out_channels; // bn2 weight
5170                                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                                // UniversalInvertedResidual: all Conv2d layers, not quantizable
5183                                if *expand_ratio != 1.0 {
5184                                    total_elems += in_chs * mid_chs; // expand conv (Conv2d)
5185                                    total_elems += mid_chs; // expand norm
5186                                }
5187                                if *start_kernel_size > 0 {
5188                                    total_elems += mid_chs * start_kernel_size * start_kernel_size; // depthwise start (Conv2d)
5189                                    total_elems += mid_chs; // norm
5190                                }
5191                                if *mid_kernel_size > 0 {
5192                                    total_elems += mid_chs * mid_kernel_size * mid_kernel_size; // depthwise mid (Conv2d)
5193                                    total_elems += mid_chs; // norm
5194                                }
5195                                total_elems += mid_chs * out_channels; // project conv (Conv2d)
5196                                total_elems += out_channels; // project norm
5197                                total_elems += out_channels; // layer scale gamma
5198                                in_chs = *out_channels;
5199                            }
5200                            BlockType::MultiQueryAttention {
5201                                num_heads,
5202                                kv_dim,
5203                                kv_stride: _,
5204                                ..
5205                            } => {
5206                                // MMQA: all Conv2d layers, not quantizable
5207                                let dw_kernel_size = 3; // Default dw_kernel_size for MMQA
5208                                total_elems += in_chs; // norm weight
5209                                total_elems += in_chs * num_heads * kv_dim; // query_proj (Conv2d)
5210                                total_elems += in_chs * kv_dim; // key_proj (Conv2d)
5211                                total_elems += in_chs * dw_kernel_size * dw_kernel_size; // key_dw_conv (Conv2d)
5212                                total_elems += *kv_dim; // value_down_conv (Conv2d)
5213                                total_elems += 1; // value_norm weight
5214                                total_elems += *kv_dim; // value_proj (Conv2d)
5215                                total_elems += num_heads * kv_dim * in_chs; // output_proj (Conv2d)
5216                                total_elems += in_chs; // layer scale
5217                            }
5218                        }
5219                    }
5220                }
5221
5222                // Multi-scale fusion adapter (msfa) - also uses Conv2d layers
5223                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                // MSFA FFN (UIR with expansion_ratio) - Conv2d layers, not quantizable
5229                total_elems += msfa_in * msfa_mid; // expand (Conv2d)
5230                total_elems += msfa_mid; // expand norm
5231                total_elems += msfa_mid * msfa_out; // project (Conv2d)
5232                total_elems += msfa_out; // project norm
5233                total_elems += msfa_out; // final norm
5234
5235                total_elems
5236            };
5237
5238            // Vision multimodal embedder components
5239            let embed_vision_elems = {
5240                // Embedding layer (not quantizable)
5241                let embedding = multimodal_cfg.vocab_size * multimodal_cfg.hidden_size;
5242
5243                // Normalization layers (not quantizable)
5244                let hard_norm = multimodal_cfg.hidden_size;
5245                let soft_norm = multimodal_cfg.hidden_size;
5246
5247                // Projection from vision to text hidden size (IS Arc<dyn QuantMethod>, so quantizable)
5248                let projection =
5249                    multimodal_cfg.hidden_size * text_cfg.hidden_size / weight_pack_factor;
5250
5251                // Post-projection norm (not quantizable)
5252                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        // Audio components - based on actual audio.rs structure
5261        let audio_elems = {
5262            let audio_cfg = &cfg.audio_config;
5263
5264            // SubSampleConvProjection components
5265            let subsample_conv_projection_elems = {
5266                // Conv blocks (Conv2d layers - NOT quantizable)
5267                let mut conv_elems = 0;
5268
5269                // conv_0: Conv2d from 1 channel to first channel size
5270                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                // conv_1: Conv2d from first to second channel size
5276                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                // CumulativeGroupNorm for each conv block (weight only, no bias by default)
5282                let norm_0 = out_ch_0; // norm weight for conv_0
5283                let norm_1 = out_ch_1; // norm weight for conv_1
5284
5285                // input_proj_linear (Arc<dyn QuantMethod> - IS quantizable)
5286                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            // Conformer blocks
5302            let conformer_elems = {
5303                let mut total = 0;
5304
5305                for _ in 0..audio_cfg.conf_num_hidden_layers {
5306                    // ConformerAttention
5307                    let attention_elems = {
5308                        // Norms (NOT quantizable)
5309                        let pre_attn_norm = audio_cfg.hidden_size;
5310                        let post_norm = audio_cfg.hidden_size;
5311
5312                        // Attention projections (Arc<dyn QuantMethod> - IS quantizable)
5313                        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                        // RelativePositionEmbedding
5323                        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; // head_dim
5327                        let inv_timescales = audio_cfg.hidden_size / 2; // num_timescales
5328                        let pos_indices = audio_cfg.conf_attention_context_left
5329                            + audio_cfg.conf_attention_context_right
5330                            + 1;
5331
5332                        // Local causal masks (precomputed tensors)
5333                        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; // U8 tensor
5337                        let invalid_logits_tensor = 1; // single f32 value
5338
5339                        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                    // ConformerFeedForward (start and end)
5354                    let ffw_elems = {
5355                        // Each FFW has:
5356                        // - pre_layer_norm (NOT quantizable)
5357                        // - ffw_layer_1 (Arc<dyn QuantMethod> - IS quantizable)
5358                        // - ffw_layer_2 (Arc<dyn QuantMethod> - IS quantizable)
5359                        // - post_layer_norm (NOT quantizable)
5360                        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; // Same structure
5373
5374                        ffw_start + ffw_end
5375                    };
5376
5377                    // ConformerLightConv1d
5378                    let lconv1d_elems = {
5379                        // Norms (NOT quantizable)
5380                        let pre_layer_norm = audio_cfg.hidden_size;
5381                        let conv_norm = audio_cfg.hidden_size;
5382
5383                        // Linear layers (Arc<dyn QuantMethod> - IS quantizable)
5384                        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                        // depthwise_conv1d (Conv1d - NOT quantizable)
5390                        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                    // Final norm for conformer block (NOT quantizable)
5396                    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            // Audio multimodal embedder (embed_audio)
5405            let embed_audio_elems = {
5406                // Embedding layer (ScaledEmbedding - NOT quantizable)
5407                let embedding = audio_cfg.vocab_size * audio_cfg.hidden_size;
5408
5409                // RMS norms (NOT quantizable)
5410                let hard_embedding_norm = audio_cfg.hidden_size; // with scale
5411                let soft_embedding_norm = audio_cfg.hidden_size; // with scale
5412                let embedding_post_projection_norm = text_cfg.hidden_size; // without scale
5413
5414                // Projection (Arc<dyn QuantMethod> - IS quantizable)
5415                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            // f16 -> f32 for vision model in particular.
5430            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        // Apply matformer slicing if configured
5452        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        // When matformer slicing is applied, we only include the layers that are kept
5474        let mut layer_sizes = Vec::new();
5475
5476        // Note: We don't need orig_intermediate_sizes anymore since the adjusted config
5477        // already has the correct intermediate sizes after matformer slicing
5478
5479        for layer_idx in 0..text_cfg.num_hidden_layers {
5480            let per_layer_elems = {
5481                // Layer norms
5482                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                // Attention components
5489                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                // Q, K, V norms
5499                let q_norm = text_cfg.head_dim;
5500                let k_norm = text_cfg.head_dim;
5501                let v_norm = text_cfg.head_dim; // No bias for v_norm
5502
5503                // MLP components - use the adjusted intermediate sizes from matformer
5504                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                // AltUp components (per layer)
5514                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                // Laurel block components
5530                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                // Per-layer input components
5539                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, // None to be more forgiving, some do not
5587            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
5600// ======================== Qwen3VL Loader
5601
5602/// [`MultimodalLoader`] for an Qwen3VL model.
5603///
5604/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
5605pub struct Qwen3VLLoader;
5606
5607pub struct Qwen3VLPrefixer;
5608
5609impl MultimodalPromptPrefixer for Qwen3VLPrefixer {
5610    // No-op: With MessagesAction::Keep, the chat template handles image tokens
5611    // when it sees {"type": "image"} entries in the content.
5612}
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            // Attention
5669            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            // MLP
5674            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        // For images, grid_t=1. After spatial merging, grid_h and grid_w are reduced.
5703        let img_seq_len = {
5704            let cfg = &cfg.vision_config;
5705            // grid_t is 1 for images (temporal dimension is for video only)
5706            let grid_t = 1;
5707            // After patch embedding and spatial merge, the effective grid dimensions are reduced
5708            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            // This model injects the vision information directly into the input embeddings
5716            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        // For the vision encoder, before spatial merging
5741        let img_seq_len = {
5742            let cfg = &cfg.vision_config;
5743            // grid_t is 1 for images
5744            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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
5771            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            // Main merger: norm uses cfg.hidden_size
5788            let ln_q = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
5789            let merger = mlp0 + mlp2 + ln_q;
5790
5791            // Deepstack mergers: norm uses merged hidden_size
5792            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
5922// ======================== Qwen3VLMoE Loader
5923
5924/// [`MultimodalLoader`] for a Qwen3VLMoE model.
5925///
5926/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
5927pub struct Qwen3VLMoELoader;
5928
5929pub struct Qwen3VLMoEPrefixer;
5930
5931impl MultimodalPromptPrefixer for Qwen3VLMoEPrefixer {
5932    // No-op: With MessagesAction::Keep, the chat template handles image tokens
5933    // when it sees {"type": "image"} entries in the content.
5934}
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            // Attention
5991            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            // MLP (dense layers)
5996            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            // MoE router
6000            Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate\.(weight|bias)$")?,
6001            // MoE experts - now unpacked into individual experts
6002            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            // MLP (dense layers)
6020            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            // MoE router
6024            Regex::new(r"model\.language_model\.layers\.(\d+)\.mlp\.gate\.(weight|bias)$")?,
6025            // MoE experts
6026            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        // For images, grid_t=1. After spatial merging, grid_h and grid_w are reduced.
6061        let img_seq_len = {
6062            let cfg = &cfg.vision_config;
6063            // grid_t is 1 for images (temporal dimension is for video only)
6064            let grid_t = 1;
6065            // After patch embedding and spatial merge, the effective grid dimensions are reduced
6066            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            // This model injects the vision information directly into the input embeddings
6074            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        // For the vision encoder, before spatial merging
6099        let img_seq_len = {
6100            let cfg = &cfg.vision_config;
6101            // grid_t is 1 for images
6102            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            // If embeddings are tied and no packing, reuse weights -> no separate lm_head needed
6129            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            // Main merger: norm uses cfg.hidden_size
6146            let ln_q = cfg.hidden_size + bias_if!(true, cfg.hidden_size);
6147            let merger = mlp0 + mlp2 + ln_q;
6148
6149            // Deepstack mergers: norm uses merged hidden_size
6150            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            // Check if this is a MoE layer
6230            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                // MoE layer: gate + experts
6236                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                // Dense MLP layer
6248                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
6302// ======================== Qwen3_5 (Dense) Loader
6303
6304/// [`MultimodalLoader`] for a Qwen3.5 dense (hybrid GDN + full attention) model.
6305///
6306/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
6307pub struct Qwen3_5Loader;
6308
6309pub struct Qwen3_5Prefixer;
6310
6311impl MultimodalPromptPrefixer for Qwen3_5Prefixer {
6312    // No-op: With MessagesAction::Keep, the chat template handles image tokens
6313    // when it sees {"type": "image"} entries in the content.
6314}
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            // Full attention projections
6371            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            // GDN linear attention output projection
6376            Regex::new(
6377                r"model\.language_model\.layers\.(\d+)\.linear_attn\.out_proj\.(weight|bias)$",
6378            )?,
6379            // Dense MLP
6380            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                    // in_proj_qkvz: (2 * key_dim + 2 * value_dim, hidden)
6576                    let in_proj_qkvz = hidden * (key_dim * 2 + value_dim * 2);
6577                    // in_proj_ba: (2 * num_v_heads, hidden)
6578                    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                    // RmsNormGated over per-head value dim
6584                    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            // Dense MLP
6590            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
6637// ======================== Qwen3_5Moe Loader
6638
6639/// [`MultimodalLoader`] for a Qwen3.5 MoE (hybrid GDN + full attention) model.
6640///
6641/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
6642pub struct Qwen3_5MoeLoader;
6643
6644pub struct Qwen3_5MoePrefixer;
6645
6646impl MultimodalPromptPrefixer for Qwen3_5MoePrefixer {
6647    // No-op: With MessagesAction::Keep, the chat template handles image tokens
6648    // when it sees {"type": "image"} entries in the content.
6649}
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            // Full attention projections
6706            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            // GDN linear attention output projection
6711            Regex::new(
6712                r"model\.language_model\.layers\.(\d+)\.linear_attn\.out_proj\.(weight|bias)$",
6713            )?,
6714            // MoE experts
6715            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            // Shared expert
6725            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            // MoE experts
6743            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            // Shared expert
6753            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                    // in_proj_qkvz: (2 * key_dim + 2 * value_dim, hidden)
6955                    let in_proj_qkvz = hidden * (key_dim * 2 + value_dim * 2);
6956                    // in_proj_ba: (2 * num_v_heads, hidden)
6957                    let in_proj_ba = hidden * (text_cfg.linear_num_value_heads * 2);
6958                    // out_proj: value_dim -> hidden
6959                    let out_proj = value_dim * hidden / weight_pack_factor;
6960                    // conv1d weight
6961                    let conv1d = conv_dim * text_cfg.linear_conv_kernel_dim;
6962                    // dt_bias, A_log, norm weight
6963                    let dt_bias = text_cfg.linear_num_value_heads;
6964                    let a_log = text_cfg.linear_num_value_heads;
6965                    // RmsNormGated over per-head value dim
6966                    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            // All layers have MoE
6972            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
7032// ─── Voxtral ────────────────────────────────────────────────────────────────
7033
7034/// [`MultimodalLoader`] for a Voxtral model.
7035///
7036/// [`MultimodalLoader`]: https://docs.rs/mistralrs/latest/mistralrs/struct.MultimodalLoader.html
7037pub 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        // Mistral v7 instruct format using [INST]/[/INST] tokens
7099        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        // Mistral tekken tokenizer: <s> = ID 1, </s> = ID 2
7103        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            // Output / lm_head (tied with tok_embeddings)
7111            Regex::new(r"lm_head\.(weight|bias)$")?,
7112            // Decoder attention (Mistral-native naming)
7113            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            // Decoder MLP (Mistral-native naming)
7118            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            // Decoder attention
7127            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            // Decoder MLP
7132            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        // Audio tokens are prepended: max audio len + text seq len
7158        // Audio: ~30s at 16kHz = 480k samples, /160 hop = 3000 frames, /2 conv stride = 1500, /4 adapter = 375 tokens
7159        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        // Encoder max activation: attention matrix
7176        // ~3000 mel frames, encoder has 32 heads, seq_len^2
7177        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        // Encoder conv layers
7195        let conv1 = enc.dim * enc.audio_encoding_args.num_mel_bins * 3 + enc.dim; // weight + bias
7196        let conv2 = enc.dim * enc.dim * 3 + enc.dim;
7197
7198        // Encoder layers
7199        let enc_attn_per_layer = 4 * enc.dim * enc.dim; // wq, wk, wv, wo (full heads)
7200        let enc_mlp_per_layer = 3 * enc.dim * enc.hidden_dim; // w1, w2, w3
7201        let enc_norm_per_layer = 2 * enc.dim; // attention_norm, ffn_norm
7202        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        // Adapter
7207        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        // Decoder embeddings
7213        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; // attention_norm + ffn_norm
7236
7237        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
7266// ── Gemma4 ─────────────────────────────────────────────────────────────────
7267
7268pub 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        // `embed_vision.embedding_projection` is intentionally excluded.
7352        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}