From d713c09294f287d378985a4cf11f2f77ecff604f Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Thu, 16 Jul 2026 21:12:52 +0900 Subject: [PATCH 1/4] feat(vlm): add MiniMax-M3-VL multimodal support Add MiniMax-M3-VL (model_type "minimax_m3_vl"), the vision-language variant on top of the merged MiniMax-M3 text backbone (#763). The only public checkpoint (MiniMaxAI/MiniMax-M3, 427B) is itself the VL model: a top-level minimax_m3_vl config nesting text_config plus vision_config. Vision tower (src/vision/encoders/minimax_m3_vl.rs): a CLIP-style ViT (hidden 1280, 16 heads, head_dim 80, 32 layers, intermediate 5120, patch 14) that structurally matches the Qwen2-VL vision tower but carries the checkpoint's real deltas: CLIP tensor naming under vision_tower.vision_model.*, a pre_layrnorm (the checkpoint's spelling) before the encoder, separate self_attn.{q,k,v,out}_proj with bias, and LayerNorm plus exact-GELU blocks. It reuses the shared Qwen2-VL vision helpers for the 2D (h, w) rotary embedding and per-image cu_seqlens variable-length attention, matching the image_grid_thw packing. A two-stage projector maps features to the text width: a per-patch multi_modal_projector (linear_1, GELU, linear_2 into projection_dim 6144) followed by a patch_merge_mlp that folds spatial_merge_size^2 = 4 adjacent patches (linear_1 [6144, 24576], GELU, linear_2). The tower runs in f32 (the vision weights are non-quantized in the checkpoint) to avoid mixed-dtype matmuls; the LLaVA-style merge casts features back to the text embedding dtype. Processor (src/vision/processors/minimax_m3.rs): a faithful port of the checkpoint image_processor.py, i.e. Qwen2-VL-style smart_resize (28-aligned, 672x672 max), CLIP mean/std normalization, temporal padding, and the exact merge-grouped patchify order, emitting pixel_values [num_patches, 1176] plus image_grid_thw. Model wrapper (src/vision/minimax_m3_vl.rs): MiniMaxM3VlModel composes the tower with the MiniMax-M3 hybrid dense/MoE backbone. Image placeholders (image 200025, vision_start 200029, vision_end 200030) expand via the shared Qwen-VL insertion helper (its t*(h/merge)*(w/merge) count and vision_start/+1 framing already match), the projected features replace them LLaVA-style, and the decoder runs its standard partial 1D RoPE. The MiniMax-M3 text model gains get_embed_tokens and forward_with_embeddings_impl for the merge path; its forward is refactored to share the embeddings path (non-regressing). Loader (src/loading/vlm_minimax_m3_vl.rs): parses the nested text_config/vision_config, builds the tower from the raw weight map, then hands the map to the MiniMax-M3 text sanitizer (which drops the vision/projector tensors and rewrites language_model.model.* to model.*). Vision weights are f32/bf16 while the text tower may be quantized in community exports; the text ModelArgs inherits the top-level quantization block when text_config omits it. Registration wires model_type detection, the ModelType/LoadedModel/VlmRuntimeRef variants and runtime dispatch arm, model_metadata, the TP arch string, docs/supported-models.md, and every exhaustive-match site (the arch-coverage macro test passes). Video is out of scope for this port (the acceptance criteria are image-based) and is declined cleanly. Validation: the 427B checkpoint exceeds the development machine, so real-image Q&A is deferred to a runtime validation once a fitting quantized conversion exists (as the issue anticipates). The merge gate is the unit tests plus the real nested-config parse. cargo test --release --features cuda --lib minimax_m3 passes 46 tests: 9 new (a tiny synthetic tower forward covering patch embed, cu_seqlens attention and 2D vision RoPE, both projector stages; projector fold ordering; processor grid math and 1176-dim rows; placeholder count; sanitizer/loader remap for verbatim tensor names including pre_layrnorm; and real vision/nested config parse) plus the 37 existing MiniMax-M3 text tests still green. Refs #764 --- docs/supported-models.md | 1 + src/distributed/tensor_parallel/inference.rs | 1 + src/loaded_model.rs | 2 + src/loaded_model_capabilities.rs | 3 + src/loading/mod.rs | 1 + src/loading/vlm.rs | 3 + src/loading/vlm_minimax_m3_vl.rs | 106 ++++ src/model_metadata.rs | 1 + src/models/detection.rs | 1 + src/models/minimax_m3.rs | 25 +- src/models/mod.rs | 7 + src/multimodal/vlm_runtime.rs | 30 + src/vision/encoders/minimax_m3_vl.rs | 633 +++++++++++++++++++ src/vision/encoders/mod.rs | 1 + src/vision/minimax_m3_vl.rs | 132 ++++ src/vision/minimax_m3_vl_tests.rs | 366 +++++++++++ src/vision/mod.rs | 2 + src/vision/processors/minimax_m3.rs | 287 +++++++++ src/vision/processors/mod.rs | 1 + 19 files changed, 1602 insertions(+), 1 deletion(-) create mode 100644 src/loading/vlm_minimax_m3_vl.rs create mode 100644 src/vision/encoders/minimax_m3_vl.rs create mode 100644 src/vision/minimax_m3_vl.rs create mode 100644 src/vision/minimax_m3_vl_tests.rs create mode 100644 src/vision/processors/minimax_m3.rs diff --git a/docs/supported-models.md b/docs/supported-models.md index e58138cc8..ee800802e 100644 --- a/docs/supported-models.md +++ b/docs/supported-models.md @@ -95,6 +95,7 @@ Implemented VLM variants include: - ERNIE-4.5 MoE VL (`ernie4_5_moe_vl`): Baidu's vision-language MoE. A DFNRope ViT (linear patch embedding over 588-wide merge-window rows, 2D vision RoPE, `cu_seqlens`-packed attention, quick_gelu MLP) feeds a variable-resolution resampler (2x2 spatial fold, temporal pair fold with single-frame duplication, GELU MLP stacks, RMSNorm) whose rows replace the `<|IMAGE_PLACEHOLDER|>` tokens. The text decoder extends ERNIE-4.5 MoE with modality-split expert banks: separate text and multimodal routers and expert stacks selected per token by token type, with a correction bias that shifts expert selection but never the mixing weights, plus a fused shared-experts MLP. Position encoding is interleaved 3D MRoPE (`[T, H, W]` axes assigned per frequency index, adjacent-pair rotation) that degenerates to traditional RoPE for text. Validated against `mlx-community/ERNIE-4.5-VL-28B-A3B-Thinking-4bit` (28B-A3B; text and resampler 4-bit, vision tower bf16). Best for general image chat and grounded reasoning; the checkpoint is a thinking variant and emits reasoning before the answer. - Qwen3-Omni MoE (`qwen3_omni_moe`, thinker): Alibaba's omni-modal MoE. Stage 1 covers the thinker: text output conditioned on text, image, and audio inputs. The vision tower and MoE text decoder are the Qwen3-VL-MoE stack (DeepStack feature injection, interleaved MRoPE) reused unchanged; the new audio tower converts 16 kHz audio to a 128-bin log-mel spectrogram, downsamples it through three stride-2 convolutions (13 output frames per second of audio), and runs 32 windowed-attention encoder layers whose output rows scatter into the token stream exactly like image features. Audio arrives via `--audio file.wav` on the CLI (combinable with `--image`). Stage 2 adds speech output: `mlxcel generate --output-audio out.wav` runs the talker and code2wav after text generation and writes 24 kHz mono PCM16. The talker is a 20-layer Qwen3-MoE codec decoder conditioned on the projected thinker token embeddings of the chat-role segments; per frame it emits the first of 16 codebooks and a 5-layer code predictor fills in the residual 15, then the code2wav vocoder (causal pre-transformer, ConvNeXt upsampling, BigVGAN-style SnakeBeta decoder) renders 1920 samples per 12.5 Hz frame. `--speaker` selects the voice (ethan default; chelsie and aiden also ship in the released checkpoints). The speech stack loads lazily and only when requested, so text/vision use keeps its memory footprint; speech currently requires a text-only, chat-templated prompt. Validated against `mlx-community/Qwen3-Omni-30B-A3B-Instruct-4bit` (text and talker 4-bit; vision, audio tower, code predictor, and code2wav bf16). - Hunyuan-VL (`hunyuan_vl`, e.g. HunyuanOCR): Tencent's vision-language family. A ViT with a per-patch conv embedding, bilinearly interpolated learned position embeddings, and full attention over the packed patch sequence feeds a `perceive` merger: a stride-2 conv pair over the raster grid, a learned `image_newline` column, a linear to the decoder width, and learned `image_begin` / `image_end` rows. Per image that yields `mh * (mw + 1) + 2` feature rows, matching the prompt placeholder count exactly. The decoder is the Hunyuan dense stack (per-head Q/K RMSNorm after the rotation, DynamicNTK-alpha rope base) with XD-RoPE at prefill: 4D `[P, T, H, W]` position ids split across the frequency dims, degenerating to the standard rotation for text; decode uses sequential positions. Validated against `hadeseus/HunyuanOCR-mlx-4bit` (text 4-bit, vision bf16). Best for OCR: text spotting, document parsing, and grounded extraction. +- MiniMax-M3-VL (`minimax_m3_vl`): MiniMax's vision-language model on top of the MiniMax-M3 text backbone. A CLIP-style ViT (hidden 1280, 16 heads, 32 layers, patch 14, `pre_layrnorm`, LayerNorm + exact-GELU blocks with separate `q/k/v/out` projections) runs native-resolution packing: Qwen2-VL-style dynamic `smart_resize` to `patch_size * spatial_merge_size = 28`-aligned dimensions, `image_grid_thw` patchify, per-image `cu_seqlens` variable-length attention, and 2D (h, w) vision RoPE. A two-stage projector then maps features to the text width: a per-patch `multi_modal_projector` (`linear_1 -> GELU -> linear_2`, into `projection_dim` 6144) followed by a `patch_merge_mlp` that folds each `spatial_merge_size^2 = 4` adjacent patches (`linear_1` [6144, 24576] `-> GELU -> linear_2`). Each `]<]image[>[` placeholder expands to `grid_t * (h/2) * (w/2)` tokens (`vision_start` 200029, `vision_end` 200030), which the merged features replace LLaVA-style; the MiniMax-M3 hybrid dense/MoE decoder then runs its standard partial 1D RoPE. The vision tower runs in f32 (non-quantized in the checkpoint) while the text tower may be quantized in community exports. Image and multi-image inputs are wired end to end (CLI and server); video is out of scope for this port. The only public checkpoint (`MiniMaxAI/MiniMax-M3`, 427B) exceeds the development machine, so image Q&A parity is deferred to a runtime validation once a fitting quantized conversion exists; the merge gate is the unit tests plus the real nested-config parse. - FastVLM (`llava_qwen2` / `fastvlm`): Apple's low-latency VLM. A FastViTHD hybrid encoder runs entirely on channels-last maps: a conv stem, three RepMixer stages (depthwise token mixing plus a BatchNorm ConvFFN), two attention stages with channel LayerNorm and `head_dim` 32 self-attention, inter-stage large-kernel PatchEmbed downsamples, RepCPE position encoders, and a squeeze-excite `conv_exp` head. Each 1024x1024 pad-to-square image becomes a `(16, 16, 3072)` map flattened to 256 tokens, projected to the Qwen2 decoder width by an `mlp2x_gelu` MLP. The `` placeholder is the fixed `-200` sentinel (not a vocabulary token); the runtime splices one sentinel per image, expands it to 256 tokens, and scatters the image embeddings (LLaVA merge). The text decoder is stock Qwen2, reused unchanged. The loader accepts both the genuine (`apple/FastVLM-0.5B`) and converted (`mlx-community/FastVLM-0.5B-bf16`) weight layouts. Best for fast image description and grounded chat. - GLM-OCR (`glm_ocr`): document-OCR sibling of GLM-4V. A 24-block ViT (3D patch embedding, per-head q/k RMSNorm on the packed `cu_seqlens` attention, 2D vision RoPE, Conv2d spatial downsample, SwiGLU patch merger) feeds a 16-layer GLM-4 text decoder driven by full-width even/odd MRoPE (`rope_parameters` with `mrope_section [16, 24, 24]`, `partial_rotary_factor 1.0`). The tower has no learned position embedding or post-conv norm, and the loader drops the next-n prediction (MTP) layer. Patches are reordered from the processor's raster order into spatial-merge-window order so the rotary, downsample, and merged-token grid stay spatially aligned (OCR reads scrambled patches wrong). Best for plain text, tables, and formula recognition. - Youtu-VL diff --git a/src/distributed/tensor_parallel/inference.rs b/src/distributed/tensor_parallel/inference.rs index b06c81976..752b15c99 100644 --- a/src/distributed/tensor_parallel/inference.rs +++ b/src/distributed/tensor_parallel/inference.rs @@ -171,6 +171,7 @@ fn fallback_architecture(model_type: ModelType) -> &'static str { ModelType::GptOss => "gpt_oss", ModelType::MiniMax => "minimax", ModelType::MiniMaxM3 => "minimax_m3", + ModelType::MiniMaxM3VL => "minimax_m3", ModelType::Mixtral => "mixtral", ModelType::Qwen2Moe => "qwen2_moe", ModelType::OLMoE => "olmoe", diff --git a/src/loaded_model.rs b/src/loaded_model.rs index f9e94722f..47907c314 100644 --- a/src/loaded_model.rs +++ b/src/loaded_model.rs @@ -84,6 +84,7 @@ pub enum LoadedModel { FastVLM(vision::VisionLanguageModel), Ernie45MoeVLM(vision::ernie4_5_moe_vl::Ernie45MoeVlModel), HunyuanVLM(vision::hunyuan_vl::HunyuanVlModel), + MiniMaxM3VL(vision::MiniMaxM3VlModel), GraniteVisionVLM(vision::GraniteVisionVLModel), Granite4VisionVLM(vision::Granite4VisionVLModel), DeepSeekOcrVLM(vision::deepseekocr::DeepSeekOcrVlModel), @@ -229,6 +230,7 @@ macro_rules! delegate_language_model { LoadedModel::FastVLM(inner) => LanguageModel::$method(inner, $($arg),*), LoadedModel::Ernie45MoeVLM(inner) => LanguageModel::$method(inner, $($arg),*), LoadedModel::HunyuanVLM(inner) => LanguageModel::$method(inner, $($arg),*), + LoadedModel::MiniMaxM3VL(inner) => LanguageModel::$method(inner, $($arg),*), LoadedModel::GraniteVisionVLM(inner) => LanguageModel::$method(inner, $($arg),*), LoadedModel::Granite4VisionVLM(inner) => LanguageModel::$method(inner, $($arg),*), LoadedModel::DeepSeekOcrVLM(inner) => LanguageModel::$method(inner, $($arg),*), diff --git a/src/loaded_model_capabilities.rs b/src/loaded_model_capabilities.rs index cffc1d8fe..61812317c 100644 --- a/src/loaded_model_capabilities.rs +++ b/src/loaded_model_capabilities.rs @@ -72,6 +72,8 @@ pub enum VlmRuntimeRef<'a> { Ernie45MoeVl(&'a vision::ernie4_5_moe_vl::Ernie45MoeVlModel), Qwen3OmniMoe(&'a vision::qwen3_omni_moe::Qwen3OmniMoeModel), HunyuanVl(&'a vision::hunyuan_vl::HunyuanVlModel), + /// MiniMax-M3-VL runtime (CLIP ViT + MiniMax-M3 hybrid dense/MoE text). + MiniMaxM3Vl(&'a vision::MiniMaxM3VlModel), /// Pixtral / Mistral3 dynamic aspect-ratio runtime. Shares the generic /// `VisionModule` storage but preserves each image's aspect ratio and emits /// `[IMG] / [IMG_BREAK] / [IMG_END]` row structure (see `pixtral_layout`). @@ -206,6 +208,7 @@ impl LoadedModel { Self::FastVLM(vlm) => Some(VlmRuntimeRef::FastVLM(&vlm.vision)), Self::Ernie45MoeVLM(model) => Some(VlmRuntimeRef::Ernie45MoeVl(model)), Self::HunyuanVLM(model) => Some(VlmRuntimeRef::HunyuanVl(model)), + Self::MiniMaxM3VL(model) => Some(VlmRuntimeRef::MiniMaxM3Vl(model)), _ => None, } } diff --git a/src/loading/mod.rs b/src/loading/mod.rs index 61616c24c..595b929e5 100644 --- a/src/loading/mod.rs +++ b/src/loading/mod.rs @@ -186,6 +186,7 @@ fn try_load_vlm_model_from_dir( ModelType::FastVLM => Some(load_fastvlm_vlm(model_path)?), ModelType::Ernie45MoeVLM => Some(load_ernie4_5_moe_vlm(model_path)?), ModelType::HunyuanVLM => Some(load_hunyuan_vlm(model_path)?), + ModelType::MiniMaxM3VL => Some(load_minimax_m3_vl(model_path)?), ModelType::LlavaBunnyVLM => Some(load_llava_bunny_vlm(model_path)?), ModelType::AyaVisionVLM => Some(load_aya_vision_vlm(model_path)?), ModelType::PaliGemmaVLM => Some(load_paligemma_vlm(model_path)?), diff --git a/src/loading/vlm.rs b/src/loading/vlm.rs index 9821d74fb..0c57c60f8 100644 --- a/src/loading/vlm.rs +++ b/src/loading/vlm.rs @@ -67,6 +67,8 @@ mod kimi_vl_loader; mod lfm2_vl; #[path = "vlm_llava.rs"] mod llava; +#[path = "vlm_minimax_m3_vl.rs"] +mod minimax_m3_vl; #[path = "vlm_mllama.rs"] mod mllama; #[path = "vlm_nemotron_h_nano_omni.rs"] @@ -103,6 +105,7 @@ pub(crate) use internvl::load_internvl_vlm; pub(crate) use kimi_vl_loader::load_kimi_vl_vlm; pub(crate) use lfm2_vl::load_lfm2_vl; pub(crate) use llava::{load_llava_bunny_vlm, load_llava_vlm}; +pub(crate) use minimax_m3_vl::load_minimax_m3_vl; pub(crate) use mllama::load_mllama_vlm; pub(crate) use nemotron_h_nano_omni::load_nemotron_h_nano_omni_vlm; pub(crate) use paddleocr::load_paddleocr_vl; diff --git a/src/loading/vlm_minimax_m3_vl.rs b/src/loading/vlm_minimax_m3_vl.rs new file mode 100644 index 000000000..7fb8ef010 --- /dev/null +++ b/src/loading/vlm_minimax_m3_vl.rs @@ -0,0 +1,106 @@ +// Copyright 2025-2026 Lablup Inc. and Jeongkyu Shin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! MiniMax-M3-VL loader (`model_type: "minimax_m3_vl"`). +//! +//! The only public checkpoint (`MiniMaxAI/MiniMax-M3`, 427B) is itself the VL +//! model: a top-level `minimax_m3_vl` config with nested `text_config` and +//! `vision_config`. Text weights live under `language_model.model.*` and the +//! CLIP-style vision tower under `vision_tower.vision_model.*`, with the +//! two-stage projector at `multi_modal_projector.*` / `patch_merge_mlp.*`. +//! +//! The tower is built from the raw weight map first (its verbatim keys are +//! intact there), after which the map is moved into the MiniMax-M3 text +//! sanitizer, which drops the vision/projector tensors and rewrites +//! `language_model.model.*` -> `model.*` for the text decoder. Vision weights +//! are f32/bf16 non-quantized while the text tower may be quantized in +//! community exports; the text `ModelArgs` inherits the top-level quantization +//! block when its own `text_config` omits it. + +use anyhow::Result; +use std::path::Path; + +use crate::LoadedModel; +use crate::models::MiniMaxM3Model; +use crate::models::minimax_m3::{ModelArgs, sanitize_weights}; +use crate::vision; +use crate::vision::encoders::minimax_m3_vl::{MiniMaxM3VisionConfig, MiniMaxM3VisionEncoder}; +use crate::vision::processors::minimax_m3::MiniMaxM3Processor; + +use super::{load_vlm_weights_common, read_sanitized_vlm_config}; + +pub(crate) fn load_minimax_m3_vl(model_path: &Path) -> Result { + let (_config_str, full_config) = read_sanitized_vlm_config(model_path)?; + + // Text config lives under `text_config`. Community 4-bit exports commonly + // keep the quantization block at the config root, so inject it when the + // nested block omits it. + let mut text_value = full_config + .get("text_config") + .cloned() + .ok_or_else(|| anyhow::anyhow!("Missing text_config in MiniMax-M3-VL config.json"))?; + if let (Some(obj), Some(quant)) = (text_value.as_object_mut(), full_config.get("quantization")) + { + obj.entry("quantization").or_insert_with(|| quant.clone()); + } + let text_args: ModelArgs = serde_json::from_value(text_value) + .map_err(|e| anyhow::anyhow!("Failed to parse MiniMax-M3-VL text_config: {}", e))?; + if let Some(sparse) = text_args.sparse_attention_config.as_ref() { + sparse.validate().map_err(|e| anyhow::anyhow!(e))?; + } + + let vision_value = full_config + .get("vision_config") + .cloned() + .ok_or_else(|| anyhow::anyhow!("Missing vision_config in MiniMax-M3-VL config.json"))?; + let vision_config: MiniMaxM3VisionConfig = serde_json::from_value(vision_value) + .map_err(|e| anyhow::anyhow!("Failed to parse MiniMax-M3-VL vision_config: {}", e))?; + + let weights = load_vlm_weights_common(model_path, None)?; + + // Build the vision tower from the raw keys first (borrowing), then move the + // map into the text sanitizer. + let vision_encoder = MiniMaxM3VisionEncoder::from_weights(&weights, &vision_config) + .map_err(|e| anyhow::anyhow!("Failed to load MiniMax-M3-VL vision tower: {}", e))?; + + let text_weights = sanitize_weights(weights, &text_args); + let text_model = MiniMaxM3Model::from_weights(&text_weights, &text_args) + .map_err(|e| anyhow::anyhow!("Failed to load MiniMax-M3-VL text model: {}", e))?; + + let read_token = |key: &str, default: i64| -> i32 { + full_config + .get(key) + .and_then(|v| v.as_i64()) + .unwrap_or(default) as i32 + }; + + let mut eos_token_ids = crate::loading::read_eos_token_ids(model_path); + if eos_token_ids.is_empty() { + // MiniMax `[e~[` sentinel at 200020 in its 200k vocab. + eos_token_ids = vec![200020]; + } + + let vlm = vision::minimax_m3_vl::MiniMaxM3VlModel { + text_model, + vision_encoder, + processor: MiniMaxM3Processor::default(), + image_token_id: read_token("image_token_id", 200025), + video_token_id: read_token("video_token_id", 200026), + vision_start_token_id: read_token("vision_start_token_id", 200029), + vision_end_token_id: read_token("vision_end_token_id", 200030), + spatial_merge_size: vision_config.spatial_merge_size() as i32, + eos_token_ids, + }; + Ok(LoadedModel::MiniMaxM3VL(vlm)) +} diff --git a/src/model_metadata.rs b/src/model_metadata.rs index aee638cd0..ff5a4a259 100644 --- a/src/model_metadata.rs +++ b/src/model_metadata.rs @@ -165,6 +165,7 @@ macro_rules! for_each_model_registration { PhiMoe => { kind: Text, directory: ConfigBacked, weight: Some(WeightLoadRoute::ConfigBacked), adapter: None, config_backed: { dir_loader: models::PhiMoeModel::load, args: models::phimoe::ModelArgs, weight_builder: models::PhiMoeModel::from_weights, wrap: LoadedModel::PhiMoe } }; MiniMax => { kind: Text, directory: ConfigBacked, weight: Some(WeightLoadRoute::ConfigBacked), adapter: None, config_backed: { dir_loader: models::MiniMaxModel::load, args: models::minimax::ModelArgs, weight_builder: models::MiniMaxModel::from_weights, wrap: LoadedModel::MiniMax } }; MiniMaxM3 => { kind: Text, directory: ConfigBacked, weight: Some(WeightLoadRoute::ConfigBacked), adapter: None, config_backed: { dir_loader: models::MiniMaxM3Model::load, args: models::minimax_m3::ModelArgs, weight_builder: models::MiniMaxM3Model::from_weights, wrap: LoadedModel::MiniMaxM3 } }; + MiniMaxM3VL => { kind: Vlm, directory: Vlm, weight: None, adapter: Some("MiniMax-M3-VL cannot be loaded with LoRA adapters yet") }; Mixtral => { kind: Text, directory: ConfigBacked, weight: Some(WeightLoadRoute::ConfigBacked), adapter: None, config_backed: { dir_loader: models::MixtralModel::load, args: models::mixtral::ModelArgs, weight_builder: models::MixtralModel::from_weights, wrap: LoadedModel::Mixtral } }; OLMoE => { kind: Text, directory: ConfigBacked, weight: Some(WeightLoadRoute::ConfigBacked), adapter: None, config_backed: { dir_loader: models::OlmoeModel::load, args: models::olmoe::ModelArgs, weight_builder: models::OlmoeModel::from_weights, wrap: LoadedModel::OLMoE } }; DeepSeek => { kind: Text, directory: ConfigBacked, weight: Some(WeightLoadRoute::ConfigBacked), adapter: None, config_backed: { dir_loader: models::DeepSeekModel::load, args: models::deepseek::ModelArgs, weight_builder: models::DeepSeekModel::from_weights, wrap: LoadedModel::DeepSeek } }; diff --git a/src/models/detection.rs b/src/models/detection.rs index dbd8e880a..3f9a1a1d5 100644 --- a/src/models/detection.rs +++ b/src/models/detection.rs @@ -160,6 +160,7 @@ pub fn get_model_type(model_path: &Path) -> Result { "phimoe" => Ok(ModelType::PhiMoe), "minimax" => Ok(ModelType::MiniMax), "minimax_m3" => Ok(ModelType::MiniMaxM3), + "minimax_m3_vl" => Ok(ModelType::MiniMaxM3VL), "gpt_oss" => Ok(ModelType::GptOss), "mixtral" => Ok(ModelType::Mixtral), "olmoe" => Ok(ModelType::OLMoE), diff --git a/src/models/minimax_m3.rs b/src/models/minimax_m3.rs index 2a585a703..c7c871597 100644 --- a/src/models/minimax_m3.rs +++ b/src/models/minimax_m3.rs @@ -167,7 +167,30 @@ impl MiniMaxM3Model { caches: &mut [KVCache], mask: Option<&MlxArray>, ) -> UniquePtr { - let mut h = self.embed_tokens.forward(input_ids); + self.forward_with_embeddings_impl(input_ids, None, caches, mask) + } + + /// Token embedding lookup, exposed for the VL wrapper's LLaVA-style merge + /// (`src/vision/minimax_m3_vl.rs`). + pub fn get_embed_tokens(&self, input_ids: &MlxArray) -> UniquePtr { + self.embed_tokens.forward(input_ids) + } + + /// Decoder forward that optionally starts from precomputed input + /// embeddings. When `input_embeddings` is `Some`, the embedding lookup is + /// skipped and the provided (already vision-merged) embeddings are decoded + /// directly; `input_ids` is then only used for shape/consistency by callers. + pub fn forward_with_embeddings_impl( + &self, + input_ids: &MlxArray, + input_embeddings: Option<&MlxArray>, + caches: &mut [KVCache], + mask: Option<&MlxArray>, + ) -> UniquePtr { + let mut h = match input_embeddings { + Some(embeds) => mlxcel_core::copy(embeds), + None => self.embed_tokens.forward(input_ids), + }; for (i, layer) in self.layers.iter().enumerate() { h = layer.forward(&h, &mut caches[i], mask); diff --git a/src/models/mod.rs b/src/models/mod.rs index c7a5355b7..00866f360 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -312,6 +312,7 @@ pub enum ModelType { GptOss, MiniMax, MiniMaxM3, + MiniMaxM3VL, // MiniMax-M3-VL (CLIP ViT + M3 hybrid dense/MoE text) Mixtral, Qwen2Moe, OLMoE, @@ -506,6 +507,7 @@ pub const ALL_MODEL_TYPES: &[ModelType] = &[ ModelType::GptOss, ModelType::MiniMax, ModelType::MiniMaxM3, + ModelType::MiniMaxM3VL, ModelType::Mixtral, ModelType::Qwen2Moe, ModelType::OLMoE, @@ -760,6 +762,10 @@ impl ModelType { "MiniMax-M3 (hybrid dense/MoE, block-sparse attention)", "MoE (other)", ), + ModelType::MiniMaxM3VL => ( + "MiniMax-M3-VL (CLIP ViT + M3 hybrid dense/MoE)", + "MiniMax VLM", + ), ModelType::Mixtral => ("Mixtral (MoE)", "MoE (other)"), ModelType::KimiLinear => ("Kimi Linear (MLA + GatedDeltaNet hybrid)", "MoE (other)"), ModelType::KimiVL => ("Kimi-VL (MoonViT + DeepSeek-V3 MoE)", "Kimi VLM"), @@ -971,6 +977,7 @@ mod metadata_tests { GptOss, MiniMax, MiniMaxM3, + MiniMaxM3VL, Mixtral, Qwen2Moe, OLMoE, diff --git a/src/multimodal/vlm_runtime.rs b/src/multimodal/vlm_runtime.rs index d83c016bd..3178218c8 100644 --- a/src/multimodal/vlm_runtime.rs +++ b/src/multimodal/vlm_runtime.rs @@ -487,6 +487,36 @@ where preparation, })) } + VlmRuntimeRef::MiniMaxM3Vl(model) => { + // Native-resolution CLIP tower: preprocess to (pixel_values, + // grid_thw), expand image placeholders with the shared Qwen-VL + // helper (its `t*(h/merge)*(w/merge)` count and vision_start/+1 + // framing already match MiniMax-M3-VL: start 200029, end 200030), + // then scatter the projected features LLaVA-style. + let (pixel_values, grid_thw) = model.processor.preprocess_with_grid(images); + let preparation = insert_qwen_vl_image_tokens( + prompt_tokens, + &grid_thw, + model.spatial_merge_size as usize, + model.vision_start_token_id, + model.image_token_id, + ) + .map(|stats| VlmPreparationSummary::QwenVlm { + image_blocks: stats.image_blocks, + total_image_tokens: stats.total_image_tokens, + }); + + let _ = active_caches; + let _ = image_cache_keys; + + let input_ids_arr = prompt_ids_array(prompt_tokens); + let embeddings = model.input_embeddings(&input_ids_arr, &pixel_values, &grid_thw); + + Ok(Some(PreparedVlmEmbeddings { + embeddings, + preparation, + })) + } VlmRuntimeRef::MiniCPMO(minicpmo) => { let prepared = prepare_minicpmo_prompt_tokens( prompt, diff --git a/src/vision/encoders/minimax_m3_vl.rs b/src/vision/encoders/minimax_m3_vl.rs new file mode 100644 index 000000000..81b643f86 --- /dev/null +++ b/src/vision/encoders/minimax_m3_vl.rs @@ -0,0 +1,633 @@ +// Copyright 2025-2026 Lablup Inc. and Jeongkyu Shin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! MiniMax-M3-VL vision tower (`model_type: "minimax_m3_vl"`, nested +//! `vision_config.model_type: "clip_vision_model"`). +//! +//! A CLIP-style ViT with native-resolution packing, structurally the Qwen2-VL +//! vision tower (hidden 1280, 16 heads, head_dim 80, 32 layers, intermediate +//! 5120, patch 14) but with the differences that live in the real checkpoint: +//! - CLIP tensor naming (`vision_tower.vision_model.encoder.layers.N.*`) with a +//! leading `pre_layrnorm` (the checkpoint's spelling) before the encoder, +//! - separate `self_attn.{q,k,v,out}_proj` (each with bias) instead of a fused +//! `qkv`, +//! - LayerNorm + exact-GELU MLP (`layer_norm1/2`, `mlp.fc1/fc2`), +//! - a two-stage projector: a per-patch `multi_modal_projector` +//! (`linear_1` -> GELU -> `linear_2`, into `projection_dim` 6144) followed by +//! a `patch_merge_mlp` that folds `spatial_merge_size^2 = 4` adjacent patches +//! (`linear_1` [6144, 24576] -> GELU -> `linear_2`) into the text hidden size. +//! +//! The tower reuses the shared Qwen2-VL vision helpers for the 2D (h, w) rotary +//! embedding and per-image `cu_seqlens` variable-length attention, matching the +//! `image_grid_thw` packing emitted by the processor. Video (`grid_t > 1`) is +//! out of scope for this port; the temporal axis of the rope reduces to the +//! well-tested (h, w) form for images (`grid_t == 1`). +//! +//! The whole tower runs in f32: the checkpoint stores the vision weights as +//! f32/bf16 non-quantized, and running the tower uniformly in f32 avoids +//! mixed-dtype matmuls while the projector output is cast back to the text +//! embedding dtype by the LLaVA-style merge. The 427B checkpoint cannot be +//! loaded on the development machine, so the validated surface is the synthetic +//! reduced-config unit tests plus the real-config parse test. + +use super::VisionEncoderOutput; +use super::qwen2_vl::{VisionRotaryEmbedding, apply_rotary_pos_emb_vision, concat_many}; +use mlxcel_core::layers::LayerNorm; +use mlxcel_core::weights::WeightMap; +use mlxcel_core::{MlxArray, UniquePtr}; +use serde::Deserialize; + +/// `img_token_compression_config` block: how the projector folds patches. +#[derive(Debug, Clone, Deserialize)] +pub struct ImgTokenCompressionConfig { + #[serde(default = "default_compression_method")] + pub image_token_compression_method: String, + #[serde(default = "default_spatial_merge_size")] + pub spatial_merge_size: usize, + #[serde(default = "default_temporal_patch_size")] + pub temporal_patch_size: usize, +} + +impl Default for ImgTokenCompressionConfig { + fn default() -> Self { + Self { + image_token_compression_method: default_compression_method(), + spatial_merge_size: default_spatial_merge_size(), + temporal_patch_size: default_temporal_patch_size(), + } + } +} + +/// MiniMax-M3-VL vision encoder configuration (nested `vision_config`). +/// +/// The real checkpoint also ships LLaVA-style keys (`image_grid_pinpoints`, +/// `vision_feature_layer`, `vision_feature_select_strategy`, `image_seq_length`) +/// that are vestigial for image processing; they are ignored here (serde drops +/// unknown fields) so the config parses permissively without building logic on +/// them. +#[derive(Debug, Clone, Deserialize)] +pub struct MiniMaxM3VisionConfig { + #[serde(default = "default_hidden_size")] + pub hidden_size: usize, + #[serde(default = "default_num_attention_heads")] + pub num_attention_heads: usize, + #[serde(default = "default_num_hidden_layers")] + pub num_hidden_layers: usize, + #[serde(default = "default_intermediate_size")] + pub intermediate_size: usize, + #[serde(default = "default_patch_size")] + pub patch_size: usize, + #[serde(default = "default_projection_dim")] + pub projection_dim: usize, + #[serde(default = "default_rope_theta")] + pub rope_theta: f32, + #[serde(default = "default_layer_norm_eps")] + pub layer_norm_eps: f32, + #[serde(alias = "num_channels", default = "default_in_channels")] + pub in_channels: usize, + #[serde(default)] + pub img_token_compression_config: ImgTokenCompressionConfig, +} + +fn default_hidden_size() -> usize { + 1280 +} +fn default_num_attention_heads() -> usize { + 16 +} +fn default_num_hidden_layers() -> usize { + 32 +} +fn default_intermediate_size() -> usize { + 5120 +} +fn default_patch_size() -> usize { + 14 +} +fn default_projection_dim() -> usize { + 6144 +} +fn default_rope_theta() -> f32 { + 10000.0 +} +fn default_layer_norm_eps() -> f32 { + 1e-5 +} +fn default_in_channels() -> usize { + 3 +} +fn default_compression_method() -> String { + "patch_merge".to_string() +} +fn default_spatial_merge_size() -> usize { + 2 +} +fn default_temporal_patch_size() -> usize { + 2 +} + +impl MiniMaxM3VisionConfig { + pub fn spatial_merge_size(&self) -> usize { + self.img_token_compression_config.spatial_merge_size + } + + pub fn temporal_patch_size(&self) -> usize { + self.img_token_compression_config.temporal_patch_size + } + + pub fn head_dim(&self) -> usize { + self.hidden_size / self.num_attention_heads + } +} + +// ============================================================================ +// Plain f32 linear / layernorm helpers +// ============================================================================ + +fn load_f32(weights: &WeightMap, key: &str) -> Result, String> { + weights + .get(key) + .map(|w| mlxcel_core::astype(w, mlxcel_core::dtype::FLOAT32)) + .ok_or_else(|| format!("Weight not found: {}", key)) +} + +fn load_f32_opt(weights: &WeightMap, key: &str) -> Option> { + weights + .get(key) + .map(|w| mlxcel_core::astype(w, mlxcel_core::dtype::FLOAT32)) +} + +/// Non-quantized f32 linear (`y = x @ W^T + b`). The vision tower is not +/// quantized in the checkpoint, so a plain matmul keeps the whole tower in a +/// single dtype and avoids the quantized-linear machinery. +struct VisionLinear { + weight: UniquePtr, + bias: Option>, +} + +impl VisionLinear { + fn load(weights: &WeightMap, prefix: &str) -> Result { + Ok(Self { + weight: load_f32(weights, &format!("{}.weight", prefix))?, + bias: load_f32_opt(weights, &format!("{}.bias", prefix)), + }) + } + + fn forward(&self, x: &MlxArray) -> UniquePtr { + let wt = mlxcel_core::transpose(&self.weight); + let y = mlxcel_core::matmul(x, &wt); + match &self.bias { + Some(b) => mlxcel_core::add(&y, b), + None => y, + } + } +} + +fn load_layer_norm(weights: &WeightMap, prefix: &str, eps: f32) -> Result { + let weight = load_f32(weights, &format!("{}.weight", prefix))?; + let bias = load_f32_opt(weights, &format!("{}.bias", prefix)); + Ok(LayerNorm::new(weight, bias, eps)) +} + +// ============================================================================ +// Patch embedding +// ============================================================================ + +/// Temporal 3D patch conv degenerated to a linear. +/// +/// The checkpoint weight is `[out, in_channels, temporal, patch_h, patch_w]` +/// (PyTorch Conv3d layout). The processor flattens each patch row in the exact +/// same `[channel, temporal, patch_h, patch_w]` order, so a row-major reshape +/// to `[out, in_features]` aligns the two without any axis permutation. +struct PatchEmbed { + proj_weight: UniquePtr, + proj_bias: Option>, +} + +impl PatchEmbed { + fn from_weights( + weights: &WeightMap, + config: &MiniMaxM3VisionConfig, + prefix: &str, + ) -> Result { + let key = format!("{}.weight", prefix); + let w = load_f32(weights, &key)?; + let out_features = config.hidden_size as i32; + let in_features = (config.in_channels + * config.temporal_patch_size() + * config.patch_size + * config.patch_size) as i32; + let shape = mlxcel_core::array_shape(&w); + let proj_weight = if shape.len() == 2 { + w + } else { + mlxcel_core::reshape(&w, &[out_features, in_features]) + }; + Ok(Self { + proj_weight, + proj_bias: load_f32_opt(weights, &format!("{}.bias", prefix)), + }) + } + + /// `hidden_states`: `[num_patches, in_features]` -> `[num_patches, hidden]`. + fn forward(&self, hidden_states: &MlxArray) -> UniquePtr { + let wt = mlxcel_core::transpose(&self.proj_weight); + let y = mlxcel_core::matmul(hidden_states, &wt); + match &self.proj_bias { + Some(b) => mlxcel_core::add(&y, b), + None => y, + } + } +} + +// ============================================================================ +// Encoder layer (CLIP style) +// ============================================================================ + +struct VisionAttention { + q_proj: VisionLinear, + k_proj: VisionLinear, + v_proj: VisionLinear, + out_proj: VisionLinear, + num_heads: i32, + head_dim: i32, + scale: f32, +} + +impl VisionAttention { + fn from_weights( + weights: &WeightMap, + config: &MiniMaxM3VisionConfig, + prefix: &str, + ) -> Result { + let head_dim = config.head_dim() as i32; + Ok(Self { + q_proj: VisionLinear::load(weights, &format!("{}.q_proj", prefix))?, + k_proj: VisionLinear::load(weights, &format!("{}.k_proj", prefix))?, + v_proj: VisionLinear::load(weights, &format!("{}.v_proj", prefix))?, + out_proj: VisionLinear::load(weights, &format!("{}.out_proj", prefix))?, + num_heads: config.num_attention_heads as i32, + head_dim, + scale: (head_dim as f32).powf(-0.5), + }) + } + + /// Packed variable-length attention: `x` is `[total_tokens, hidden]`, + /// `cu_seqlens` marks the per-image segment boundaries, and full attention + /// runs independently within each segment. + fn forward( + &self, + x: &MlxArray, + cu_seqlens: &[i32], + rotary_pos_emb: &MlxArray, + ) -> UniquePtr { + let shape = mlxcel_core::array_shape(x); + let seq_length = shape[0]; + + let reshape_heads = |proj: UniquePtr| { + mlxcel_core::reshape(&proj, &[seq_length, self.num_heads, self.head_dim]) + }; + let q = reshape_heads(self.q_proj.forward(x)); + let k = reshape_heads(self.k_proj.forward(x)); + let v = reshape_heads(self.v_proj.forward(x)); + + let q = apply_rotary_pos_emb_vision(&q, rotary_pos_emb); + let k = apply_rotary_pos_emb_vision(&k, rotary_pos_emb); + + // [seq, heads, head_dim] -> [1, heads, seq, head_dim] + let to_bhsd = |t: &MlxArray| { + let t = mlxcel_core::transpose_axes(t, &[1, 0, 2]); + mlxcel_core::expand_dims(&t, 0) + }; + let q = to_bhsd(&q); + let k = to_bhsd(&k); + let v = to_bhsd(&v); + + let num_segments = cu_seqlens.len() - 1; + let mut attn_outputs = Vec::with_capacity(num_segments); + for seg in 0..num_segments { + let start = cu_seqlens[seg]; + let end = cu_seqlens[seg + 1]; + let take = |t: &MlxArray| { + mlxcel_core::slice( + t, + &[0, 0, start, 0], + &[1, self.num_heads, end, self.head_dim], + ) + }; + let attn = unsafe { + mlxcel_core::layers::attention_from_ptr( + &take(&q), + &take(&k), + &take(&v), + self.scale, + std::ptr::null(), + 0.0, + 0, + ) + }; + attn_outputs.push(attn); + } + + let output = if attn_outputs.len() == 1 { + attn_outputs.into_iter().next().unwrap() + } else { + concat_many(&attn_outputs, 2) + }; + + let output = mlxcel_core::squeeze_axis(&output, 0); + let output = mlxcel_core::transpose_axes(&output, &[1, 0, 2]); + let output = mlxcel_core::reshape(&output, &[seq_length, -1]); + self.out_proj.forward(&output) + } +} + +struct VisionMlp { + fc1: VisionLinear, + fc2: VisionLinear, +} + +impl VisionMlp { + fn from_weights(weights: &WeightMap, prefix: &str) -> Result { + Ok(Self { + fc1: VisionLinear::load(weights, &format!("{}.fc1", prefix))?, + fc2: VisionLinear::load(weights, &format!("{}.fc2", prefix))?, + }) + } + + fn forward(&self, x: &MlxArray) -> UniquePtr { + let h = self.fc1.forward(x); + let h = mlxcel_core::gelu(&h); + self.fc2.forward(&h) + } +} + +struct VisionLayer { + layer_norm1: LayerNorm, + layer_norm2: LayerNorm, + attn: VisionAttention, + mlp: VisionMlp, +} + +impl VisionLayer { + fn from_weights( + weights: &WeightMap, + config: &MiniMaxM3VisionConfig, + prefix: &str, + ) -> Result { + let eps = config.layer_norm_eps; + Ok(Self { + layer_norm1: load_layer_norm(weights, &format!("{}.layer_norm1", prefix), eps)?, + layer_norm2: load_layer_norm(weights, &format!("{}.layer_norm2", prefix), eps)?, + attn: VisionAttention::from_weights(weights, config, &format!("{}.self_attn", prefix))?, + mlp: VisionMlp::from_weights(weights, &format!("{}.mlp", prefix))?, + }) + } + + fn forward( + &self, + hidden_states: &MlxArray, + cu_seqlens: &[i32], + rotary_pos_emb: &MlxArray, + ) -> UniquePtr { + let normed = self.layer_norm1.forward(hidden_states); + let attn_out = self.attn.forward(&normed, cu_seqlens, rotary_pos_emb); + let h = mlxcel_core::add(hidden_states, &attn_out); + let normed = self.layer_norm2.forward(&h); + let mlp_out = self.mlp.forward(&normed); + mlxcel_core::add(&h, &mlp_out) + } +} + +// ============================================================================ +// Two-stage projector +// ============================================================================ + +/// Per-patch projector into `projection_dim` (`linear_1` -> GELU -> `linear_2`). +struct MultiModalProjector { + linear_1: VisionLinear, + linear_2: VisionLinear, +} + +impl MultiModalProjector { + fn from_weights(weights: &WeightMap, prefix: &str) -> Result { + Ok(Self { + linear_1: VisionLinear::load(weights, &format!("{}.linear_1", prefix))?, + linear_2: VisionLinear::load(weights, &format!("{}.linear_2", prefix))?, + }) + } + + fn forward(&self, x: &MlxArray) -> UniquePtr { + let h = self.linear_1.forward(x); + let h = mlxcel_core::gelu(&h); + self.linear_2.forward(&h) + } +} + +/// Patch-merge MLP: folds `spatial_merge_size^2` adjacent patches +/// (`linear_1` [projection_dim, merge^2 * projection_dim] -> GELU -> +/// `linear_2`) into the text hidden size. +struct PatchMergeMlp { + linear_1: VisionLinear, + linear_2: VisionLinear, + fold: i32, + projection_dim: i32, +} + +impl PatchMergeMlp { + fn from_weights( + weights: &WeightMap, + prefix: &str, + merge: usize, + projection_dim: usize, + ) -> Result { + Ok(Self { + linear_1: VisionLinear::load(weights, &format!("{}.linear_1", prefix))?, + linear_2: VisionLinear::load(weights, &format!("{}.linear_2", prefix))?, + fold: (merge * merge) as i32, + projection_dim: projection_dim as i32, + }) + } + + /// `x`: `[num_patches, projection_dim]`. Adjacent groups of `fold` patches + /// are the same 2x2 spatial-merge cell (the processor emits them + /// contiguously), so a row-major reshape concatenates their features. + fn forward(&self, x: &MlxArray) -> UniquePtr { + let merged = mlxcel_core::reshape(x, &[-1, self.fold * self.projection_dim]); + let h = self.linear_1.forward(&merged); + let h = mlxcel_core::gelu(&h); + self.linear_2.forward(&h) + } +} + +// ============================================================================ +// Encoder +// ============================================================================ + +pub struct MiniMaxM3VisionEncoder { + patch_embed: PatchEmbed, + pre_layrnorm: LayerNorm, + rotary_pos_emb: VisionRotaryEmbedding, + layers: Vec, + projector: MultiModalProjector, + patch_merge: PatchMergeMlp, + spatial_merge_size: usize, +} + +impl MiniMaxM3VisionEncoder { + /// Build the tower from the full (raw) weight map. The vision layers live + /// under `vision_tower.vision_model.*` while the two-stage projector lives + /// at the top level (`multi_modal_projector.*`, `patch_merge_mlp.*`). + pub fn from_weights( + weights: &WeightMap, + config: &MiniMaxM3VisionConfig, + ) -> Result { + let vt = "vision_tower.vision_model"; + let patch_embed = PatchEmbed::from_weights( + weights, + config, + &format!("{}.embeddings.patch_embedding", vt), + )?; + let pre_layrnorm = load_layer_norm( + weights, + &format!("{}.pre_layrnorm", vt), + config.layer_norm_eps, + )?; + + // The vision rope table covers head_dim/2 frequencies (h and w halves), + // matching the shared Qwen2-VL helper. + let rotary_pos_emb = VisionRotaryEmbedding::new(config.head_dim() / 2); + + let mut layers = Vec::with_capacity(config.num_hidden_layers); + for i in 0..config.num_hidden_layers { + layers.push(VisionLayer::from_weights( + weights, + config, + &format!("{}.encoder.layers.{}", vt, i), + )?); + } + + let projector = MultiModalProjector::from_weights(weights, "multi_modal_projector")?; + let patch_merge = PatchMergeMlp::from_weights( + weights, + "patch_merge_mlp", + config.spatial_merge_size(), + config.projection_dim, + )?; + + Ok(Self { + patch_embed, + pre_layrnorm, + rotary_pos_emb, + layers, + projector, + patch_merge, + spatial_merge_size: config.spatial_merge_size(), + }) + } + + /// 2D (h, w) rotary position embeddings, merge-grouped to match the + /// processor patch order. Identical construction to the Qwen2-VL tower. + fn rot_pos_emb(&self, grid_thw: &[(i32, i32, i32)]) -> UniquePtr { + let mut all_pos_ids: Vec> = Vec::new(); + let mut max_grid_dim: i32 = 0; + let merge = self.spatial_merge_size as i32; + + for &(t, h, w) in grid_thw { + max_grid_dim = max_grid_dim.max(h).max(w); + + let h_arange = mlxcel_core::arange_i32(0, h, 1); + let h_col = mlxcel_core::reshape(&h_arange, &[h, 1]); + let hpos = mlxcel_core::repeat(&h_col, w, 1); + let hpos = mlxcel_core::reshape(&hpos, &[h / merge, merge, w / merge, merge]); + let hpos = mlxcel_core::transpose_axes(&hpos, &[0, 2, 1, 3]); + let hpos = mlxcel_core::flatten(&hpos); + + let w_arange = mlxcel_core::arange_i32(0, w, 1); + let w_row = mlxcel_core::reshape(&w_arange, &[1, w]); + let wpos = mlxcel_core::repeat(&w_row, h, 0); + let wpos = mlxcel_core::reshape(&wpos, &[h / merge, merge, w / merge, merge]); + let wpos = mlxcel_core::transpose_axes(&wpos, &[0, 2, 1, 3]); + let wpos = mlxcel_core::flatten(&wpos); + + let stacked = mlxcel_core::stack_owned(&[hpos, wpos], -1); + let tiled = mlxcel_core::tile(&stacked, &[t, 1]); + all_pos_ids.push(tiled); + } + + let pos_ids = if all_pos_ids.len() == 1 { + all_pos_ids.into_iter().next().unwrap() + } else { + concat_many(&all_pos_ids, 0) + }; + + let rotary_table = self.rotary_pos_emb.forward(max_grid_dim); + let pos_ids_flat = mlxcel_core::flatten(&pos_ids); + let all_freqs = mlxcel_core::take(&rotary_table, &pos_ids_flat, 0); + let total_shape = mlxcel_core::array_shape(&pos_ids); + let total_tokens = total_shape[0]; + let freq_shape = mlxcel_core::array_shape(&all_freqs); + let half_dim = freq_shape[1]; + let all_freqs = mlxcel_core::reshape(&all_freqs, &[total_tokens, 2, half_dim]); + mlxcel_core::reshape(&all_freqs, &[total_tokens, 2 * half_dim]) + } + + /// Per-image `cu_seqlens` (full attention over each image's patches): + /// `h * w` tokens per temporal frame. + fn compute_cu_seqlens(grid_thw: &[(i32, i32, i32)]) -> Vec { + let mut cu_seqlens = vec![0i32]; + let mut cumulative = 0i32; + for &(t, h, w) in grid_thw { + let tokens_per_frame = h * w; + for _ in 0..t { + cumulative += tokens_per_frame; + cu_seqlens.push(cumulative); + } + } + cu_seqlens + } + + /// `hidden_states`: `[num_patches, in_features]` (channels-last patch rows), + /// `grid_thw`: per-image `(t, h, w)`. Returns `[num_merged_tokens, + /// text_hidden]` where `num_merged_tokens = sum(t * h * w) / merge^2`. + pub fn forward_with_grid( + &self, + hidden_states: &MlxArray, + grid_thw: &[(i32, i32, i32)], + ) -> VisionEncoderOutput { + let hidden_states = mlxcel_core::astype(hidden_states, mlxcel_core::dtype::FLOAT32); + let mut h = self.patch_embed.forward(&hidden_states); + h = self.pre_layrnorm.forward(&h); + + let rotary_pos_emb = self.rot_pos_emb(grid_thw); + let cu_seqlens = Self::compute_cu_seqlens(grid_thw); + + for layer in &self.layers { + h = layer.forward(&h, &cu_seqlens, &rotary_pos_emb); + } + + // Per-patch projection, then fold spatial_merge_size^2 patches. + h = self.projector.forward(&h); + h = self.patch_merge.forward(&h); + + VisionEncoderOutput { hidden_states: h } + } +} + +/// VisionEncoder trait - panics since grid_thw is required. +impl super::VisionEncoder for MiniMaxM3VisionEncoder { + fn forward(&self, _pixel_values: &MlxArray) -> VisionEncoderOutput { + panic!("MiniMax-M3-VL vision encoder requires grid_thw; use forward_with_grid() instead"); + } +} diff --git a/src/vision/encoders/mod.rs b/src/vision/encoders/mod.rs index c6ccb46cb..487a0f29c 100644 --- a/src/vision/encoders/mod.rs +++ b/src/vision/encoders/mod.rs @@ -35,6 +35,7 @@ pub mod lfm2_vl; pub mod llama4; pub mod minicpmo; pub mod minicpmv4_6; +pub mod minimax_m3_vl; pub mod mllama; pub mod molmo; pub mod molmo2; diff --git a/src/vision/minimax_m3_vl.rs b/src/vision/minimax_m3_vl.rs new file mode 100644 index 000000000..2df008033 --- /dev/null +++ b/src/vision/minimax_m3_vl.rs @@ -0,0 +1,132 @@ +// Copyright 2025-2026 Lablup Inc. and Jeongkyu Shin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! MiniMax-M3-VL Vision-Language Model (`model_type: "minimax_m3_vl"`). +//! +//! Composes the CLIP-style `MiniMaxM3VisionEncoder` (tower + two-stage +//! projector, which already projects to the text hidden size) with the +//! MiniMax-M3 hybrid dense/MoE text backbone. There is no MRoPE and no separate +//! connector: the projector output is scattered LLaVA-style into the image +//! placeholder positions of the embedded prompt, and the decoder runs its +//! standard partial 1D RoPE. +//! +//! The 427B checkpoint cannot be loaded on the development machine, so the +//! validated surface is the synthetic reduced-config unit tests plus the +//! real-config parse test. + +use super::{encoders, merge, processors}; +use crate::LanguageModel; +use crate::models::MiniMaxM3Model; +use mlxcel_core::cache::SequenceId; +use mlxcel_core::layers::KVCache; +use mlxcel_core::{MlxArray, UniquePtr}; + +pub struct MiniMaxM3VlModel { + pub text_model: MiniMaxM3Model, + pub vision_encoder: encoders::minimax_m3_vl::MiniMaxM3VisionEncoder, + pub processor: processors::minimax_m3::MiniMaxM3Processor, + /// `]<]image[>[` scatter target (200025 in the real checkpoint). + pub image_token_id: i32, + /// `]<]video[>[` (200026). Video is out of scope for this port. + pub video_token_id: i32, + /// `]<]start of image[>[` (200029). + pub vision_start_token_id: i32, + /// `]<]end of image[>[` (200030). + pub vision_end_token_id: i32, + pub spatial_merge_size: i32, + pub eos_token_ids: Vec, +} + +impl MiniMaxM3VlModel { + /// Encode images and scatter the merged features at image-placeholder + /// positions. The vision tower emits `grid.prod() / merge^2` tokens per + /// image, which matches the placeholder expansion count produced by the + /// shared Qwen-VL insertion helper (`t * (h/merge) * (w/merge)`). + pub fn input_embeddings( + &self, + input_ids: &MlxArray, + pixel_values: &MlxArray, + grid_thw: &[(i32, i32, i32)], + ) -> merge::InputEmbeddings { + let inputs_embeds = self.text_model.get_embed_tokens(input_ids); + // forward_with_grid casts pixel_values to f32 internally; merge_llava + // casts the projected features back to the text embedding dtype. + let vision_output = self + .vision_encoder + .forward_with_grid(pixel_values, grid_thw); + merge::merge_llava( + self.image_token_id, + &vision_output.hidden_states, + &inputs_embeds, + input_ids, + ) + } +} + +impl LanguageModel for MiniMaxM3VlModel { + fn forward( + &self, + input_ids: &MlxArray, + caches: &mut [KVCache], + mask: Option<&MlxArray>, + ) -> UniquePtr { + self.text_model.forward(input_ids, caches, mask) + } + + fn forward_with_embeddings( + &self, + input_ids: &MlxArray, + input_embeddings: Option<&MlxArray>, + caches: &mut [KVCache], + mask: Option<&MlxArray>, + ) -> UniquePtr { + self.text_model + .forward_with_embeddings_impl(input_ids, input_embeddings, caches, mask) + } + + fn embed_tokens(&self, input_ids: &MlxArray) -> Option> { + Some(self.text_model.get_embed_tokens(input_ids)) + } + + fn forward_with_sequence_id( + &self, + input_ids: &MlxArray, + _seq_id: Option, + caches: &mut [KVCache], + mask: Option<&MlxArray>, + ) -> UniquePtr { + self.text_model.forward(input_ids, caches, mask) + } + + fn make_caches(&self) -> Vec { + self.text_model.make_caches() + } + + fn num_layers(&self) -> usize { + mlxcel_core::generate::LanguageModel::num_layers(&self.text_model) + } + + fn eos_token_ids(&self) -> Vec { + self.eos_token_ids.clone() + } + + fn output_suppressed_token_ids(&self) -> Vec { + // The image placeholder id must never be sampled during decode. + vec![self.image_token_id] + } +} + +#[cfg(test)] +#[path = "minimax_m3_vl_tests.rs"] +mod minimax_m3_vl_tests; diff --git a/src/vision/minimax_m3_vl_tests.rs b/src/vision/minimax_m3_vl_tests.rs new file mode 100644 index 000000000..9feccc5c5 --- /dev/null +++ b/src/vision/minimax_m3_vl_tests.rs @@ -0,0 +1,366 @@ +// Copyright 2025-2026 Lablup Inc. and Jeongkyu Shin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Checkpoint-free unit tests for MiniMax-M3-VL. +//! +//! The 427B checkpoint cannot be loaded on the development machine, so these +//! cover the surface reachable without weights: real nested-config parsing +//! (`vision_config` and `text_config`), the sanitizer's vision/projector skip +//! against verbatim checkpoint keys, the projector fold ordering +//! (`merge^2 * projection_dim`), a tiny synthetic tower forward (patch embed -> +//! pre_layrnorm -> CLIP layers with `cu_seqlens` + 2D vision RoPE -> two-stage +//! projector), and the placeholder-count invariant against the shared Qwen-VL +//! insertion helper. Run serially (`--test-threads=1`); the MLX ops touch the +//! device. + +use crate::models::minimax_m3; +use crate::vision::encoders::minimax_m3_vl::{MiniMaxM3VisionConfig, MiniMaxM3VisionEncoder}; +use mlxcel_core::MlxArray; +use mlxcel_core::weights::WeightMap; + +fn filled(shape: &[i32], val: f32) -> mlxcel_core::UniquePtr { + let n: i32 = shape.iter().product(); + mlxcel_core::from_slice_f32(&vec![val; n as usize], shape) +} + +fn reduce_max_abs(a: &MlxArray) -> f32 { + let flat = mlxcel_core::reshape(a, &[-1]); + let m = mlxcel_core::max_axis(&mlxcel_core::abs(&flat), 0, false); + mlxcel_core::eval(&m); + mlxcel_core::item_f32(&m) +} + +// A reduced tower config: hidden 128, 2 heads (head_dim 64, a fused-SDPA +// supported width), 1 layer, patch 2, projection_dim 16, merge 2, +// temporal_patch 2. in_features = 3*2*2*2 = 24. +const TINY_VISION_CONFIG: &str = r#"{ + "model_type": "clip_vision_model", + "hidden_size": 128, + "num_attention_heads": 2, + "num_hidden_layers": 1, + "intermediate_size": 64, + "patch_size": 2, + "projection_dim": 16, + "rope_theta": 10000.0, + "layer_norm_eps": 1e-05, + "img_token_compression_config": { + "image_token_compression_method": "patch_merge", + "spatial_merge_size": 2, + "temporal_patch_size": 2 + } +}"#; + +fn tiny_vision_config() -> MiniMaxM3VisionConfig { + serde_json::from_str(TINY_VISION_CONFIG).expect("tiny vision config parses") +} + +/// Verbatim-named synthetic weights for the tiny tower. +fn tiny_tower_weights(cfg: &MiniMaxM3VisionConfig) -> WeightMap { + let hidden = cfg.hidden_size as i32; + let inter = cfg.intermediate_size as i32; + let proj = cfg.projection_dim as i32; + let in_features = + (cfg.in_channels * cfg.temporal_patch_size() * cfg.patch_size * cfg.patch_size) as i32; + let fold = (cfg.spatial_merge_size() * cfg.spatial_merge_size()) as i32; + + let vt = "vision_tower.vision_model"; + let mut w = WeightMap::new(); + + // Patch embedding (2D form) + pre_layrnorm. + w.insert( + format!("{vt}.embeddings.patch_embedding.weight"), + filled(&[hidden, in_features], 0.02), + ); + w.insert(format!("{vt}.pre_layrnorm.weight"), filled(&[hidden], 1.0)); + w.insert(format!("{vt}.pre_layrnorm.bias"), filled(&[hidden], 0.0)); + + // Encoder layer 0. + let l = format!("{vt}.encoder.layers.0"); + for ln in ["layer_norm1", "layer_norm2"] { + w.insert(format!("{l}.{ln}.weight"), filled(&[hidden], 1.0)); + w.insert(format!("{l}.{ln}.bias"), filled(&[hidden], 0.0)); + } + for p in ["q_proj", "k_proj", "v_proj", "out_proj"] { + w.insert( + format!("{l}.self_attn.{p}.weight"), + filled(&[hidden, hidden], 0.05), + ); + w.insert(format!("{l}.self_attn.{p}.bias"), filled(&[hidden], 0.0)); + } + w.insert( + format!("{l}.mlp.fc1.weight"), + filled(&[inter, hidden], 0.03), + ); + w.insert(format!("{l}.mlp.fc1.bias"), filled(&[inter], 0.0)); + w.insert( + format!("{l}.mlp.fc2.weight"), + filled(&[hidden, inter], 0.03), + ); + w.insert(format!("{l}.mlp.fc2.bias"), filled(&[hidden], 0.0)); + + // Two-stage projector. + w.insert( + "multi_modal_projector.linear_1.weight".into(), + filled(&[proj, hidden], 0.04), + ); + w.insert( + "multi_modal_projector.linear_1.bias".into(), + filled(&[proj], 0.0), + ); + w.insert( + "multi_modal_projector.linear_2.weight".into(), + filled(&[proj, proj], 0.04), + ); + w.insert( + "multi_modal_projector.linear_2.bias".into(), + filled(&[proj], 0.0), + ); + w.insert( + "patch_merge_mlp.linear_1.weight".into(), + filled(&[proj, fold * proj], 0.02), + ); + w.insert("patch_merge_mlp.linear_1.bias".into(), filled(&[proj], 0.0)); + w.insert( + "patch_merge_mlp.linear_2.weight".into(), + filled(&[proj, proj], 0.02), + ); + w.insert("patch_merge_mlp.linear_2.bias".into(), filled(&[proj], 0.0)); + + w +} + +#[test] +fn tower_loads_verbatim_keys_and_emits_merged_token_shape() { + // A single 2x2-patch-grid image (grid (1,2,2)) -> 4 patches -> 1 merged + // token after the merge^2 fold. Exercises patch embed, pre_layrnorm, the + // CLIP layer (cu_seqlens attention + 2D vision RoPE), and both projector + // stages, and proves the tower resolves every verbatim key (including the + // `pre_layrnorm` spelling). + let cfg = tiny_vision_config(); + let weights = tiny_tower_weights(&cfg); + let encoder = MiniMaxM3VisionEncoder::from_weights(&weights, &cfg).expect("tiny tower loads"); + + let in_features = + (cfg.in_channels * cfg.temporal_patch_size() * cfg.patch_size * cfg.patch_size) as i32; + let grid = vec![(1i32, 2i32, 2i32)]; + let num_patches = 4i32; + let pixel_values = filled(&[num_patches, in_features], 0.1); + + let out = encoder.forward_with_grid(&pixel_values, &grid); + mlxcel_core::eval(&out.hidden_states); + + let merge = cfg.spatial_merge_size() as i32; + let (t, h, w) = grid[0]; + let expected_tokens = t * (h / merge) * (w / merge); // grid.prod()/merge^2 = 1 + assert_eq!( + mlxcel_core::array_shape(&out.hidden_states), + vec![expected_tokens, cfg.projection_dim as i32] + ); + assert!( + reduce_max_abs(&out.hidden_states).is_finite(), + "tower forward must be finite" + ); +} + +#[test] +fn projector_fold_concatenates_four_adjacent_patches_in_order() { + // The patch-merge fold is a row-major reshape [P, D] -> [P/4, 4*D]. The four + // patches of a 2x2 merge cell are contiguous in the processor's patch + // order, so the reshape must place patch i's features in the i-th D-slice. + let d = 3i32; + let rows = 4i32; + // Row i is the constant (i+1); after fold the 4*D vector must read + // [1,1,1, 2,2,2, 3,3,3, 4,4,4]. + let mut data = Vec::new(); + for i in 0..rows { + for _ in 0..d { + data.push((i + 1) as f32); + } + } + let patches = mlxcel_core::from_slice_f32(&data, &[rows, d]); + let folded = mlxcel_core::reshape(&patches, &[-1, 4 * d]); + mlxcel_core::eval(&folded); + assert_eq!(mlxcel_core::array_shape(&folded), vec![1, 4 * d]); + + for (patch_idx, &expected) in [1.0f32, 2.0, 3.0, 4.0].iter().enumerate() { + let start = patch_idx as i32 * d; + let slice = mlxcel_core::slice(&folded, &[0, start], &[1, start + d]); + let diff = mlxcel_core::subtract(&slice, &filled(&[1, d], expected)); + assert!( + reduce_max_abs(&diff) < 1e-6, + "patch {patch_idx} features must occupy fold slice {patch_idx}" + ); + } +} + +#[test] +fn vision_config_parses_real_nested_values() { + // The real MiniMaxAI/MiniMax-M3 vision_config (with the vestigial LLaVA-style + // keys present, which must be ignored). + let json = r#"{ + "model_type": "clip_vision_model", + "hidden_size": 1280, + "num_attention_heads": 16, + "num_hidden_layers": 32, + "intermediate_size": 5120, + "patch_size": 14, + "image_size": 2016, + "projection_dim": 6144, + "position_embedding_type": "rope", + "rope_mode": "3d", + "rope_theta": 10000.0, + "hidden_act": "gelu", + "layer_norm_eps": 1e-05, + "img_token_compression_config": { + "image_token_compression_method": "patch_merge", + "spatial_merge_size": 2, + "temporal_patch_size": 2 + }, + "vision_segment_max_frames": 4, + "image_grid_pinpoints": [[336, 336]], + "vision_feature_layer": -1, + "vision_feature_select_strategy": "full", + "image_seq_length": 576 + }"#; + let cfg: MiniMaxM3VisionConfig = serde_json::from_str(json).expect("real vision_config parses"); + assert_eq!(cfg.hidden_size, 1280); + assert_eq!(cfg.num_attention_heads, 16); + assert_eq!(cfg.num_hidden_layers, 32); + assert_eq!(cfg.intermediate_size, 5120); + assert_eq!(cfg.patch_size, 14); + assert_eq!(cfg.projection_dim, 6144); + assert_eq!(cfg.head_dim(), 80); + assert!((cfg.rope_theta - 10000.0).abs() < 1.0); + assert!((cfg.layer_norm_eps - 1e-5).abs() < 1e-9); + assert_eq!(cfg.spatial_merge_size(), 2); + assert_eq!(cfg.temporal_patch_size(), 2); + // 24576 = merge^2 * projection_dim, the patch-merge fold width. + assert_eq!( + cfg.spatial_merge_size() * cfg.spatial_merge_size() * cfg.projection_dim, + 24576 + ); +} + +#[test] +fn nested_config_parses_text_and_vision_blocks() { + // The top-level minimax_m3_vl config nests text_config + vision_config. + let json = r#"{ + "model_type": "minimax_m3_vl", + "text_config": { + "model_type": "minimax_m3", + "hidden_size": 6144, + "intermediate_size": 3072, + "num_hidden_layers": 60, + "num_attention_heads": 64, + "num_key_value_heads": 4, + "head_dim": 128, + "vocab_size": 200064, + "num_local_experts": 128, + "num_experts_per_tok": 4 + }, + "vision_config": { + "model_type": "clip_vision_model", + "hidden_size": 1280, + "num_attention_heads": 16, + "num_hidden_layers": 32, + "intermediate_size": 5120, + "patch_size": 14, + "projection_dim": 6144 + } + }"#; + let full: serde_json::Value = serde_json::from_str(json).expect("nested config parses"); + + let text_args: minimax_m3::ModelArgs = + serde_json::from_value(full.get("text_config").cloned().unwrap()) + .expect("text_config parses into MiniMax-M3 ModelArgs"); + assert_eq!(text_args.hidden_size, 6144); + assert_eq!(text_args.num_hidden_layers, 60); + assert_eq!(text_args.num_local_experts, 128); + + let vcfg: MiniMaxM3VisionConfig = + serde_json::from_value(full.get("vision_config").cloned().unwrap()) + .expect("vision_config parses"); + assert_eq!(vcfg.hidden_size, 1280); + assert_eq!(vcfg.projection_dim, 6144); + // Vision projector output must equal the text hidden size for the merge. + assert_eq!(vcfg.projection_dim, text_args.hidden_size); +} + +#[test] +fn text_sanitizer_drops_verbatim_vision_and_projector_keys() { + // The VL loader builds the tower from the raw keys, then hands the map to + // the MiniMax-M3 text sanitizer, which must drop every vision/projector + // tensor (verbatim names, including the pre_layrnorm spelling) and rewrite + // language_model.model.* -> model.*. + let mut weights = WeightMap::new(); + for key in [ + "language_model.model.embed_tokens.weight", + "language_model.model.norm.weight", + "language_model.lm_head.weight", + "vision_tower.vision_model.pre_layrnorm.weight", + "vision_tower.vision_model.embeddings.patch_embedding.weight", + "vision_tower.vision_model.encoder.layers.0.self_attn.q_proj.weight", + "multi_modal_projector.linear_1.weight", + "patch_merge_mlp.linear_1.weight", + ] { + weights.insert(key.to_string(), filled(&[1], 0.0)); + } + + let args: minimax_m3::ModelArgs = serde_json::from_str( + r#"{"model_type":"minimax_m3","hidden_size":8,"intermediate_size":8,"num_hidden_layers":4,"num_attention_heads":4,"num_key_value_heads":4,"vocab_size":16,"num_local_experts":4,"num_experts_per_tok":2}"#, + ) + .expect("args parse"); + let out = minimax_m3::sanitize_weights(weights, &args); + + assert!(out.contains_key("model.embed_tokens.weight")); + assert!(out.contains_key("model.norm.weight")); + assert!(out.contains_key("model.lm_head.weight")); + assert!(!out.contains_key("vision_tower.vision_model.pre_layrnorm.weight")); + assert!(!out.contains_key("vision_tower.vision_model.embeddings.patch_embedding.weight")); + assert!( + !out.contains_key("vision_tower.vision_model.encoder.layers.0.self_attn.q_proj.weight") + ); + assert!(!out.contains_key("multi_modal_projector.linear_1.weight")); + assert!(!out.contains_key("patch_merge_mlp.linear_1.weight")); + // Only the 3 text tensors survive. + assert_eq!(out.len(), 3); +} + +#[test] +fn placeholder_expansion_matches_merged_token_count() { + // The shared Qwen-VL insertion helper expands one placeholder per image to + // t*(h/merge)*(w/merge) image tokens, which must equal the tower's merged + // token count (grid.prod()/merge^2). + let merge = 2usize; + let grid = vec![(1i32, 6i32, 8i32)]; + let image_token_id = 200025; + let vision_start = 200029; + let mut prompt = vec![1i32, image_token_id, 42i32]; + + let stats = crate::qwen_vl::insert_qwen_vl_image_tokens( + &mut prompt, + &grid, + merge, + vision_start, + image_token_id, + ) + .expect("insertion succeeds for a one-placeholder-per-image prompt"); + + let (t, h, w) = grid[0]; + let merged_tokens = t * (h / merge as i32) * (w / merge as i32); + assert_eq!(stats.total_image_tokens, merged_tokens); + // Expanded prompt: bos + merged_tokens image tokens + trailing text. + let placeholders = prompt.iter().filter(|&&t| t == image_token_id).count() as i32; + assert_eq!(placeholders, merged_tokens); +} diff --git a/src/vision/mod.rs b/src/vision/mod.rs index 025c4b0a1..17df138db 100644 --- a/src/vision/mod.rs +++ b/src/vision/mod.rs @@ -61,6 +61,7 @@ pub mod kimi_vl; pub mod lfm2_vl; pub mod minicpmo_vl; pub mod minicpmv4_6_vl; +pub mod minimax_m3_vl; pub mod mllama_vl; pub mod molmo2_vl; pub mod molmo_point_vl; @@ -103,6 +104,7 @@ pub use kimi_vl::{KimiVLModel, KimiVLMultiModalProjector}; pub use lfm2_vl::Lfm2VlModel; pub use minicpmo_vl::MiniCPMOVLModel; pub use minicpmv4_6_vl::MiniCPMV46VLModel; +pub use minimax_m3_vl::MiniMaxM3VlModel; pub use mllama_vl::MllamaVLModel; pub use molmo_point_vl::MolmoPointVLModel; pub use molmo_vl::MolmoVLModel; diff --git a/src/vision/processors/minimax_m3.rs b/src/vision/processors/minimax_m3.rs new file mode 100644 index 000000000..78e7c6977 --- /dev/null +++ b/src/vision/processors/minimax_m3.rs @@ -0,0 +1,287 @@ +// Copyright 2025-2026 Lablup Inc. and Jeongkyu Shin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! MiniMax-M3-VL image processor. +//! +//! A faithful port of the checkpoint's `image_processor.py` +//! (`MiniMaxM3VLImageProcessor`), which is Qwen2-VL-style dynamic-resolution +//! preprocessing: +//! 1. `smart_resize` to dimensions that are multiples of +//! `factor = patch_size * merge_size = 28`, bounded by `min_pixels` / +//! `max_pixels` (672x672), +//! 2. CLIP mean/std normalization, +//! 3. temporal padding: repeat the last frame to a multiple of +//! `temporal_patch_size` (a single image becomes 2 identical frames), +//! 4. patchify into the exact `(grid_t, grid_h/m, grid_w/m, m, m, C, temporal, +//! patch_h, patch_w)` order, emitting `pixel_values` of shape +//! `[num_patches, C * temporal * patch * patch]` (= `[num_patches, 1176]`) +//! plus `image_grid_thw`. +//! +//! The patch sequence order groups each 2x2 spatial-merge cell contiguously and +//! the per-row feature order is `[channel, temporal, patch_h, patch_w]`, which +//! is what the tower's patch embedding and patch-merge fold both assume. + +use super::ImageProcessor; +use image::imageops::FilterType; +use mlxcel_core::{MlxArray, UniquePtr}; + +const MAX_RATIO: f64 = 200.0; + +pub struct MiniMaxM3Processor { + pub patch_size: usize, + pub temporal_patch_size: usize, + pub spatial_merge_size: usize, + pub min_pixels: usize, + pub max_pixels: usize, + pub mean: [f32; 3], + pub std: [f32; 3], +} + +impl Default for MiniMaxM3Processor { + fn default() -> Self { + Self { + patch_size: 14, + temporal_patch_size: 2, + spatial_merge_size: 2, + min_pixels: 4 * 28 * 28, // 3136 + max_pixels: 451_584, // 672 * 672 + mean: [0.481_454_66, 0.457_827_5, 0.408_210_73], + std: [0.268_629_54, 0.261_302_58, 0.275_777_11], + } + } +} + +fn round_by_factor(n: f64, factor: f64) -> f64 { + (n / factor).round() * factor +} +fn ceil_by_factor(n: f64, factor: f64) -> f64 { + (n / factor).ceil() * factor +} +fn floor_by_factor(n: f64, factor: f64) -> f64 { + (n / factor).floor() * factor +} + +impl MiniMaxM3Processor { + fn factor(&self) -> f64 { + (self.patch_size * self.spatial_merge_size) as f64 + } + + /// Port of `smart_resize`. Returns `(height, width)`, both multiples of + /// `factor`. Aspect ratios beyond `MAX_RATIO` are clamped to `MAX_RATIO` + /// rather than raising, so a degenerate input still yields a valid grid. + pub fn smart_resize(&self, height: u32, width: u32) -> (u32, u32) { + let factor = self.factor(); + let mut h = height as f64; + let mut w = width as f64; + + let ratio = h.max(w) / h.min(w).max(1.0); + if ratio > MAX_RATIO { + // Clamp the long side so the ratio equals MAX_RATIO. + if h > w { + h = w * MAX_RATIO; + } else { + w = h * MAX_RATIO; + } + } + + let mut h_bar = factor.max(round_by_factor(h, factor)); + let mut w_bar = factor.max(round_by_factor(w, factor)); + + if h_bar * w_bar > self.max_pixels as f64 { + let beta = (h * w / self.max_pixels as f64).sqrt(); + h_bar = floor_by_factor(h / beta, factor).max(factor); + w_bar = floor_by_factor(w / beta, factor).max(factor); + } else if h_bar * w_bar < self.min_pixels as f64 { + let beta = (self.min_pixels as f64 / (h * w)).sqrt(); + h_bar = ceil_by_factor(h * beta, factor); + w_bar = ceil_by_factor(w * beta, factor); + } + + (h_bar as u32, w_bar as u32) + } + + /// Per-image `(temporal, grid_h, grid_w)` in post-resize patch units. + pub fn compute_grid_thw(&self, images: &[image::DynamicImage]) -> Vec<(i32, i32, i32)> { + images + .iter() + .map(|img| { + let (h, w) = self.smart_resize(img.height(), img.width()); + let grid_h = h as i32 / self.patch_size as i32; + let grid_w = w as i32 / self.patch_size as i32; + (1i32, grid_h, grid_w) + }) + .collect() + } + + /// Preprocess images into `(pixel_values, grid_thw)`. + /// + /// `pixel_values` has shape `[sum(t * grid_h * grid_w), 1176]` with patch + /// rows in merge-grouped order and per-row feature order + /// `[channel, temporal, patch_h, patch_w]`. + pub fn preprocess_with_grid( + &self, + images: &[image::DynamicImage], + ) -> (UniquePtr, Vec<(i32, i32, i32)>) { + let grid_thw = self.compute_grid_thw(images); + let patch = self.patch_size; + let merge = self.spatial_merge_size; + let in_channels = 3usize; + let features_per_row = in_channels * self.temporal_patch_size * patch * patch; + + let mut all_patches: Vec = Vec::new(); + + for (img_idx, img) in images.iter().enumerate() { + let (_, grid_h, grid_w) = grid_thw[img_idx]; + let grid_h = grid_h as usize; + let grid_w = grid_w as usize; + let target_h = (grid_h * patch) as u32; + let target_w = (grid_w * patch) as u32; + + let resized = img.resize_exact(target_w, target_h, FilterType::Lanczos3); + let rgb = resized.to_rgb8(); + let h = target_h as usize; + let w = target_w as usize; + + // Channels-first normalized buffer: normalized[c*h*w + y*w + x]. + let mut normalized = vec![0f32; in_channels * h * w]; + for y in 0..h { + for x in 0..w { + let pixel = rgb.get_pixel(x as u32, y as u32); + for c in 0..in_channels { + let val = pixel[c] as f32 / 255.0; + normalized[c * h * w + y * w + x] = (val - self.mean[c]) / self.std[c]; + } + } + } + + // Merge-grouped patch order: for each 2x2 merge cell, emit its + // `merge^2` sub-patches contiguously. Per-row features run + // [channel, temporal, patch_h, patch_w]; the temporal frames are + // identical copies (single-image temporal padding). + let hm = grid_h / merge; + let wm = grid_w / merge; + for hh in 0..hm { + for ww in 0..wm { + for mh in 0..merge { + for mw in 0..merge { + let py = hh * merge + mh; + let px = ww * merge + mw; + let y0 = py * patch; + let x0 = px * patch; + for c in 0..in_channels { + for _tp in 0..self.temporal_patch_size { + for dy in 0..patch { + for dx in 0..patch { + let y = y0 + dy; + let x = x0 + dx; + all_patches.push(normalized[c * h * w + y * w + x]); + } + } + } + } + } + } + } + } + } + + let total_rows: usize = grid_thw + .iter() + .map(|&(t, gh, gw)| (t * gh * gw) as usize) + .sum(); + + let pixel_values = mlxcel_core::from_slice_f32( + &all_patches, + &[total_rows as i32, features_per_row as i32], + ); + + (pixel_values, grid_thw) + } +} + +impl ImageProcessor for MiniMaxM3Processor { + fn preprocess(&self, images: &[image::DynamicImage]) -> UniquePtr { + let (pixel_values, _) = self.preprocess_with_grid(images); + pixel_values + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn solid_image(w: u32, h: u32) -> image::DynamicImage { + image::DynamicImage::ImageRgb8(image::RgbImage::from_pixel( + w, + h, + image::Rgb([120, 60, 200]), + )) + } + + #[test] + fn smart_resize_aligns_to_factor_and_bounds() { + let p = MiniMaxM3Processor::default(); + let factor = (p.patch_size * p.spatial_merge_size) as u32; // 28 + for &(h, w) in &[(100u32, 100u32), (37, 512), (1024, 300), (13, 640)] { + let (rh, rw) = p.smart_resize(h, w); + assert_eq!(rh % factor, 0, "height {rh} not multiple of {factor}"); + assert_eq!(rw % factor, 0, "width {rw} not multiple of {factor}"); + let pixels = (rh as usize) * (rw as usize); + assert!(pixels >= p.min_pixels, "pixels {pixels} below min"); + assert!(pixels <= p.max_pixels, "pixels {pixels} above max"); + } + } + + #[test] + fn grid_thw_matches_resized_patch_grid() { + let p = MiniMaxM3Processor::default(); + let img = solid_image(200, 100); // w=200, h=100 + let grid = p.compute_grid_thw(&[img]); + assert_eq!(grid.len(), 1); + let (t, gh, gw) = grid[0]; + assert_eq!(t, 1); + let (rh, rw) = p.smart_resize(100, 200); + assert_eq!(gh, rh as i32 / p.patch_size as i32); + assert_eq!(gw, rw as i32 / p.patch_size as i32); + // grid dims are even (multiple of merge) so the fold-by-merge^2 is exact. + assert_eq!(gh % p.spatial_merge_size as i32, 0); + assert_eq!(gw % p.spatial_merge_size as i32, 0); + } + + #[test] + fn preprocess_emits_1176_dim_rows_and_grid_prod_patches() { + let p = MiniMaxM3Processor::default(); + let img = solid_image(140, 84); // small, non-square + let (pixel_values, grid) = p.preprocess_with_grid(&[img]); + mlxcel_core::eval(&pixel_values); + let shape = mlxcel_core::array_shape(&pixel_values); + let (t, gh, gw) = grid[0]; + let expected_rows = (t * gh * gw) as i32; + assert_eq!(shape, vec![expected_rows, 1176]); + } + + #[test] + fn placeholder_count_equals_grid_prod_over_merge_squared() { + // The vision tower emits grid.prod()/merge^2 merged tokens, which must + // equal the placeholder expansion count. + let p = MiniMaxM3Processor::default(); + let img = solid_image(280, 196); + let grid = p.compute_grid_thw(&[img]); + let (t, gh, gw) = grid[0]; + let merge = p.spatial_merge_size as i32; + let merged_tokens = t * (gh / merge) * (gw / merge); + let grid_prod_over_merge2 = (t * gh * gw) / (merge * merge); + assert_eq!(merged_tokens, grid_prod_over_merge2); + } +} diff --git a/src/vision/processors/mod.rs b/src/vision/processors/mod.rs index 410e22467..a5d90e577 100644 --- a/src/vision/processors/mod.rs +++ b/src/vision/processors/mod.rs @@ -30,6 +30,7 @@ pub mod internvl; pub mod kimi_vl; pub mod lfm2_vl; pub mod minicpmo; +pub mod minimax_m3; pub mod mllama; pub mod molmo; pub mod molmo2; From 6453f552a2482a38fe465925d23e9c30fc9f20d4 Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Thu, 16 Jul 2026 21:51:09 +0900 Subject: [PATCH 2/4] fix(vlm): use 3D vision RoPE and suppress sentinels in MiniMax-M3-VL Address two review findings on the MiniMax-M3-VL port (PR #800). The vision tower reused the shared Qwen2-VL 2D (h, w) rotary embedding, but the reference MiniMaxVLVisionTransformer uses a genuine 3D (t, h, w) split: axis_dim = 2 * ((head_dim / 2) / 3 / 2) head dims per axis (26 for head_dim 80), rot_dim = 3 * axis_dim head dims rotated, and the remaining head_dim - rot_dim trailing dims pass through unrotated. The 2D reuse allocated the head dimensions across only h and w with a different frequency base, so every vision embedding was wrong even for images: with grid_t == 1 the temporal section is inert, but it still reserves its slice of the head dim, so the split cannot collapse to the 2D form. rot_pos_emb now builds the temporal ids and emits the concatenated t, h, w frequency sections, and VisionAttention applies a partial rotation to the leading rot_dim dims while leaving the trailing tail untouched. output_suppressed_token_ids masked only the image placeholder id, leaving video_token_id, vision_start_token_id, and vision_end_token_id samplable. Those are input-alignment sentinels that must never appear in decoded text, so they are now all suppressed, matching the ernie4_5_moe_vl and hunyuan_vl convention (video reuses the image vision_start/vision_end framing in MiniMax-M3-VL). Validation: cargo fmt clean, release lib build with --features cuda exit 0, and cargo test --release --features cuda --lib minimax_m3 -- --test-threads=1 passes 46/46. The 427B checkpoint cannot run on the development machine, so image parity still defers to a runtime check once a fitting quantized conversion exists. Refs #764 --- docs/supported-models.md | 2 +- src/vision/encoders/minimax_m3_vl.rs | 105 ++++++++++++++++++++++----- src/vision/minimax_m3_vl.rs | 12 ++- src/vision/minimax_m3_vl_tests.rs | 4 +- 4 files changed, 100 insertions(+), 23 deletions(-) diff --git a/docs/supported-models.md b/docs/supported-models.md index ee800802e..8b02c9249 100644 --- a/docs/supported-models.md +++ b/docs/supported-models.md @@ -95,7 +95,7 @@ Implemented VLM variants include: - ERNIE-4.5 MoE VL (`ernie4_5_moe_vl`): Baidu's vision-language MoE. A DFNRope ViT (linear patch embedding over 588-wide merge-window rows, 2D vision RoPE, `cu_seqlens`-packed attention, quick_gelu MLP) feeds a variable-resolution resampler (2x2 spatial fold, temporal pair fold with single-frame duplication, GELU MLP stacks, RMSNorm) whose rows replace the `<|IMAGE_PLACEHOLDER|>` tokens. The text decoder extends ERNIE-4.5 MoE with modality-split expert banks: separate text and multimodal routers and expert stacks selected per token by token type, with a correction bias that shifts expert selection but never the mixing weights, plus a fused shared-experts MLP. Position encoding is interleaved 3D MRoPE (`[T, H, W]` axes assigned per frequency index, adjacent-pair rotation) that degenerates to traditional RoPE for text. Validated against `mlx-community/ERNIE-4.5-VL-28B-A3B-Thinking-4bit` (28B-A3B; text and resampler 4-bit, vision tower bf16). Best for general image chat and grounded reasoning; the checkpoint is a thinking variant and emits reasoning before the answer. - Qwen3-Omni MoE (`qwen3_omni_moe`, thinker): Alibaba's omni-modal MoE. Stage 1 covers the thinker: text output conditioned on text, image, and audio inputs. The vision tower and MoE text decoder are the Qwen3-VL-MoE stack (DeepStack feature injection, interleaved MRoPE) reused unchanged; the new audio tower converts 16 kHz audio to a 128-bin log-mel spectrogram, downsamples it through three stride-2 convolutions (13 output frames per second of audio), and runs 32 windowed-attention encoder layers whose output rows scatter into the token stream exactly like image features. Audio arrives via `--audio file.wav` on the CLI (combinable with `--image`). Stage 2 adds speech output: `mlxcel generate --output-audio out.wav` runs the talker and code2wav after text generation and writes 24 kHz mono PCM16. The talker is a 20-layer Qwen3-MoE codec decoder conditioned on the projected thinker token embeddings of the chat-role segments; per frame it emits the first of 16 codebooks and a 5-layer code predictor fills in the residual 15, then the code2wav vocoder (causal pre-transformer, ConvNeXt upsampling, BigVGAN-style SnakeBeta decoder) renders 1920 samples per 12.5 Hz frame. `--speaker` selects the voice (ethan default; chelsie and aiden also ship in the released checkpoints). The speech stack loads lazily and only when requested, so text/vision use keeps its memory footprint; speech currently requires a text-only, chat-templated prompt. Validated against `mlx-community/Qwen3-Omni-30B-A3B-Instruct-4bit` (text and talker 4-bit; vision, audio tower, code predictor, and code2wav bf16). - Hunyuan-VL (`hunyuan_vl`, e.g. HunyuanOCR): Tencent's vision-language family. A ViT with a per-patch conv embedding, bilinearly interpolated learned position embeddings, and full attention over the packed patch sequence feeds a `perceive` merger: a stride-2 conv pair over the raster grid, a learned `image_newline` column, a linear to the decoder width, and learned `image_begin` / `image_end` rows. Per image that yields `mh * (mw + 1) + 2` feature rows, matching the prompt placeholder count exactly. The decoder is the Hunyuan dense stack (per-head Q/K RMSNorm after the rotation, DynamicNTK-alpha rope base) with XD-RoPE at prefill: 4D `[P, T, H, W]` position ids split across the frequency dims, degenerating to the standard rotation for text; decode uses sequential positions. Validated against `hadeseus/HunyuanOCR-mlx-4bit` (text 4-bit, vision bf16). Best for OCR: text spotting, document parsing, and grounded extraction. -- MiniMax-M3-VL (`minimax_m3_vl`): MiniMax's vision-language model on top of the MiniMax-M3 text backbone. A CLIP-style ViT (hidden 1280, 16 heads, 32 layers, patch 14, `pre_layrnorm`, LayerNorm + exact-GELU blocks with separate `q/k/v/out` projections) runs native-resolution packing: Qwen2-VL-style dynamic `smart_resize` to `patch_size * spatial_merge_size = 28`-aligned dimensions, `image_grid_thw` patchify, per-image `cu_seqlens` variable-length attention, and 2D (h, w) vision RoPE. A two-stage projector then maps features to the text width: a per-patch `multi_modal_projector` (`linear_1 -> GELU -> linear_2`, into `projection_dim` 6144) followed by a `patch_merge_mlp` that folds each `spatial_merge_size^2 = 4` adjacent patches (`linear_1` [6144, 24576] `-> GELU -> linear_2`). Each `]<]image[>[` placeholder expands to `grid_t * (h/2) * (w/2)` tokens (`vision_start` 200029, `vision_end` 200030), which the merged features replace LLaVA-style; the MiniMax-M3 hybrid dense/MoE decoder then runs its standard partial 1D RoPE. The vision tower runs in f32 (non-quantized in the checkpoint) while the text tower may be quantized in community exports. Image and multi-image inputs are wired end to end (CLI and server); video is out of scope for this port. The only public checkpoint (`MiniMaxAI/MiniMax-M3`, 427B) exceeds the development machine, so image Q&A parity is deferred to a runtime validation once a fitting quantized conversion exists; the merge gate is the unit tests plus the real nested-config parse. +- MiniMax-M3-VL (`minimax_m3_vl`): MiniMax's vision-language model on top of the MiniMax-M3 text backbone. A CLIP-style ViT (hidden 1280, 16 heads, 32 layers, patch 14, `pre_layrnorm`, LayerNorm + exact-GELU blocks with separate `q/k/v/out` projections) runs native-resolution packing: Qwen2-VL-style dynamic `smart_resize` to `patch_size * spatial_merge_size = 28`-aligned dimensions, `image_grid_thw` patchify, per-image `cu_seqlens` variable-length attention, and 3D (t, h, w) vision RoPE (temporal axis inert for images, trailing dims unrotated). A two-stage projector then maps features to the text width: a per-patch `multi_modal_projector` (`linear_1 -> GELU -> linear_2`, into `projection_dim` 6144) followed by a `patch_merge_mlp` that folds each `spatial_merge_size^2 = 4` adjacent patches (`linear_1` [6144, 24576] `-> GELU -> linear_2`). Each `]<]image[>[` placeholder expands to `grid_t * (h/2) * (w/2)` tokens (`vision_start` 200029, `vision_end` 200030), which the merged features replace LLaVA-style; the MiniMax-M3 hybrid dense/MoE decoder then runs its standard partial 1D RoPE. The vision tower runs in f32 (non-quantized in the checkpoint) while the text tower may be quantized in community exports. Image and multi-image inputs are wired end to end (CLI and server); video is out of scope for this port. The only public checkpoint (`MiniMaxAI/MiniMax-M3`, 427B) exceeds the development machine, so image Q&A parity is deferred to a runtime validation once a fitting quantized conversion exists; the merge gate is the unit tests plus the real nested-config parse. - FastVLM (`llava_qwen2` / `fastvlm`): Apple's low-latency VLM. A FastViTHD hybrid encoder runs entirely on channels-last maps: a conv stem, three RepMixer stages (depthwise token mixing plus a BatchNorm ConvFFN), two attention stages with channel LayerNorm and `head_dim` 32 self-attention, inter-stage large-kernel PatchEmbed downsamples, RepCPE position encoders, and a squeeze-excite `conv_exp` head. Each 1024x1024 pad-to-square image becomes a `(16, 16, 3072)` map flattened to 256 tokens, projected to the Qwen2 decoder width by an `mlp2x_gelu` MLP. The `` placeholder is the fixed `-200` sentinel (not a vocabulary token); the runtime splices one sentinel per image, expands it to 256 tokens, and scatters the image embeddings (LLaVA merge). The text decoder is stock Qwen2, reused unchanged. The loader accepts both the genuine (`apple/FastVLM-0.5B`) and converted (`mlx-community/FastVLM-0.5B-bf16`) weight layouts. Best for fast image description and grounded chat. - GLM-OCR (`glm_ocr`): document-OCR sibling of GLM-4V. A 24-block ViT (3D patch embedding, per-head q/k RMSNorm on the packed `cu_seqlens` attention, 2D vision RoPE, Conv2d spatial downsample, SwiGLU patch merger) feeds a 16-layer GLM-4 text decoder driven by full-width even/odd MRoPE (`rope_parameters` with `mrope_section [16, 24, 24]`, `partial_rotary_factor 1.0`). The tower has no learned position embedding or post-conv norm, and the loader drops the next-n prediction (MTP) layer. Patches are reordered from the processor's raster order into spatial-merge-window order so the rotary, downsample, and merged-token grid stay spatially aligned (OCR reads scrambled patches wrong). Best for plain text, tables, and formula recognition. - Youtu-VL diff --git a/src/vision/encoders/minimax_m3_vl.rs b/src/vision/encoders/minimax_m3_vl.rs index 81b643f86..6e8b623c7 100644 --- a/src/vision/encoders/minimax_m3_vl.rs +++ b/src/vision/encoders/minimax_m3_vl.rs @@ -28,11 +28,19 @@ //! a `patch_merge_mlp` that folds `spatial_merge_size^2 = 4` adjacent patches //! (`linear_1` [6144, 24576] -> GELU -> `linear_2`) into the text hidden size. //! -//! The tower reuses the shared Qwen2-VL vision helpers for the 2D (h, w) rotary -//! embedding and per-image `cu_seqlens` variable-length attention, matching the -//! `image_grid_thw` packing emitted by the processor. Video (`grid_t > 1`) is -//! out of scope for this port; the temporal axis of the rope reduces to the -//! well-tested (h, w) form for images (`grid_t == 1`). +//! The tower reuses the shared Qwen2-VL `cu_seqlens` variable-length attention +//! and the frequency-table / `apply_rotary_pos_emb_vision` helpers, but drives +//! them with a genuine 3D (t, h, w) vision RoPE that matches the reference +//! `MiniMaxVLVisionTransformer`. The head dimension is split into three equal +//! axis sections (t, h, w) plus a trailing unrotated tail: `axis_dim = +//! 2 * ((head_dim / 2) / 3 / 2)` head dims per axis, `rot_dim = 3 * axis_dim` +//! head dims rotated, and the remaining `head_dim - rot_dim` trailing dims pass +//! through untouched (2 dims for head_dim 80, 4 for the reduced test head_dim +//! 64). For images (`grid_t == 1`) the temporal section is inert (all-zero t +//! ids), but it still reserves its slice of the head dimension, so the split +//! cannot collapse to the 2D (h, w) form. Video (`grid_t > 1`) is out of scope +//! for this port. The `image_grid_thw` packing emitted by the processor drives +//! both the position ids and the per-image `cu_seqlens`. //! //! The whole tower runs in f32: the checkpoint stores the vision weights as //! f32/bf16 non-quantized, and running the tower uniformly in f32 avoids @@ -151,6 +159,26 @@ impl MiniMaxM3VisionConfig { } } +/// Per-axis rotary width of the 3D vision RoPE, shared by the t/h/w sections. +/// +/// Matches the reference `MiniMaxVLVisionTransformer`: `rope_dims` rounds the +/// head dim down to an even width, then each of the three axis sections gets an +/// even slice `axis_dim = 2 * ((rope_dims / 3) / 2)` (integer division). The +/// frequency table therefore holds `axis_dim / 2` entries per axis. head_dim 80 +/// -> axis_dim 26; head_dim 64 -> axis_dim 20. +fn rope_axis_dim(head_dim: i32) -> i32 { + let rope_dims = 2 * (head_dim / 2); + 2 * ((rope_dims / 3) / 2) +} + +/// Number of head dims actually rotated by the 3D vision RoPE (`3 * axis_dim`). +/// The remaining `head_dim - rot_dim` trailing dims pass through unrotated. +/// head_dim 80 -> rot_dim 78 (2 pass-through); head_dim 64 -> rot_dim 60 (4 +/// pass-through). +fn rope_rot_dim(head_dim: i32) -> i32 { + 3 * rope_axis_dim(head_dim) +} + // ============================================================================ // Plain f32 linear / layernorm helpers // ============================================================================ @@ -262,6 +290,9 @@ struct VisionAttention { out_proj: VisionLinear, num_heads: i32, head_dim: i32, + /// Leading head dims rotated by the 3D vision RoPE (`3 * axis_dim`); the + /// trailing `head_dim - rot_dim` dims pass through unrotated. + rot_dim: i32, scale: f32, } @@ -279,6 +310,7 @@ impl VisionAttention { out_proj: VisionLinear::load(weights, &format!("{}.out_proj", prefix))?, num_heads: config.num_attention_heads as i32, head_dim, + rot_dim: rope_rot_dim(head_dim), scale: (head_dim as f32).powf(-0.5), }) } @@ -302,8 +334,25 @@ impl VisionAttention { let k = reshape_heads(self.k_proj.forward(x)); let v = reshape_heads(self.v_proj.forward(x)); - let q = apply_rotary_pos_emb_vision(&q, rotary_pos_emb); - let k = apply_rotary_pos_emb_vision(&k, rotary_pos_emb); + // Partial 3D vision RoPE: rotate only the leading `rot_dim` head dims + // (the concatenated t/h/w axis sections) and pass the trailing + // `head_dim - rot_dim` dims through unrotated. v is not rotated. + let apply_rope = |t: &MlxArray| -> UniquePtr { + if self.rot_dim >= self.head_dim { + return apply_rotary_pos_emb_vision(t, rotary_pos_emb); + } + let rot = + mlxcel_core::slice(t, &[0, 0, 0], &[seq_length, self.num_heads, self.rot_dim]); + let pass = mlxcel_core::slice( + t, + &[0, 0, self.rot_dim], + &[seq_length, self.num_heads, self.head_dim], + ); + let rot = apply_rotary_pos_emb_vision(&rot, rotary_pos_emb); + mlxcel_core::concatenate(&rot, &pass, 2) + }; + let q = apply_rope(&q); + let k = apply_rope(&k); // [seq, heads, head_dim] -> [1, heads, seq, head_dim] let to_bhsd = |t: &MlxArray| { @@ -505,9 +554,12 @@ impl MiniMaxM3VisionEncoder { config.layer_norm_eps, )?; - // The vision rope table covers head_dim/2 frequencies (h and w halves), - // matching the shared Qwen2-VL helper. - let rotary_pos_emb = VisionRotaryEmbedding::new(config.head_dim() / 2); + // The 3D vision RoPE splits the head dim into three equal (t, h, w) axis + // sections plus an unrotated tail. All three axes share one frequency + // table of `axis_dim / 2` entries (their widths are identical), so a + // single `VisionRotaryEmbedding::new(axis_dim)` covers t, h, and w. + let axis_dim = rope_axis_dim(config.head_dim() as i32); + let rotary_pos_emb = VisionRotaryEmbedding::new(axis_dim as usize); let mut layers = Vec::with_capacity(config.num_hidden_layers); for i in 0..config.num_hidden_layers { @@ -537,15 +589,19 @@ impl MiniMaxM3VisionEncoder { }) } - /// 2D (h, w) rotary position embeddings, merge-grouped to match the - /// processor patch order. Identical construction to the Qwen2-VL tower. + /// 3D (t, h, w) rotary position embeddings, merge-grouped to match the + /// processor patch order. The (h, w) grouping is identical to the Qwen2-VL + /// tower; the temporal ids repeat each frame index over its `h * w` tokens + /// and are all-zero for images (`grid_t == 1`). Emits + /// `[total_tokens, 3 * (axis_dim / 2)]` (the t, h, w frequency sections + /// concatenated) for the partial rotation in `VisionAttention::forward`. fn rot_pos_emb(&self, grid_thw: &[(i32, i32, i32)]) -> UniquePtr { let mut all_pos_ids: Vec> = Vec::new(); let mut max_grid_dim: i32 = 0; let merge = self.spatial_merge_size as i32; for &(t, h, w) in grid_thw { - max_grid_dim = max_grid_dim.max(h).max(w); + max_grid_dim = max_grid_dim.max(t).max(h).max(w); let h_arange = mlxcel_core::arange_i32(0, h, 1); let h_col = mlxcel_core::reshape(&h_arange, &[h, 1]); @@ -561,9 +617,20 @@ impl MiniMaxM3VisionEncoder { let wpos = mlxcel_core::transpose_axes(&wpos, &[0, 2, 1, 3]); let wpos = mlxcel_core::flatten(&wpos); - let stacked = mlxcel_core::stack_owned(&[hpos, wpos], -1); - let tiled = mlxcel_core::tile(&stacked, &[t, 1]); - all_pos_ids.push(tiled); + // Stack the spatial (h, w) ids and tile them across the t frames. + let hw = mlxcel_core::stack_owned(&[hpos, wpos], -1); + let hw = mlxcel_core::tile(&hw, &[t, 1]); + + // Temporal ids: each frame index repeated over its h * w tokens + // (all-zero for images, where grid_t == 1). + let t_arange = mlxcel_core::arange_i32(0, t, 1); + let t_col = mlxcel_core::reshape(&t_arange, &[t, 1]); + let tpos = mlxcel_core::repeat(&t_col, h * w, 1); + let tpos = mlxcel_core::reshape(&tpos, &[t * h * w, 1]); + + // Concatenate the (t, h, w) columns in that order -> [t*h*w, 3]. + let stacked = mlxcel_core::concatenate(&tpos, &hw, 1); + all_pos_ids.push(stacked); } let pos_ids = if all_pos_ids.len() == 1 { @@ -579,8 +646,10 @@ impl MiniMaxM3VisionEncoder { let total_tokens = total_shape[0]; let freq_shape = mlxcel_core::array_shape(&all_freqs); let half_dim = freq_shape[1]; - let all_freqs = mlxcel_core::reshape(&all_freqs, &[total_tokens, 2, half_dim]); - mlxcel_core::reshape(&all_freqs, &[total_tokens, 2 * half_dim]) + // [total*3, axis_dim/2] -> [total, 3, axis_dim/2] -> [total, 3*axis_dim/2], + // concatenating the t, h, w frequency sections in that order. + let all_freqs = mlxcel_core::reshape(&all_freqs, &[total_tokens, 3, half_dim]); + mlxcel_core::reshape(&all_freqs, &[total_tokens, 3 * half_dim]) } /// Per-image `cu_seqlens` (full attention over each image's patches): diff --git a/src/vision/minimax_m3_vl.rs b/src/vision/minimax_m3_vl.rs index 2df008033..552708d17 100644 --- a/src/vision/minimax_m3_vl.rs +++ b/src/vision/minimax_m3_vl.rs @@ -122,8 +122,16 @@ impl LanguageModel for MiniMaxM3VlModel { } fn output_suppressed_token_ids(&self) -> Vec { - // The image placeholder id must never be sampled during decode. - vec![self.image_token_id] + // Image/video placeholders and their vision framing markers are + // input-alignment ids and must never be sampled during decode. Video + // reuses the image vision_start/vision_end framing in MiniMax-M3-VL, so + // there are no separate video framing tokens. + vec![ + self.image_token_id, + self.video_token_id, + self.vision_start_token_id, + self.vision_end_token_id, + ] } } diff --git a/src/vision/minimax_m3_vl_tests.rs b/src/vision/minimax_m3_vl_tests.rs index 9feccc5c5..c6686c723 100644 --- a/src/vision/minimax_m3_vl_tests.rs +++ b/src/vision/minimax_m3_vl_tests.rs @@ -19,7 +19,7 @@ //! (`vision_config` and `text_config`), the sanitizer's vision/projector skip //! against verbatim checkpoint keys, the projector fold ordering //! (`merge^2 * projection_dim`), a tiny synthetic tower forward (patch embed -> -//! pre_layrnorm -> CLIP layers with `cu_seqlens` + 2D vision RoPE -> two-stage +//! pre_layrnorm -> CLIP layers with `cu_seqlens` + 3D vision RoPE -> two-stage //! projector), and the placeholder-count invariant against the shared Qwen-VL //! insertion helper. Run serially (`--test-threads=1`); the MLX ops touch the //! device. @@ -144,7 +144,7 @@ fn tiny_tower_weights(cfg: &MiniMaxM3VisionConfig) -> WeightMap { fn tower_loads_verbatim_keys_and_emits_merged_token_shape() { // A single 2x2-patch-grid image (grid (1,2,2)) -> 4 patches -> 1 merged // token after the merge^2 fold. Exercises patch embed, pre_layrnorm, the - // CLIP layer (cu_seqlens attention + 2D vision RoPE), and both projector + // CLIP layer (cu_seqlens attention + 3D vision RoPE), and both projector // stages, and proves the tower resolves every verbatim key (including the // `pre_layrnorm` spelling). let cfg = tiny_vision_config(); From 03bdcf53af5f06481cf3fdaa984a85b559a584fc Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Thu, 16 Jul 2026 22:25:17 +0900 Subject: [PATCH 3/4] test(vlm): pin 3D vision-RoPE axis/rot-dim math for MiniMax-M3-VL The HIGH-severity review fix (3D vision RoPE, PR #800) rewrote rot_pos_emb and added the rope_axis_dim/rope_rot_dim helpers, but only updated existing docstrings; no assertion pinned the new axis/rot-dim arithmetic or the emitted frequency-table shape. Add one test that checks rope_axis_dim/rope_rot_dim against the reference values (26/78 for head_dim 80, 20/60 for head_dim 64) and, using the existing tiny-tower fixture, checks that rot_pos_emb emits [total_tokens, rot_dim / 2] with the leading (temporal) section exactly zero when grid_t == 1. rope_axis_dim, rope_rot_dim, and rot_pos_emb widen from private to pub(crate) so the existing vision::minimax_m3_vl_tests module (a sibling of vision::encoders::minimax_m3_vl, not a descendant) can call them directly instead of duplicating the tower-building fixture in a second test module. Validation: cargo test --release --features cuda --lib minimax_m3 -- --test-threads=1, 47/47 passed (10 pre-existing MiniMax-M3-VL tests plus this one). Refs #764 --- src/vision/encoders/minimax_m3_vl.rs | 11 ++++-- src/vision/minimax_m3_vl_tests.rs | 59 ++++++++++++++++++++++++++-- 2 files changed, 63 insertions(+), 7 deletions(-) diff --git a/src/vision/encoders/minimax_m3_vl.rs b/src/vision/encoders/minimax_m3_vl.rs index 6e8b623c7..c76d29067 100644 --- a/src/vision/encoders/minimax_m3_vl.rs +++ b/src/vision/encoders/minimax_m3_vl.rs @@ -166,7 +166,7 @@ impl MiniMaxM3VisionConfig { /// even slice `axis_dim = 2 * ((rope_dims / 3) / 2)` (integer division). The /// frequency table therefore holds `axis_dim / 2` entries per axis. head_dim 80 /// -> axis_dim 26; head_dim 64 -> axis_dim 20. -fn rope_axis_dim(head_dim: i32) -> i32 { +pub(crate) fn rope_axis_dim(head_dim: i32) -> i32 { let rope_dims = 2 * (head_dim / 2); 2 * ((rope_dims / 3) / 2) } @@ -175,7 +175,7 @@ fn rope_axis_dim(head_dim: i32) -> i32 { /// The remaining `head_dim - rot_dim` trailing dims pass through unrotated. /// head_dim 80 -> rot_dim 78 (2 pass-through); head_dim 64 -> rot_dim 60 (4 /// pass-through). -fn rope_rot_dim(head_dim: i32) -> i32 { +pub(crate) fn rope_rot_dim(head_dim: i32) -> i32 { 3 * rope_axis_dim(head_dim) } @@ -595,7 +595,12 @@ impl MiniMaxM3VisionEncoder { /// and are all-zero for images (`grid_t == 1`). Emits /// `[total_tokens, 3 * (axis_dim / 2)]` (the t, h, w frequency sections /// concatenated) for the partial rotation in `VisionAttention::forward`. - fn rot_pos_emb(&self, grid_thw: &[(i32, i32, i32)]) -> UniquePtr { + /// + /// `pub(crate)` (rather than private) so the unit tests in + /// `vision::minimax_m3_vl_tests` can pin the emitted shape and the + /// all-zero temporal section for `grid_t == 1` directly, without + /// duplicating this method's logic. + pub(crate) fn rot_pos_emb(&self, grid_thw: &[(i32, i32, i32)]) -> UniquePtr { let mut all_pos_ids: Vec> = Vec::new(); let mut max_grid_dim: i32 = 0; let merge = self.spatial_merge_size as i32; diff --git a/src/vision/minimax_m3_vl_tests.rs b/src/vision/minimax_m3_vl_tests.rs index c6686c723..785976a35 100644 --- a/src/vision/minimax_m3_vl_tests.rs +++ b/src/vision/minimax_m3_vl_tests.rs @@ -20,12 +20,15 @@ //! against verbatim checkpoint keys, the projector fold ordering //! (`merge^2 * projection_dim`), a tiny synthetic tower forward (patch embed -> //! pre_layrnorm -> CLIP layers with `cu_seqlens` + 3D vision RoPE -> two-stage -//! projector), and the placeholder-count invariant against the shared Qwen-VL -//! insertion helper. Run serially (`--test-threads=1`); the MLX ops touch the -//! device. +//! projector), the 3D vision RoPE axis/rot-dim arithmetic and the emitted +//! `rot_pos_emb` shape with its all-zero temporal section for `grid_t == 1`, +//! and the placeholder-count invariant against the shared Qwen-VL insertion +//! helper. Run serially (`--test-threads=1`); the MLX ops touch the device. use crate::models::minimax_m3; -use crate::vision::encoders::minimax_m3_vl::{MiniMaxM3VisionConfig, MiniMaxM3VisionEncoder}; +use crate::vision::encoders::minimax_m3_vl::{ + MiniMaxM3VisionConfig, MiniMaxM3VisionEncoder, rope_axis_dim, rope_rot_dim, +}; use mlxcel_core::MlxArray; use mlxcel_core::weights::WeightMap; @@ -173,6 +176,54 @@ fn tower_loads_verbatim_keys_and_emits_merged_token_shape() { ); } +#[test] +fn rope_axis_and_rot_dim_match_reference_arithmetic() { + // Pins the 3D vision RoPE split introduced by the review fix: axis_dim = + // 2 * ((head_dim / 2) / 3 / 2) head dims per (t, h, w) axis, rot_dim = + // 3 * axis_dim rotated, with head_dim - rot_dim trailing dims untouched. + // head_dim 80 is the real MiniMaxAI/MiniMax-M3 vision head_dim (hidden + // 1280 / 16 heads); head_dim 64 is the reduced tiny test config's + // (hidden 128 / 2 heads). + assert_eq!(rope_axis_dim(80), 26); + assert_eq!(rope_rot_dim(80), 78); + assert_eq!(rope_axis_dim(64), 20); + assert_eq!(rope_rot_dim(64), 60); + + // rot_pos_emb must emit [total_tokens, rot_dim / 2] (the concatenated t, h, + // w frequency sections), and for grid_t == 1 (images) the temporal ids are + // all-zero, so the leading axis_dim/2 columns (the t section) must be + // exactly zero even though the axis still reserves its slice of head_dim. + let cfg = tiny_vision_config(); + let weights = tiny_tower_weights(&cfg); + let encoder = MiniMaxM3VisionEncoder::from_weights(&weights, &cfg).expect("tiny tower loads"); + + let head_dim = cfg.head_dim() as i32; + let axis_dim = rope_axis_dim(head_dim); + let rot_dim = rope_rot_dim(head_dim); + assert_eq!(head_dim, 64); + assert_eq!(axis_dim, 20); + assert_eq!(rot_dim, 60); + + let grid = vec![(1i32, 2i32, 2i32)]; // grid_t == 1: temporal section is inert. + let (t, h, w) = grid[0]; + let total_tokens = t * h * w; + + let freqs = encoder.rot_pos_emb(&grid); + mlxcel_core::eval(&freqs); + assert_eq!( + mlxcel_core::array_shape(&freqs), + vec![total_tokens, rot_dim / 2] + ); + + let half_axis = axis_dim / 2; + let t_section = mlxcel_core::slice(&freqs, &[0, 0], &[total_tokens, half_axis]); + assert_eq!( + reduce_max_abs(&t_section), + 0.0, + "t-section of rot_pos_emb must be exactly zero for grid_t == 1" + ); +} + #[test] fn projector_fold_concatenates_four_adjacent_patches_in_order() { // The patch-merge fold is a row-major reshape [P, D] -> [P/4, 4*D]. The four From 90e8b7e9224ef7de33706e58ba077f917a7cb624 Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Thu, 16 Jul 2026 22:25:38 +0900 Subject: [PATCH 4/4] fix(vlm): trim excessive float precision in MiniMax-M3-VL CLIP std cargo clippy --release --features cuda --lib -p mlxcel -- -D warnings (the only feature set this CUDA-only machine can build) fails on two of the three CLIP std constants added for the vision processor: 0.261_302_58 and 0.275_777_11 both carry more decimal digits than an f32 needs to round-trip, tripping clippy::excessive_precision under -D warnings. Truncate both to the shortest literal that parses to the identical f32 bit pattern (verified via Rust's f32 parser), so the constants are unchanged in value. Validation: cargo clippy --release --features cuda --lib -p mlxcel -- -D warnings now exits clean, and cargo test --release --features cuda --lib minimax_m3 -- --test-threads=1 still passes 47/47. Refs #764 --- src/vision/processors/minimax_m3.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/vision/processors/minimax_m3.rs b/src/vision/processors/minimax_m3.rs index 78e7c6977..bd78e656a 100644 --- a/src/vision/processors/minimax_m3.rs +++ b/src/vision/processors/minimax_m3.rs @@ -57,7 +57,7 @@ impl Default for MiniMaxM3Processor { min_pixels: 4 * 28 * 28, // 3136 max_pixels: 451_584, // 672 * 672 mean: [0.481_454_66, 0.457_827_5, 0.408_210_73], - std: [0.268_629_54, 0.261_302_58, 0.275_777_11], + std: [0.268_629_54, 0.261_302_6, 0.275_777_1], } } }