@@ -95,6 +95,10 @@ class VisionConfig:
9595 norm_eps : float = 1e-6
9696 mm_tokens_per_image : int | None = None
9797 image_token_id : int | None = None
98+ # Pixtral / Mistral-3 vision fields
99+ model_type : str | None = None
100+ head_dim : int | None = None
101+ rope_theta : float | None = None
98102 # Qwen VL-specific
99103 out_hidden_size : int | None = None
100104 in_channels : int = 3
@@ -273,7 +277,13 @@ def _extract_rope_config(config) -> RoPEConfig:
273277 _nested_rope_theta (rope_scaling , "full_attention" ),
274278 default = 10_000.0 ,
275279 ),
276- rope_scaling = (_normalize_rope_scaling (rope_scaling ) or None ),
280+ # Some models (e.g. Ministral-3) store YaRN config under
281+ # rope_parameters instead of rope_scaling; fall back accordingly.
282+ rope_scaling = (
283+ _normalize_rope_scaling (rope_scaling )
284+ or _normalize_rope_scaling (rope_parameters )
285+ or None
286+ ),
277287 partial_rotary_factor = _first_not_none (
278288 getattr (config , "partial_rotary_factor" , None ),
279289 rope_scaling .get ("partial_rotary_factor" , None ),
@@ -284,12 +294,11 @@ def _extract_rope_config(config) -> RoPEConfig:
284294 getattr (config , "rope_local_base_freq" , None ),
285295 _nested_rope_theta (rope_scaling , "sliding_attention" ),
286296 ),
287- original_max_position_embeddings = (
288- getattr (
289- config ,
290- "original_max_position_embeddings" ,
291- rope_scaling .get ("original_max_position_embeddings" , None ),
292- )
297+ original_max_position_embeddings = _first_not_none (
298+ getattr (config , "original_max_position_embeddings" , None ),
299+ rope_scaling .get ("original_max_position_embeddings" , None ),
300+ # Also check rope_parameters (see rope_scaling comment above).
301+ rope_parameters .get ("original_max_position_embeddings" , None ),
293302 ),
294303 )
295304
@@ -351,10 +360,22 @@ def _extract_vision_config(config, parent_config, model_type: str) -> dict:
351360 ),
352361 image_size = getattr (vc , "image_size" , None ),
353362 patch_size = getattr (vc , "patch_size" , None ),
354- norm_eps = getattr (vc , "layer_norm_eps" , 1e-6 ),
363+ norm_eps = _first_not_none (
364+ getattr (vc , "layer_norm_eps" , None ),
365+ getattr (vc , "norm_eps" , None ),
366+ default = 1e-6 ,
367+ ),
368+ # Pixtral / Mistral-3 vision fields
369+ model_type = getattr (vc , "model_type" , None ),
370+ head_dim = getattr (vc , "head_dim" , None ),
371+ rope_theta = getattr (vc , "rope_theta" , None ),
355372 # Qwen VL-specific vision fields
356373 out_hidden_size = getattr (vc , "out_hidden_size" , None ),
357- in_channels = getattr (vc , "in_channels" , 3 ),
374+ in_channels = _first_not_none (
375+ getattr (vc , "in_channels" , None ),
376+ getattr (vc , "num_channels" , None ),
377+ default = 3 ,
378+ ),
358379 spatial_merge_size = getattr (vc , "spatial_merge_size" , 2 ),
359380 temporal_patch_size = getattr (vc , "temporal_patch_size" , 2 ),
360381 num_position_embeddings = getattr (vc , "num_position_embeddings" , None ),
@@ -573,6 +594,11 @@ def from_transformers(cls, hf_config) -> QuantizationConfig | None:
573594 method = qc .get ("quant_method" , "none" )
574595 if method == "none" :
575596 return None
597+ # FP8 per-tensor quantization (float8_e4m3fn + scalar scale)
598+ # is handled by dtype casting in _assign_weight(), not by
599+ # QuantizedLinear block quantization.
600+ if method == "fp8" :
601+ return None
576602 return cls (
577603 bits = qc .get ("bits" , 4 ),
578604 group_size = qc .get ("group_size" , 128 ),
0 commit comments