From c245d35ff467bb3e9a73fcb3c4b02e6c7a3d2964 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 27 Apr 2026 21:26:51 +0800 Subject: [PATCH] [Model] Add MiMo-V2.5 support (#40967) Signed-off-by: Jee Jee Li Signed-off-by: Isotr0py Signed-off-by: Isotr0py Signed-off-by: zjy0516 Co-authored-by: Jee Jee Li Co-authored-by: zjy0516 Co-authored-by: zjy0516 Co-authored-by: yasong Co-authored-by: Claude Opus 4.6 Co-authored-by: Copilot --- docs/models/supported_models.md | 2 + tests/models/registry.py | 18 + vllm/config/model.py | 2 +- vllm/config/speculative.py | 43 + vllm/model_executor/layers/linear.py | 7 +- vllm/model_executor/models/mimo_audio.py | 1389 +++++++++++++++ .../models/{mimo_v2_flash.py => mimo_v2.py} | 24 +- vllm/model_executor/models/mimo_v2_mtp.py | 373 +++++ vllm/model_executor/models/mimo_v2_omni.py | 1488 +++++++++++++++++ vllm/model_executor/models/registry.py | 6 +- .../configs/mimo_v2_omni.py | 65 + .../model_arch_config_convertor.py | 35 + .../transformers_utils/processors/__init__.py | 2 + .../processors/mimo_v2_omni.py | 1285 ++++++++++++++ vllm/v1/core/kv_cache_utils.py | 2 + vllm/v1/spec_decode/llm_base_proposer.py | 1 + 16 files changed, 4737 insertions(+), 5 deletions(-) create mode 100644 vllm/model_executor/models/mimo_audio.py rename vllm/model_executor/models/{mimo_v2_flash.py => mimo_v2.py} (96%) create mode 100644 vllm/model_executor/models/mimo_v2_mtp.py create mode 100644 vllm/model_executor/models/mimo_v2_omni.py create mode 100644 vllm/transformers_utils/configs/mimo_v2_omni.py create mode 100644 vllm/transformers_utils/processors/mimo_v2_omni.py diff --git a/docs/models/supported_models.md b/docs/models/supported_models.md index 7c87359afb2..e12b83e25d4 100644 --- a/docs/models/supported_models.md +++ b/docs/models/supported_models.md @@ -439,6 +439,7 @@ th { | `Mamba2ForCausalLM` | Mamba2 | `mistralai/Mamba-Codestral-7B-v0.1`, etc. | | ✅︎ | | `MiMoForCausalLM` | MiMo | `XiaomiMiMo/MiMo-7B-RL`, etc. | ✅︎ | ✅︎ | | `MiMoV2FlashForCausalLM` | MiMoV2Flash | `XiaomiMiMo/MiMo-V2-Flash`, etc. | | ✅︎ | +| `MiMoV2ProForCausalLM` | MiMoV2Pro | `XiaomiMiMo/MiMo-V2.5-Pro`, etc. | | ✅︎ | | `MiniCPMForCausalLM` | MiniCPM | `openbmb/MiniCPM-2B-sft-bf16`, `openbmb/MiniCPM-2B-dpo-bf16`, `openbmb/MiniCPM-S-1B-sft`, etc. | ✅︎ | ✅︎ | | `MiniCPM3ForCausalLM` | MiniCPM3 | `openbmb/MiniCPM3-4B`, etc. | ✅︎ | ✅︎ | | `MiniMaxForCausalLM` | MiniMax-Text | `MiniMaxAI/MiniMax-Text-01-hf`, etc. | | | @@ -590,6 +591,7 @@ These models primarily accept the [`LLM.generate`](./generative_models.md#llmgen | `LlavaNextVideoForConditionalGeneration` | LLaVA-NeXT-Video | T + V | `llava-hf/LLaVA-NeXT-Video-7B-hf`, etc. | | ✅︎ | | `LlavaOnevisionForConditionalGeneration` | LLaVA-Onevision | T + I+ + V+ | `llava-hf/llava-onevision-qwen2-7b-ov-hf`, `llava-hf/llava-onevision-qwen2-0.5b-ov-hf`, etc. | | ✅︎ | | `MiDashengLMModel` | MiDashengLM | T + A+ | `mispeech/midashenglm-7b` | | ✅︎ | +| `MiMoV2OmniForCausalLM` | MiMo-V2.5-Omni | T + IE+ + VE+ + A+ | `XiaomiMiMo/MiMo-V2.5-Omni` | | ✅︎ | | `MiniCPMO` | MiniCPM-O | T + IE+ + VE+ + AE+ | `openbmb/MiniCPM-o-2_6`, etc. | ✅︎ | ✅︎ | | `MiniCPMV` | MiniCPM-V | T + IE+ + VE+ | `openbmb/MiniCPM-V-2` (see note), `openbmb/MiniCPM-Llama3-V-2_5`, `openbmb/MiniCPM-V-2_6`, `openbmb/MiniCPM-V-4`, `openbmb/MiniCPM-V-4_5`, etc. | ✅︎ | | | `MiniMaxVL01ForConditionalGeneration` | MiniMax-VL | T + IE+ | `MiniMaxAI/MiniMax-VL-01`, etc. | | ✅︎ | diff --git a/tests/models/registry.py b/tests/models/registry.py index 05097fc77f7..e499d34d05f 100644 --- a/tests/models/registry.py +++ b/tests/models/registry.py @@ -594,6 +594,9 @@ _TEXT_GENERATION_EXAMPLE_MODELS = { "MiMoV2FlashForCausalLM": _HfExamplesInfo( "XiaomiMiMo/MiMo-V2-Flash", trust_remote_code=True ), + "MiMoV2ProForCausalLM": _HfExamplesInfo( + "XiaomiMiMo/MiMo-V2.5-Pro", trust_remote_code=True, is_available_online=False + ), "Dots1ForCausalLM": _HfExamplesInfo("rednote-hilab/dots.llm1.inst"), } @@ -1069,6 +1072,9 @@ _MULTIMODAL_EXAMPLE_MODELS = { "MiDashengLMModel": _HfExamplesInfo( "mispeech/midashenglm-7b", trust_remote_code=True ), + "MiMoV2OmniForCausalLM": _HfExamplesInfo( + "XiaomiMiMo/MiMo-V2.5-Omni", trust_remote_code=True, is_available_online=False + ), "MiniCPMO": _HfExamplesInfo( "openbmb/MiniCPM-o-2_6", trust_remote_code=True, @@ -1552,6 +1558,18 @@ _SPECULATIVE_DECODING_EXAMPLE_MODELS = { trust_remote_code=True, speculative_model="XiaomiMiMo/MiMo-7B-RL", ), + "MiMoV2MTPModel": _HfExamplesInfo( + "XiaomiMiMo/MiMo-V2.5-Pro", + trust_remote_code=True, + speculative_model="XiaomiMiMo/MiMo-V2.5-Pro", + is_available_online=False, + ), + "MiMoV2OmniMTPModel": _HfExamplesInfo( + "XiaomiMiMo/MiMo-V2.5-Omni", + trust_remote_code=True, + speculative_model="XiaomiMiMo/MiMo-V2.5-Omni", + is_available_online=False, + ), "NemotronHMTPModel": _HfExamplesInfo( "nvidia/Nemotron-Super-Placeholder", speculative_model="nvidia/Nemotron-Super-Placeholder", diff --git a/vllm/config/model.py b/vllm/config/model.py index b032302c04c..470e8091cc2 100644 --- a/vllm/config/model.py +++ b/vllm/config/model.py @@ -521,6 +521,7 @@ class ModelConfig: if dict_overrides: self._apply_dict_overrides(hf_config, dict_overrides) self.hf_text_config = get_hf_text_config(self.hf_config) + self.model_arch_config = self.get_model_arch_config() self.attention_chunk_size = getattr( self.hf_text_config, "attention_chunk_size", None ) @@ -528,7 +529,6 @@ class ModelConfig: self.hf_image_processor_config = get_hf_image_processor_config( self.model, hf_token=self.hf_token, revision=self.revision ) - self.model_arch_config = self.get_model_arch_config() architectures = self.architectures registry = self.registry diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index 54e9e7c0187..007fe2c8c36 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -34,6 +34,7 @@ logger = init_logger(__name__) MTPModelTypes = Literal[ "deepseek_mtp", "mimo_mtp", + "mimo_v2_mtp", "glm4_moe_mtp", "glm4_moe_lite_mtp", "glm_ocr_mtp", @@ -332,6 +333,48 @@ class SpeculativeConfig: } ) + if (arch := hf_config.architectures[0]) in ( + "MiMoV2ProForCausalLM", + "MiMoV2OmniForCausalLM", + ): + from vllm.model_executor.models.mimo_v2_mtp import ( + _MIMO_V2_PRO_NUM_MTP_LAYERS, + ) + + mtp_arch_maps = { + "MiMoV2ProForCausalLM": "MiMoV2MTPModel", + "MiMoV2OmniForCausalLM": "MiMoV2OmniMTPModel", + } + + hf_config.model_type = "mimo_v2_mtp" + # vLLM currently supports only the first MiMo-V2 MTP layer. + n_predict = _MIMO_V2_PRO_NUM_MTP_LAYERS + hf_config.update( + { + "num_hidden_layers": 0, + "n_predict": n_predict, + "num_nextn_predict_layers": n_predict, + "architectures": [mtp_arch_maps[arch]], + } + ) + + if hf_config.architectures[0] == "MiMoV2FlashForCausalLM": + from vllm.model_executor.models.mimo_v2_mtp import ( + _MIMO_V2_FLASH_NUM_MTP_LAYERS, + ) + + hf_config.model_type = "mimo_v2_mtp" + # vLLM currently supports only the first MiMo-V2 MTP layer. + n_predict = _MIMO_V2_FLASH_NUM_MTP_LAYERS + hf_config.update( + { + "num_hidden_layers": 0, + "n_predict": n_predict, + "num_nextn_predict_layers": n_predict, + "architectures": ["MiMoV2MTPModel"], + } + ) + if hf_config.architectures[0] == "Glm4MoeForCausalLM": hf_config.model_type = "glm4_moe_mtp" n_predict = getattr(hf_config, "num_nextn_predict_layers", None) diff --git a/vllm/model_executor/layers/linear.py b/vllm/model_executor/layers/linear.py index 241b32656e5..6a4c1f3c47e 100644 --- a/vllm/model_executor/layers/linear.py +++ b/vllm/model_executor/layers/linear.py @@ -1120,7 +1120,12 @@ class QKVParallelLinear(ColumnParallelLinear): # Special case for Quantization. # If quantized, we need to adjust the offset and size to account # for the packing. - if ( + if isinstance(param, BlockQuantScaleParameter): + weight_block_size = getattr(self, "weight_block_size", None) + shard_size, shard_offset = adjust_block_scale_shard( + weight_block_size, shard_size, shard_offset + ) + elif ( isinstance(param, (PackedColumnParameter, PackedvLLMParameter)) and param.packed_dim == param.output_dim ): diff --git a/vllm/model_executor/models/mimo_audio.py b/vllm/model_executor/models/mimo_audio.py new file mode 100644 index 00000000000..91d46b1ceac --- /dev/null +++ b/vllm/model_executor/models/mimo_audio.py @@ -0,0 +1,1389 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""MiMo audio: tokenizer, encoding utilities, and audio encoder. + +Ported from SGLang's mimo_audio.py. +Audio tokenizer adapted from https://github.com/XiaomiMiMo/MiMo-Audio-Tokenizer.git +""" + +import dataclasses +import json +import logging +import math +import os +import typing as tp +from dataclasses import dataclass +from functools import wraps + +import torch +import torch.distributed as dist +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange, repeat +from transformers.activations import ACT2FN +from transformers.configuration_utils import PretrainedConfig +from transformers.modeling_utils import PreTrainedModel +from transformers.models.qwen2.configuration_qwen2 import Qwen2Config +from transformers.models.qwen2.modeling_qwen2 import Qwen2Model + +logger = logging.getLogger(__name__) + + +# --------------------------------------------------------------------------- +# Vector quantization (from MiMo-Audio-Tokenizer) +# --------------------------------------------------------------------------- + + +def _vq_default(val: tp.Any, d: tp.Any) -> tp.Any: + return val if val is not None else d + + +def _ema_inplace(moving_avg, new, decay: float): + if dist.is_initialized(): + dist.all_reduce(new, op=dist.ReduceOp.SUM) + moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay)) + + +def _laplace_smoothing(x, n_categories: int, epsilon: float = 1e-5): + return (x + epsilon) / (x.sum() + n_categories * epsilon) + + +def _uniform_init(*shape: int): + t = torch.empty(shape) + nn.init.kaiming_uniform_(t) + return t + + +def _sample_vectors(samples, num: int): + num_samples, device = samples.shape[0], samples.device + + if num_samples >= num: + indices = torch.randperm(num_samples, device=device)[:num] + else: + indices = torch.randint(0, num_samples, (num,), device=device) + + selected_samples = samples[indices] + + if dist.is_initialized(): + dist.broadcast(selected_samples, src=0) + + return selected_samples + + +def _kmeans(samples, num_clusters: int, num_iters: int = 10): + dim, dtype = samples.shape[-1], samples.dtype + + means = _sample_vectors(samples, num_clusters) + + for _ in range(num_iters): + dists = -( + samples.pow(2).sum(1, keepdim=True) + - 2 * samples @ means.t() + + means.t().pow(2).sum(0, keepdim=True) + ) + + buckets = dists.max(dim=-1).indices + bins = torch.bincount(buckets, minlength=num_clusters) + + new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype) + new_means = new_means.scatter_add_( + 0, repeat(buckets, "n -> n d", d=dim), samples + ) + + if dist.is_initialized(): + dist.all_reduce(bins, op=dist.ReduceOp.SUM) + dist.all_reduce(new_means, op=dist.ReduceOp.SUM) + + zero_mask = bins == 0 + bins_min_clamped = bins.masked_fill(zero_mask, 1) + + new_means = new_means / bins_min_clamped[..., None] + + means = torch.where(zero_mask[..., None], means, new_means) + + return means, bins + + +def _rotate_half(x): + x1 = x[..., : x.shape[-1] // 2] + x2 = x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + + +def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + q_embed = (q * cos) + (_rotate_half(q) * sin) + k_embed = (k * cos) + (_rotate_half(k) * sin) + return q_embed, k_embed + + +def _compute_default_rope_parameters( + config=None, device=None, seq_len=None, **rope_kwargs +): + if config is not None and len(rope_kwargs) > 0: + raise ValueError( + "Unexpected arguments: `**rope_kwargs` and `config` are mutually exclusive" + ) + if len(rope_kwargs) > 0: + base = rope_kwargs["base"] + dim = rope_kwargs["dim"] + elif config is not None: + base = config.rope_theta + partial_rotary_factor = ( + config.partial_rotary_factor + if hasattr(config, "partial_rotary_factor") + else 1.0 + ) + head_dim = ( + getattr(config, "head_dim", None) + or config.hidden_size // config.num_attention_heads + ) + dim = int(head_dim * partial_rotary_factor) + attention_factor = 1.0 + inv_freq = 1.0 / ( + base + ** ( + torch.arange(0, dim, 2, dtype=torch.int64).to( + device=device, dtype=torch.float + ) + / dim + ) + ) + return inv_freq, attention_factor + + +_ROPE_INIT_FUNCTIONS = { + "default": _compute_default_rope_parameters, +} + + +def _dynamic_rope_update(rope_forward): + def dynamic_frequency_update(self, position_ids, device): + seq_len = torch.max(position_ids) + 1 + if seq_len > self.max_seq_len_cached: + inv_freq, self.attention_scaling = self.rope_init_fn( + self.config, device, seq_len=seq_len + ) + self.register_buffer("inv_freq", inv_freq, persistent=False) + self.max_seq_len_cached = seq_len + + if ( + seq_len < self.original_max_seq_len + and self.max_seq_len_cached > self.original_max_seq_len + ): + self.original_inv_freq = self.original_inv_freq.to(device) + self.register_buffer("inv_freq", self.original_inv_freq, persistent=False) + self.max_seq_len_cached = self.original_max_seq_len + + @wraps(rope_forward) + def wrapper(self, x, position_ids): + if "dynamic" in self.rope_type: + dynamic_frequency_update(self, position_ids, device=x.device) + return rope_forward(self, x, position_ids) + + return wrapper + + +class AudioRotaryEmbedding(nn.Module): + def __init__(self, base, dim, max_seq_len, rope_type="default", device=None): + super().__init__() + self.max_seq_len = max_seq_len + self.rope_type = rope_type + self.rope_init_fn = _ROPE_INIT_FUNCTIONS[self.rope_type] + inv_freq, self.attention_scaling = self.rope_init_fn( + device=device, base=base, dim=dim + ) + self.register_buffer("inv_freq", inv_freq, persistent=False) + self.original_inv_freq = self.inv_freq + + @torch.no_grad() + @_dynamic_rope_update + def forward(self, x, position_ids): + inv_freq_expanded = self.inv_freq[:, None].float().expand(-1, 1).to(x.device) + position_ids_expanded = position_ids[None, :].float() + device_type = ( + x.device.type + if isinstance(x.device.type, str) and x.device.type != "mps" + else "cpu" + ) + with torch.autocast(device_type=device_type, enabled=False): + freqs = ( + inv_freq_expanded.float() @ position_ids_expanded.float() + ).transpose(0, 1) + emb = torch.cat((freqs, freqs), dim=-1) + cos = emb.cos() * self.attention_scaling + sin = emb.sin() * self.attention_scaling + return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) + + +class EuclideanCodebook(nn.Module): + def __init__( + self, + dim: int, + codebook_size: int, + kmeans_init: int = False, + kmeans_iters: int = 10, + decay: float = 0.99, + epsilon: float = 1e-5, + threshold_ema_dead_code: int = 2, + ): + super().__init__() + self.decay = decay + init_fn: tp.Callable[..., torch.Tensor] | tp.Any = ( + _uniform_init if not kmeans_init else torch.zeros + ) + embed = init_fn(codebook_size, dim) + + self.codebook_size = codebook_size + self.kmeans_iters = kmeans_iters + self.epsilon = epsilon + self.threshold_ema_dead_code = threshold_ema_dead_code + + self.register_buffer("inited", torch.Tensor([not kmeans_init])) + self.register_buffer("cluster_size", torch.zeros(codebook_size)) + self.register_buffer("embed", embed) + self.register_buffer("embed_avg", embed.clone()) + + @torch.jit.ignore + def init_embed_(self, data): + if self.inited: + return + + embed, cluster_size = _kmeans(data, self.codebook_size, self.kmeans_iters) + self.embed.data.copy_(embed) + self.embed_avg.data.copy_(embed.clone()) + self.cluster_size.data.copy_(cluster_size) + self.inited.data.copy_(torch.Tensor([True])) + + def replace_(self, samples, mask): + replace_num = mask.sum() + modified_codebook = self.embed.clone() + modified_codebook[mask] = _sample_vectors(samples, replace_num) + self.embed.data.copy_(modified_codebook) + + def expire_codes_(self, batch_samples): + if self.threshold_ema_dead_code == 0: + return + + expired_codes = self.cluster_size < self.threshold_ema_dead_code + if not torch.any(expired_codes): + return + + batch_samples = rearrange(batch_samples, "... d -> (...) d") + self.replace_(batch_samples, mask=expired_codes) + + def preprocess(self, x): + x = rearrange(x, "... d -> (...) d") + return x + + def quantize(self, x): + embed = self.embed.t() + dist_val = -( + x.pow(2).sum(1, keepdim=True) + - 2 * x @ embed + + embed.pow(2).sum(0, keepdim=True) + ) + embed_ind = dist_val.max(dim=-1).indices + return embed_ind + + def postprocess_emb(self, embed_ind, shape): + return embed_ind.view(*shape[:-1]) + + def dequantize(self, embed_ind): + quantize = F.embedding(embed_ind, self.embed) + return quantize + + def encode(self, x): + shape = x.shape + x = self.preprocess(x) + embed_ind = self.quantize(x) + embed_ind = self.postprocess_emb(embed_ind, shape) + return embed_ind + + def decode(self, embed_ind): + quantize = self.dequantize(embed_ind) + return quantize + + def forward(self, x): + shape, dtype = x.shape, x.dtype + x = self.preprocess(x) + + self.init_embed_(x) + + embed_ind = self.quantize(x) + embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype) + embed_ind = self.postprocess_emb(embed_ind, shape) + quantize = self.dequantize(embed_ind) + + if self.training: + self.expire_codes_(x) + _ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay) + embed_sum = x.t() @ embed_onehot + _ema_inplace(self.embed_avg, embed_sum.t().contiguous(), self.decay) + cluster_size = ( + _laplace_smoothing(self.cluster_size, self.codebook_size, self.epsilon) + * self.cluster_size.sum() + ) + embed_normalized = self.embed_avg / cluster_size.unsqueeze(1) + self.embed.data.copy_(embed_normalized) + + return quantize, embed_ind + + +class VectorQuantization(nn.Module): + def __init__( + self, + dim: int, + codebook_size: int, + codebook_dim: int | None = None, + decay: float = 0.99, + epsilon: float = 1e-5, + kmeans_init: bool = True, + kmeans_iters: int = 50, + threshold_ema_dead_code: int = 2, + commitment_weight: float = 1.0, + ): + super().__init__() + _codebook_dim: int = _vq_default(codebook_dim, dim) + + requires_projection = _codebook_dim != dim + self.project_in = ( + nn.Linear(dim, _codebook_dim) if requires_projection else nn.Identity() + ) + self.project_out = ( + nn.Linear(_codebook_dim, dim) if requires_projection else nn.Identity() + ) + + self.epsilon = epsilon + self.commitment_weight = commitment_weight + + self._codebook = EuclideanCodebook( + dim=_codebook_dim, + codebook_size=codebook_size, + kmeans_init=kmeans_init, + kmeans_iters=kmeans_iters, + decay=decay, + epsilon=epsilon, + threshold_ema_dead_code=threshold_ema_dead_code, + ) + self.codebook_size = codebook_size + + @property + def codebook(self): + return self._codebook.embed + + def encode(self, x): + x = self.project_in(x) + embed_in = self._codebook.encode(x) + return embed_in + + def decode(self, embed_ind): + quantize = self._codebook.decode(embed_ind) + quantize = self.project_out(quantize) + return quantize + + def forward(self, x): + device = x.device + x = self.project_in(x) + + quantize, embed_ind = self._codebook(x) + + if self.training: + quantize = x + (quantize - x).detach() + + loss = torch.tensor([0.0], device=device, requires_grad=self.training) + + quantize = self.project_out(quantize) + return quantize, embed_ind, loss + + +class ResidualVectorQuantization(nn.Module): + def __init__(self, *, num_quantizers, codebook_size, **kwargs): + super().__init__() + if isinstance(codebook_size, int): + codebook_size = [codebook_size] * num_quantizers + elif len(codebook_size) < num_quantizers: + codebook_size += [codebook_size[-1]] * (num_quantizers - len(codebook_size)) + self.layers = nn.ModuleList( + [ + VectorQuantization(codebook_size=codebook_size[i], **kwargs) + for i in range(num_quantizers) + ] + ) + + def forward(self, x, n_q: int | None = None, layers: list | None = None): + quantized_out = 0.0 + residual = x + + all_losses = [] + all_indices = [] + out_quantized = [] + + n_q = n_q or len(self.layers) + + for i, layer in enumerate(self.layers[:n_q]): + quantized, indices, loss = layer(residual) + residual = residual - quantized + quantized_out = quantized_out + quantized + + all_indices.append(indices) + all_losses.append(loss) + if layers and i in layers: + out_quantized.append(quantized_out) + + out_losses, out_indices = map(torch.stack, (all_losses, all_indices)) + return quantized_out, out_indices, out_losses, out_quantized + + def encode( + self, x: torch.Tensor, n_q: int | None = None, st: int | None = None + ) -> torch.Tensor: + residual = x + all_indices = [] + n_q = len(self.layers) if n_q is None else n_q + st = 0 if st is None else st + for layer in self.layers[st:n_q]: + indices = layer.encode(residual) + quantized = layer.decode(indices) + residual = residual - quantized + all_indices.append(indices) + out_indices = torch.stack(all_indices) + return out_indices + + def decode(self, q_indices: torch.Tensor, st: int = 0) -> torch.Tensor: + quantized_out = self.layers[st].decode(q_indices[0]) + for i in range(1, len(q_indices)): + layer = self.layers[st + i] + quantized = layer.decode(q_indices[i]) + quantized_out = quantized_out + quantized + return quantized_out + + +class ResidualVectorQuantizer(nn.Module): + def __init__( + self, + dimension: int = 256, + n_q: int = 8, + bins: int | list = 1024, + decay: float = 0.99, + kmeans_init: bool = True, + kmeans_iters: int = 50, + threshold_ema_dead_code: int = 2, + ): + super().__init__() + self.n_q = n_q + self.dimension = dimension + self.bins = bins + self.decay = decay + self.kmeans_init = kmeans_init + self.kmeans_iters = kmeans_iters + self.threshold_ema_dead_code = threshold_ema_dead_code + self.vq = ResidualVectorQuantization( + dim=self.dimension, + codebook_size=self.bins, + num_quantizers=self.n_q, + decay=self.decay, + kmeans_init=self.kmeans_init, + kmeans_iters=self.kmeans_iters, + threshold_ema_dead_code=self.threshold_ema_dead_code, + ) + + def forward( + self, + x: torch.Tensor, + n_q: int | None = None, + layers: list | None = None, + ): + n_q = n_q if n_q else self.n_q + quantized, codes, commit_loss, quantized_list = self.vq( + x, n_q=n_q, layers=layers + ) + return quantized, codes, torch.mean(commit_loss), quantized_list + + def encode( + self, x: torch.Tensor, n_q: int | None = None, st: int | None = None + ) -> torch.Tensor: + n_q = n_q if n_q else self.n_q + st = st or 0 + codes = self.vq.encode(x, n_q=n_q, st=st) + return codes + + def decode(self, codes: torch.Tensor, st: int = 0) -> torch.Tensor: + quantized = self.vq.decode(codes, st=st) + return quantized + + +# --------------------------------------------------------------------------- +# Audio tokenizer +# --------------------------------------------------------------------------- + + +class MiMoAudioTokenizerConfig(PretrainedConfig): + model_type = "mimo_audio_tokenizer" + + def __init__( + self, + max_audio_seconds: int = 1800, + stride_size: int = 2, + avg_pooler: int = 1, + d_model: int = 768, + scale_embedding: bool = True, + kernel_size: int = 3, + activation_function: str = "gelu", + encoder_layers: int = 8, + encoder_skip_layer_id: int = None, + encoder_attention_heads: int = 12, + encoder_ffn_dim: int = 3072, + encoder_causal: bool = False, + encoder_attn_window_size: list = None, + decoder_layers: int = 8, + decoder_attention_heads: int = 12, + decoder_ffn_dim: int = 3072, + decoder_kernel_size: int = 3, + decoder_stride_size: int = 2, + decoder_causal: bool = True, + decoder_attn_window_size: list = None, + nfft: int = 1024, + vocoder_dim: int = 512, + vocoder_intermediate_dim: int = 4096, + vocoder_num_layers: int = 30, + n_mels: int = 80, + sampling_rate: int = 24000, + hop_length: int = 240, + window_size: int = 1024, + vocoder_padding: str = "same", + fmin: int = 0, + fmax: int = None, + num_quantizers: int = 12, + codebook_size: list = None, + threshold_ema_dead_code: int = 10, + position_embedding_type: str = "rope", + rope_theta: int = 10000, + rope_type: str = "default", + ln_type: str = "LayerNorm", + vocoder_attention_heads: int = 4, + vocoder_attn_window_size: list = None, + use_istft_only: bool = False, + hybrid_attention: bool = False, + hybrid_block_size: int = 8, + swa_per_block: int = 2, + **kwargs, + ): + super().__init__(**kwargs) + self.max_audio_seconds = max_audio_seconds + self.stride_size = stride_size + self.avg_pooler = avg_pooler + self.d_model = d_model + self.scale_embedding = scale_embedding + self.kernel_size = kernel_size + self.activation_function = activation_function + self.encoder_layers = encoder_layers + self.encoder_skip_layer_id = encoder_skip_layer_id + self.encoder_attention_heads = encoder_attention_heads + self.encoder_ffn_dim = encoder_ffn_dim + self.encoder_causal = encoder_causal + self.encoder_attn_window_size = ( + encoder_attn_window_size + if encoder_attn_window_size is not None + else [-1, -1] + ) + self.decoder_layers = decoder_layers + self.decoder_attention_heads = decoder_attention_heads + self.decoder_ffn_dim = decoder_ffn_dim + self.decoder_kernel_size = decoder_kernel_size + self.decoder_stride_size = decoder_stride_size + self.decoder_causal = decoder_causal + self.decoder_attn_window_size = ( + decoder_attn_window_size + if decoder_attn_window_size is not None + else [-1, -1] + ) + self.nfft = nfft + self.vocoder_dim = vocoder_dim + self.vocoder_intermediate_dim = vocoder_intermediate_dim + self.vocoder_num_layers = vocoder_num_layers + self.n_mels = n_mels + self.sampling_rate = sampling_rate + self.hop_length = hop_length + self.window_size = window_size + self.vocoder_padding = vocoder_padding + self.fmin = fmin + self.fmax = fmax + self.num_quantizers = num_quantizers + self.codebook_size = codebook_size if codebook_size is not None else [1024] + self.threshold_ema_dead_code = threshold_ema_dead_code + self.position_embedding_type = position_embedding_type + self.rope_theta = rope_theta + self.rope_type = rope_type + self.ln_type = ln_type + self.vocoder_attention_heads = vocoder_attention_heads + self.vocoder_attn_window_size = ( + vocoder_attn_window_size + if vocoder_attn_window_size is not None + else [40, 10] + ) + self.use_istft_only = use_istft_only + self.hybrid_attention = hybrid_attention + self.hybrid_block_size = hybrid_block_size + self.swa_per_block = swa_per_block + + +def get_sequence_mask(inputs, inputs_length): + if inputs.dim() == 3: + bsz, tgt_len, _ = inputs.size() + else: + bsz, tgt_len = inputs_length.shape[0], torch.max(inputs_length) + sequence_mask = torch.arange(0, tgt_len).to(inputs.device) + sequence_mask = torch.lt(sequence_mask, inputs_length.reshape(bsz, 1)).view( + bsz, tgt_len, 1 + ) + unpacking_index = torch.cumsum(sequence_mask.to(torch.int64).view(-1), dim=0) - 1 + return sequence_mask, unpacking_index + + +def unpack_hidden_states( + hidden_states, lengths, sequence_mask=None, unpacking_index=None +): + bsz = lengths.shape[0] + if sequence_mask is None or unpacking_index is None: + sequence_mask, unpacking_index = get_sequence_mask(hidden_states, lengths) + hidden_states = torch.index_select(hidden_states, 0, unpacking_index).view( + bsz, torch.max(lengths), hidden_states.shape[-1] + ) + return torch.where(sequence_mask, hidden_states, 0) + + +def get_position_ids(lengths): + total_len = lengths.sum() + offset = torch.cat([torch.zeros(1).to(lengths), lengths[:-1].cumsum(dim=0)]) + offset = torch.repeat_interleave(offset, lengths) + return torch.arange(0, total_len).to(offset) - offset + + +LAYER_NORM = {"LayerNorm": nn.LayerNorm} + + +class AudioEncoderAttention(nn.Module): + def __init__( + self, + embed_dim: int, + num_heads: int, + window_size: tuple[int, int] = (-1, -1), + causal: bool = False, + ): + super().__init__() + self.embed_dim = embed_dim + self.num_heads = num_heads + self.head_dim = embed_dim // num_heads + self.window_size = window_size + self.causal = causal + + self.k_proj = nn.Linear(embed_dim, embed_dim, bias=False) + self.v_proj = nn.Linear(embed_dim, embed_dim, bias=True) + self.q_proj = nn.Linear(embed_dim, embed_dim, bias=True) + self.out_proj = nn.Linear(embed_dim, embed_dim, bias=True) + + def forward( + self, + hidden_states: torch.Tensor, + cu_seqlens: torch.Tensor, + max_seqlen: int, + rope_position_embeddings=None, + ): + from vllm.vllm_flash_attn import flash_attn_varlen_func + + bsz, _ = hidden_states.size() + + query_states = self.q_proj(hidden_states).view( + bsz, self.num_heads, self.head_dim + ) + key_states = self.k_proj(hidden_states).view(bsz, self.num_heads, self.head_dim) + value_states = self.v_proj(hidden_states).view( + bsz, self.num_heads, self.head_dim + ) + + if rope_position_embeddings is not None: + cos, sin = rope_position_embeddings + query_states, key_states = apply_rotary_pos_emb( + query_states, key_states, cos, sin + ) + + attn_output = flash_attn_varlen_func( + query_states, + key_states, + value_states, + cu_seqlens_q=cu_seqlens, + cu_seqlens_k=cu_seqlens, + max_seqlen_q=max_seqlen, + max_seqlen_k=max_seqlen, + causal=self.causal, + window_size=list(self.window_size), + ) + + attn_output = attn_output.reshape(bsz, self.embed_dim) + attn_output = self.out_proj(attn_output) + return attn_output + + +class AudioEncoderTransformerLayer(nn.Module): + def __init__( + self, + config: MiMoAudioTokenizerConfig, + causal: bool, + attn_window_size: tuple[int, int] = (-1, -1), + ): + super().__init__() + self.embed_dim = config.d_model + + self.self_attn = AudioEncoderAttention( + embed_dim=self.embed_dim, + num_heads=config.encoder_attention_heads, + window_size=attn_window_size, + causal=causal, + ) + self.self_attn_layer_norm = LAYER_NORM[config.ln_type](self.embed_dim) + + self.activation_fn = ACT2FN[config.activation_function] + self.fc1 = nn.Linear(self.embed_dim, config.encoder_ffn_dim) + self.fc2 = nn.Linear(config.encoder_ffn_dim, self.embed_dim) + self.final_layer_norm = LAYER_NORM[config.ln_type](self.embed_dim) + + def forward( + self, + hidden_states: torch.Tensor, + cu_seqlens: torch.Tensor, + max_seqlen: int, + rope_position_embeddings: tuple[torch.Tensor, torch.Tensor], + ) -> torch.Tensor: + residual = hidden_states + hidden_states = self.self_attn_layer_norm(hidden_states) + hidden_states = self.self_attn( + hidden_states, + cu_seqlens, + max_seqlen, + rope_position_embeddings=rope_position_embeddings, + ) + hidden_states = residual + hidden_states + + residual = hidden_states + hidden_states = self.final_layer_norm(hidden_states) + hidden_states = self.activation_fn(self.fc1(hidden_states)) + hidden_states = self.fc2(hidden_states) + hidden_states = residual + hidden_states + + return hidden_states + + +class AudioEncoder(nn.Module): + def __init__( + self, + config: MiMoAudioTokenizerConfig, + ): + super().__init__() + self.config = config + self.max_source_positions = ( + config.max_audio_seconds * config.sampling_rate // config.hop_length + ) // config.stride_size + self.embed_scale = math.sqrt(config.d_model) if config.scale_embedding else 1.0 + self.skip_layer_idx = config.encoder_skip_layer_id + + self.conv1 = nn.Conv1d( + config.n_mels, + config.d_model, + kernel_size=config.kernel_size, + padding=1, + ) + self.conv2 = nn.Conv1d( + config.d_model, + config.d_model, + kernel_size=config.kernel_size, + stride=config.stride_size, + padding=1, + ) + + self.position_embedding = AudioRotaryEmbedding( + config.rope_theta, + config.d_model // config.encoder_attention_heads, + self.max_source_positions, + config.rope_type, + ) + + attn_window_sizes = [] + if config.hybrid_attention: + for i in range(config.encoder_layers): + if i % config.swa_per_block < config.swa_per_block - 1: + attn_window_sizes.append(tuple(config.encoder_attn_window_size)) + else: + attn_window_sizes.append((-1, -1)) + else: + attn_window_sizes = [ + tuple(config.encoder_attn_window_size) + ] * config.encoder_layers + + self.layers = nn.ModuleList( + [ + AudioEncoderTransformerLayer( + config=config, + causal=config.encoder_causal, + attn_window_size=attn_window_sizes[i], + ) + for i in range(config.encoder_layers) + ] + ) + + self.layer_norm = LAYER_NORM[config.ln_type](config.d_model) + + if config.avg_pooler != 1: + self.down_sample_layer = nn.Sequential( + nn.Conv1d( + config.d_model, + config.d_model, + config.avg_pooler, + config.avg_pooler, + bias=False, + ), + nn.GELU(), + ) + self.down_sample_norm = LAYER_NORM[config.ln_type](config.d_model) + else: + self.down_sample_layer = None + + if config.num_quantizers != 0: + self.quantizer = ResidualVectorQuantizer( + dimension=config.d_model, + n_q=config.num_quantizers, + bins=config.codebook_size, + threshold_ema_dead_code=config.threshold_ema_dead_code, + ) + else: + self.quantizer = None + + def get_features(self, input_features, output_length): + input_features = input_features.to(self.conv1.weight) + inputs_embeds = nn.functional.gelu(self.conv1(input_features)) + inputs_embeds = nn.functional.gelu(self.conv2(inputs_embeds)) + inputs_embeds = inputs_embeds.permute(0, 2, 1) + bsz, tgt_len, _ = inputs_embeds.size() + hidden_states = inputs_embeds + + position_ids = get_position_ids(output_length).long().to(input_features.device) + rope_position_embeddings = self.position_embedding(input_features, position_ids) + + attention_mask, unpacking_index = get_sequence_mask( + hidden_states, output_length + ) + hidden_states = torch.masked_select(hidden_states, attention_mask).view( + torch.sum(output_length), self.config.d_model + ) + + cu_seqlens = F.pad( + torch.cumsum(output_length, dim=0), (1, 0), "constant", 0 + ).to(device=hidden_states.device, dtype=torch.int32) + max_seqlen = torch.max(output_length).to(torch.int32).item() + + skip_connect_hidden_states = 0.0 + for idx, encoder_layer in enumerate(self.layers): + hidden_states = encoder_layer( + hidden_states, + cu_seqlens, + max_seqlen, + rope_position_embeddings=rope_position_embeddings, + ) + if (self.skip_layer_idx is not None) and idx == self.skip_layer_idx - 1: + skip_connect_hidden_states = hidden_states.clone() + + hidden_states += skip_connect_hidden_states + hidden_states = self.layer_norm(hidden_states) + + if self.down_sample_layer is not None: + hidden_states = torch.index_select(hidden_states, 0, unpacking_index).view( + bsz, tgt_len, self.config.d_model + ) + if hidden_states.size(1) % self.config.avg_pooler: + pad_len = ( + self.config.avg_pooler + - hidden_states.size(1) % self.config.avg_pooler + ) + hidden_states = torch.nn.functional.pad( + hidden_states, (0, 0, 0, pad_len), mode="constant", value=0.0 + ) + tgt_len += pad_len + tgt_len = tgt_len // self.config.avg_pooler + hidden_states = self.down_sample_layer(hidden_states.transpose(1, 2)) + output_length = ( + output_length // self.config.avg_pooler + + (output_length % self.config.avg_pooler != 0).int() + ) + hidden_states = hidden_states.transpose(1, 2) + attention_mask, unpacking_index = get_sequence_mask( + hidden_states, output_length + ) + hidden_states = torch.masked_select(hidden_states, attention_mask).view( + torch.sum(output_length), self.config.d_model + ) + hidden_states = self.down_sample_norm(hidden_states) + + return ( + hidden_states, + output_length, + attention_mask, + unpacking_index, + tgt_len, + bsz, + ) + + def get_output_length(self, mel_len): + tgt_len = mel_len + 3 - self.config.kernel_size + return (tgt_len + 2 - self.config.kernel_size) // self.config.stride_size + 1 + + @torch.no_grad() + def encode( + self, + input_features, + input_lens=None, + output_length=None, + return_codes_only=False, + n_q=None, + use_quantizer=True, + ): + if output_length is None: + output_length = self.get_output_length(input_lens) + input_features = unpack_hidden_states(input_features, input_lens) + hidden_states, output_length, attention_mask, unpacking_index, tgt_len, bsz = ( + self.get_features( + input_features=input_features.transpose(1, 2), + output_length=output_length, + ) + ) + + dtype = hidden_states.dtype + if use_quantizer and self.quantizer is not None: + self.quantizer.float() + codes = self.quantizer.encode(hidden_states.float(), n_q=n_q) + if return_codes_only: + return codes, output_length + hidden_states = self.quantizer.decode(codes) + hidden_states = hidden_states.to(dtype) + else: + codes = None + + hidden_states_packed = hidden_states.clone() + hidden_states = torch.index_select(hidden_states, 0, unpacking_index).view( + bsz, tgt_len, self.config.d_model + ) + hidden_states = torch.where(attention_mask, hidden_states, 0) + return hidden_states, hidden_states_packed, output_length, codes + + @torch.no_grad() + def decode_vq(self, codes): + self.quantizer.float() + return self.quantizer.decode(codes) + + +class MiMoAudioTokenizer(PreTrainedModel): + config_class = MiMoAudioTokenizerConfig + + def __init__(self, config: MiMoAudioTokenizerConfig): + super().__init__(config) + self.config = config + self.sampling_rate = config.sampling_rate + self.encoder = AudioEncoder(config=config) + self.downsample_rate = int(config.hop_length * 2 * config.avg_pooler) + + def get_output_length(self, mel_len): + tgt_len = mel_len + 3 - self.config.kernel_size + return (tgt_len + 2 - self.config.kernel_size) // self.config.stride_size + 1 + + @torch.no_grad() + def encode(self, mels, input_lens, use_quantizer=True): + input_features = mels + encoder_output_length = self.get_output_length(input_lens) + hidden_states, hidden_states_packed, encoder_output_length, codes = ( + self.encoder.encode( + input_features, input_lens=input_lens, use_quantizer=use_quantizer + ) + ) + return hidden_states, hidden_states_packed, encoder_output_length, codes + + +# --------------------------------------------------------------------------- +# Audio encoding utilities +# --------------------------------------------------------------------------- + + +def group_by_length(features: torch.Tensor, lengths: torch.Tensor, max_length: int): + if features.size(0) != lengths.sum().item(): + raise ValueError( + f"Feature size mismatch: {features.size(0)} vs {lengths.sum().item()}" + ) + + split_points = [] + current_sum = 0 + + for i, seq_len in enumerate(lengths): + if current_sum + seq_len > max_length and current_sum > 0: + split_points.append(i) + current_sum = seq_len.item() + else: + current_sum += seq_len.item() + + group_sizes = [] + prev = 0 + for point in split_points: + group_sizes.append(point - prev) + prev = point + if prev < len(lengths): + group_sizes.append(len(lengths) - prev) + + len_groups = torch.split(lengths, group_sizes) + feature_sizes = [group.sum().item() for group in len_groups] + feature_groups = torch.split(features, feature_sizes) + + return feature_groups, len_groups + + +@torch.no_grad() +def encode_batch( + audio_tokenizer_encoder, + input_features: torch.Tensor, + input_lens: torch.Tensor, + max_length: int = 256000, +): + feature_groups, len_groups = group_by_length(input_features, input_lens, max_length) + + encoded_parts = [] + for features, lengths in zip(feature_groups, len_groups): + codes, _ = audio_tokenizer_encoder.encode( + input_features=features, input_lens=lengths, return_codes_only=True + ) + encoded_parts.append(codes) + + return torch.cat(encoded_parts, dim=-1) + + +def _segment_lengths_for_mel(mel: torch.Tensor, segment_size: int): + """Split mel into segments of segment_size with a possible shorter remainder.""" + input_len = mel.size(0) + segs = [segment_size] * (input_len // segment_size) + if input_len % segment_size > 0: + segs.append(input_len % segment_size) + return segs + + +@torch.no_grad() +def tokenize_audio_batch(mels, audio_tokenizer_encoder, segment_size=6000, device=None): + """Tokenize multiple mels in one encode_batch call. + + Returns list of code tensors, each [T_i, C] for that mel. + """ + if not mels: + return [] + if device is None: + device = next(audio_tokenizer_encoder.parameters()).device + input_len_seg_per_mel = [_segment_lengths_for_mel(m, segment_size) for m in mels] + input_lens_flat = [s for segs in input_len_seg_per_mel for s in segs] + input_features = torch.cat([m.to(device) for m in mels], dim=0) + input_lens_t = torch.tensor(input_lens_flat, dtype=torch.long, device=device) + codes_packed = encode_batch( + audio_tokenizer_encoder, + input_features=input_features, + input_lens=input_lens_t, + ) + codes = codes_packed.transpose(0, 1).detach() # [total_code_T, C] + code_lengths = [] + for segs in input_len_seg_per_mel: + out_len = audio_tokenizer_encoder.get_output_length( + torch.tensor(segs, dtype=torch.long, device=device) + ) + if getattr(audio_tokenizer_encoder, "down_sample_layer", None) is not None: + avg = audio_tokenizer_encoder.config.avg_pooler + out_len = out_len // avg + (out_len % avg != 0).long() + code_lengths.append(out_len.sum().item()) + code_list = torch.split(codes, code_lengths) + return list(code_list) + + +# --------------------------------------------------------------------------- +# MimoAudioEncoderConfig +# --------------------------------------------------------------------------- + + +@dataclass +class MimoAudioEncoderConfig: + """Config for MimoAudioEncoder. + + Field names match the audio_config dict in the model checkpoint. + """ + + speech_vocab_size: str = "1025-1025-129-129-129-129-129-129" + speech_zeroemb_idx: str = "1024-1024-128-128-128-128-128-128" + group_size: int = 4 + audio_channels: int = 8 + input_local_layers: int = 6 + input_local_dim: int = 1024 + input_full_attention: bool = True + input_local_attn_heads: int = 64 + input_local_head_dim: int = 16 + input_local_intermediate_size: int = 4096 + input_local_hidden_dropout: float = 0.0 + out_hidden_size: int = 4096 + rope_theta: float = 640000.0 + partial_rotary_factor: float = 0.334 + projection_layers: int = 1 + add_post_norm: bool = False + audio_segment_size: int = 6000 + + @classmethod + def from_dict(cls, d: dict) -> "MimoAudioEncoderConfig": + known = {f.name for f in dataclasses.fields(cls)} + return cls(**{k: v for k, v in d.items() if k in known}) + + +# --------------------------------------------------------------------------- +# AudioProjection +# --------------------------------------------------------------------------- + + +class AudioProjection(nn.Module): + def __init__( + self, + input_size: int, + hidden_size: int, + output_size: int, + ) -> None: + super().__init__() + self.mlp = nn.Sequential( + nn.Linear(input_size, hidden_size, bias=False), + nn.GELU(), + nn.Linear(hidden_size, output_size, bias=False), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.mlp(x) + + +# --------------------------------------------------------------------------- +# MimoAudioEncoder +# --------------------------------------------------------------------------- + + +class MimoAudioEncoder(nn.Module): + """Audio encoder for MiMo-V2-Omni. + + Encodes mel spectrograms into LLM-compatible embeddings via: + 1. Audio tokenizer (VQ codes) + 2. Speech embeddings lookup + 3. Local Qwen2 transformer + 4. Linear projection + """ + + def __init__(self, config, model_path: str = "") -> None: + super().__init__() + if isinstance(config, dict): + config = MimoAudioEncoderConfig.from_dict(config) + self.config = config + self.audio_channels = config.audio_channels + self.audio_group_size = config.group_size + self.audio_segment_size = config.audio_segment_size + + speech_vocab_sizes = self._parse_maybe_list( + config.speech_vocab_size, config.audio_channels + ) + speech_empty_ids = self._parse_maybe_list( + config.speech_zeroemb_idx, config.audio_channels + ) + + input_local_config = Qwen2Config( + hidden_size=config.input_local_dim, + num_hidden_layers=config.input_local_layers, + num_attention_heads=config.input_local_attn_heads, + num_key_value_heads=config.input_local_attn_heads, + intermediate_size=config.input_local_intermediate_size, + attention_dropout=config.input_local_hidden_dropout, + rope_theta=config.rope_theta, + partial_rotary_factor=config.partial_rotary_factor, + ) + + self.input_local_transformer = Qwen2Model(input_local_config) + + if not config.add_post_norm: + self.input_local_transformer.norm = nn.Identity() + + self.speech_embeddings = nn.ModuleList( + [ + nn.Embedding( + speech_vocab_sizes[i], + config.input_local_dim, + padding_idx=speech_empty_ids[i], + ) + for i in range(config.audio_channels) + ] + ) + + if config.projection_layers == 1: + self.projection = nn.Linear( + config.input_local_dim * config.group_size, + config.out_hidden_size, + bias=False, + ) + elif config.projection_layers == 2: + self.projection = AudioProjection( + config.input_local_dim * config.group_size, + config.input_local_dim * config.group_size * 4, + config.out_hidden_size, + ) + else: + raise ValueError(f"Invalid projection_layers: {config.projection_layers}") + + self.audio_tokenizer: MiMoAudioTokenizer | None = None + if model_path: + audio_tokenizer_path = os.path.join(model_path, "audio_tokenizer") + if os.path.exists(audio_tokenizer_path): + dev = torch.get_default_device() + self.audio_tokenizer = self._load_audio_tokenizer( + audio_tokenizer_path, dev + ) + else: + logger.warning( + "Audio tokenizer not found at %s, audio encoding disabled", + audio_tokenizer_path, + ) + + @staticmethod + def _load_audio_tokenizer(path: str, device: torch.device) -> MiMoAudioTokenizer: + """Load MiMoAudioTokenizer from directory.""" + from safetensors.torch import load_file + + config_path = os.path.join(path, "config.json") + with open(config_path) as f: + config_dict = json.load(f) + config = MiMoAudioTokenizer.config_class(**config_dict) + model = MiMoAudioTokenizer(config) + safetensors_path = os.path.join(path, "model.safetensors") + bin_path = os.path.join(path, "pytorch_model.bin") + if os.path.exists(safetensors_path): + state_dict = load_file(safetensors_path, device="cpu") + elif os.path.exists(bin_path): + state_dict = torch.load(bin_path, map_location="cpu", weights_only=True) + else: + raise FileNotFoundError( + f"No model weights found in {path} " + "(expected model.safetensors or pytorch_model.bin)" + ) + model.load_state_dict(state_dict, strict=False) + model = model.to(device=device, dtype=torch.bfloat16) + model.eval() + model.requires_grad_(False) + return model + + def _parse_maybe_list(self, value, length: int) -> list[int]: + if isinstance(value, str) and "-" in value: + return [int(s) for s in value.split("-")] + return [int(value)] * length + + def apply_input_local_transformer(self, speech_embeddings: torch.Tensor): + output = self.input_local_transformer( + inputs_embeds=speech_embeddings, + return_dict=True, + is_causal=not self.config.input_full_attention, + ) + return output.last_hidden_state + + def apply_speech_embeddings(self, audio_codes: torch.Tensor) -> torch.Tensor: + num_segments = audio_codes.shape[0] + _audio_embeddings = torch.zeros( + (num_segments, self.config.group_size, self.config.input_local_dim), + dtype=next(self.speech_embeddings[0].parameters()).dtype, + device=audio_codes.device, + ) + for i in range(self.config.audio_channels): + _audio_embeddings.add_(self.speech_embeddings[i](audio_codes[:, :, i])) + return _audio_embeddings + + def process_audio(self, audio: torch.Tensor) -> torch.Tensor: + """Pad audio codes to group_size boundary. + + Args: + audio: [T, audio_channels] code tensor + + Returns: + [T//group_size, group_size, audio_channels] + """ + T = audio.shape[0] + audio = audio[:, : self.audio_channels] + padded_T = ( + (T + self.audio_group_size - 1) + // self.audio_group_size + * self.audio_group_size + ) + padded_audio = torch.cat( + [ + audio, + torch.zeros( + padded_T - T, + self.audio_channels, + dtype=torch.int32, + device=audio.device, + ) + + audio[-1, :], + ], + dim=0, + ) + padded_audio = padded_audio.reshape( + padded_T // self.audio_group_size, + self.audio_group_size, + self.audio_channels, + ) + return padded_audio + + def get_audio_feature( + self, mel_specs: list[torch.Tensor] + ) -> tuple[torch.Tensor, list[int]]: + """Encode mel spectrograms into LLM embedding space. + + Args: + mel_specs: list of mel spectrogram tensors, each [T, n_mels] + + Returns: + Tuple of: + - audio_embeds: [total_tokens, out_hidden_size] concatenated embeddings + - item_token_lens: list of int, number of tokens per input item + """ + if self.audio_tokenizer is None: + raise RuntimeError( + "audio_tokenizer is not loaded. " + "Ensure model_path points to a directory containing audio_tokenizer/." + ) + + if not mel_specs: + device = next(self.projection.parameters()).device + dtype = next(self.projection.parameters()).dtype + return ( + torch.empty(0, self.config.out_hidden_size, device=device, dtype=dtype), + [], + ) + + device = next(self.audio_tokenizer.encoder.parameters()).device + code_list = tokenize_audio_batch( + mel_specs, + self.audio_tokenizer.encoder, + segment_size=self.audio_segment_size, + device=device, + ) + + item_token_lens: list[int] = [] + codecs_to_concat = [] + for codecs in code_list: + padded_codes = self.process_audio(codecs) + codecs_to_concat.append(padded_codes) + item_token_lens.append(padded_codes.shape[0]) + + audio_codes = torch.cat( + codecs_to_concat, dim=0 + ) # [total_T//group_size, group_size, audio_channels] + + _audio_embeddings = self.apply_speech_embeddings(audio_codes) + audio_embeds = self.apply_input_local_transformer(_audio_embeddings) + B = audio_embeds.shape[0] + audio_embeds = self.projection(audio_embeds.reshape(B, -1)) + return audio_embeds, item_token_lens diff --git a/vllm/model_executor/models/mimo_v2_flash.py b/vllm/model_executor/models/mimo_v2.py similarity index 96% rename from vllm/model_executor/models/mimo_v2_flash.py rename to vllm/model_executor/models/mimo_v2.py index 380e3098204..c572df25ce0 100644 --- a/vllm/model_executor/models/mimo_v2_flash.py +++ b/vllm/model_executor/models/mimo_v2.py @@ -6,6 +6,7 @@ from itertools import islice import torch from torch import nn +from vllm.compilation.decorators import support_torch_compile from vllm.config import ( CacheConfig, VllmConfig, @@ -268,7 +269,7 @@ class MiMoV2Attention(nn.Module): self.total_num_heads * self.v_head_dim, hidden_size, bias=False, - quant_config=quant_config, + quant_config=quant_config if "mtp.layers" not in prefix else None, reduce_results=True, prefix=f"{prefix}.o_proj", ) @@ -440,6 +441,7 @@ class MiMoV2FlashDecoderLayer(nn.Module): return self.config.hybrid_layer_pattern[self.layer_id] == 1 +@support_torch_compile class MiMoV2Model(nn.Module): def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() @@ -603,7 +605,13 @@ class MiMoV2Model(nn.Module): if expert_matched: continue - + # Support fused qkv_proj checkpoint (Pro format) + if "qkv_proj" in name: + if name in params_dict: + param = params_dict[name] + loaded_weight = loaded_weight.chunk(tp_size, dim=0)[tp_rank] + default_weight_loader(param, loaded_weight) + continue stacked_matched = False for param_name, weight_name, shard_id in stacked_params_mapping: if weight_name not in name: @@ -662,6 +670,11 @@ class MiMoV2Model(nn.Module): class MiMoV2FlashForCausalLM(nn.Module, SupportsPP, MixtureOfExperts): + packed_modules_mapping = { + "qkv_proj": ["q_proj", "k_proj", "v_proj"], + "gate_up_proj": ["gate_proj", "up_proj"], + } + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() config = vllm_config.model_config.hf_config @@ -718,3 +731,10 @@ class MiMoV2FlashForCausalLM(nn.Module, SupportsPP, MixtureOfExperts): def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) return loader.load_weights(weights) + + +class MiMoV2ProForCausalLM(MiMoV2FlashForCausalLM): + packed_modules_mapping = { + "qkv_proj": ["qkv_proj"], + "gate_up_proj": ["gate_proj", "up_proj"], + } diff --git a/vllm/model_executor/models/mimo_v2_mtp.py b/vllm/model_executor/models/mimo_v2_mtp.py new file mode 100644 index 00000000000..442f4986b66 --- /dev/null +++ b/vllm/model_executor/models/mimo_v2_mtp.py @@ -0,0 +1,373 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +"""Inference-only MiMo-V2 MTP (Multi-Token Prediction) draft model. + +Supports both MiMo-V2-Pro and MiMo-V2-Flash checkpoints. + +Checkpoint weight layout (model.mtp.layers.{idx}.*): + enorm - RMSNorm for token embeddings + hnorm - RMSNorm for previous hidden states + eh_proj - ReplicatedLinear(hidden*2 -> hidden) + input_layernorm - pre-attention RMSNorm + self_attn.* - attention weights; format differs by variant: + Pro: fused qkv_proj [Q;K;V] concatenated + Flash: separate q_proj, k_proj, v_proj + pre_mlp_layernorm - post-attention / pre-MLP RMSNorm + mlp.* - dense MLP (gate_proj / up_proj / down_proj) + final_layernorm - norm applied before logit computation +""" + +from collections.abc import Iterable + +import torch +import torch.nn as nn +from transformers import PretrainedConfig + +from vllm.config import VllmConfig +from vllm.distributed import ( + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, +) +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ReplicatedLinear +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.model_executor.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) +from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm.sequence import IntermediateTensors + +from .interfaces import ( + MultiModalEmbeddings, + SupportsMultiModal, + _require_is_multimodal, +) +from .mimo_v2 import MiMoV2Attention, MiMoV2MLP +from .utils import _merge_multimodal_embeddings, maybe_prefix + +# MiMo-V2 checkpoints contain multiple MTP layers, but vLLM currently supports +# only the first layer and only one speculative token. +_MIMO_V2_PRO_NUM_MTP_LAYERS = 1 +_MIMO_V2_FLASH_NUM_MTP_LAYERS = 1 + + +class MiMoV2MTPLayer(nn.Module): + """Single MTP predictor layer for MiMo-V2 (Pro and Flash). + + Mirrors the single-layer MiMo-V2 nextn reference implementation. + """ + + def __init__( + self, + config: PretrainedConfig, + prefix: str, + quant_config: QuantizationConfig | None = None, + ) -> None: + super().__init__() + + # Predictor head components + self.enorm = RMSNorm(config.hidden_size, eps=config.layernorm_epsilon) + self.hnorm = RMSNorm(config.hidden_size, eps=config.layernorm_epsilon) + self.eh_proj = ReplicatedLinear( + config.hidden_size * 2, config.hidden_size, bias=False + ) + + # MTP uses the SWA attention configuration + # implementation. + swa_rope_theta = getattr( + config, + "swa_rope_theta", + getattr(config, "rope_theta", 1000000), + ) + sliding_window_size = getattr(config, "sliding_window_size", -1) + + self.input_layernorm = RMSNorm(config.hidden_size, eps=config.layernorm_epsilon) + self.self_attn = MiMoV2Attention( + hidden_size=config.hidden_size, + num_heads=config.swa_num_attention_heads, + num_kv_heads=config.swa_num_key_value_heads, + head_dim=config.swa_head_dim, + v_head_dim=getattr(config, "swa_v_head_dim", None), + v_scale=getattr(config, "attention_value_scale", None), + sliding_window_size=sliding_window_size, + attention_bias=config.attention_bias, + add_swa_attention_sink_bias=getattr( + config, "add_swa_attention_sink_bias", False + ), + layer_id=0, + rope_theta=swa_rope_theta, + max_position_embeddings=getattr(config, "max_position_embeddings", 32768), + quant_config=quant_config, + partial_rotary_factor=getattr(config, "partial_rotary_factor", 1.0), + prefix=f"{prefix}.self_attn", + ) + self.pre_mlp_layernorm = RMSNorm( + config.hidden_size, eps=config.layernorm_epsilon + ) + self.mlp = MiMoV2MLP( + hidden_size=config.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=f"{prefix}.mlp", + ) + self.final_layernorm = RMSNorm(config.hidden_size, eps=config.layernorm_epsilon) + + def forward( + self, + inputs_embeds: torch.Tensor, + positions: torch.Tensor, + previous_hidden_states: torch.Tensor, + ) -> torch.Tensor: + # Combine token embedding and previous hidden state + h, _ = self.eh_proj( + torch.cat( + [self.enorm(inputs_embeds), self.hnorm(previous_hidden_states)], dim=-1 + ) + ) + + # Transformer block with fused residual norms + residual = h + h = self.input_layernorm(h) + h = self.self_attn(positions=positions, hidden_states=h) + h, residual = self.pre_mlp_layernorm(h, residual) + h = self.mlp(h) + h = h + residual + + return self.final_layernorm(h) + + +class _MiMoV2MTPLayers(nn.Module): + """Thin wrapper so parameter paths match checkpoint: model.mtp.layers.*""" + + def __init__( + self, + config: PretrainedConfig, + num_mtp_layers: int, + quant_config: QuantizationConfig | None, + prefix: str, + ) -> None: + super().__init__() + self.layers = nn.ModuleDict( + { + str(i): MiMoV2MTPLayer( + config=config, + prefix=f"{prefix}.{i}", + quant_config=quant_config, + ) + for i in range(num_mtp_layers) + } + ) + + +class MiMoV2MultiTokenPredictor(nn.Module): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + super().__init__() + + config = vllm_config.model_config.hf_config + spec_cfg = vllm_config.speculative_config + assert spec_cfg is not None + if spec_cfg.num_speculative_tokens != 1: + raise ValueError( + "MiMo-V2 MTP in vLLM only supports num_speculative_tokens=1." + ) + num_mtp_layers = 1 + + self.num_mtp_layers = num_mtp_layers + + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + ) + + self.mtp = _MiMoV2MTPLayers( + config=config, + num_mtp_layers=num_mtp_layers, + quant_config=vllm_config.quant_config, + prefix=maybe_prefix(prefix, "mtp.layers"), + ) + + self.logits_processor = LogitsProcessor(config.vocab_size) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + previous_hidden_states: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + spec_step_idx: int = 0, + ) -> torch.Tensor: + assert spec_step_idx == 0, "MiMo-V2 MTP only supports one speculative token." + if inputs_embeds is None: + inputs_embeds = self.embed_input_ids(input_ids) + return self.mtp.layers[str(spec_step_idx)]( + inputs_embeds, positions, previous_hidden_states + ) + + def compute_logits( + self, + hidden_states: torch.Tensor, + lm_head: ParallelLMHead, + spec_step_idx: int = 0, + ) -> torch.Tensor: + assert spec_step_idx == 0, "MiMo-V2 MTP only supports one speculative token." + return self.logits_processor(lm_head, hidden_states) + + +class MiMoV2MTP(nn.Module): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + super().__init__() + self.config = vllm_config.model_config.hf_config + self.model = MiMoV2MultiTokenPredictor( + vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model") + ) + self.lm_head = ParallelLMHead( + self.config.vocab_size, + self.config.hidden_size, + prefix=maybe_prefix(prefix, "lm_head"), + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_input_ids(input_ids) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + hidden_states: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + spec_step_idx: int = 0, + ) -> torch.Tensor: + assert spec_step_idx == 0, "MiMo-V2 MTP only supports one speculative token." + return self.model( + input_ids, positions, hidden_states, inputs_embeds, spec_step_idx + ) + + def compute_logits( + self, + hidden_states: torch.Tensor, + spec_step_idx: int = 0, + ) -> torch.Tensor | None: + assert spec_step_idx == 0, "MiMo-V2 MTP only supports one speculative token." + return self.model.compute_logits(hidden_states, self.lm_head, spec_step_idx) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + tp_rank = get_tensor_model_parallel_rank() + tp_size = get_tensor_model_parallel_world_size() + + stacked_params_mapping = [ + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + # Flash format: separate projections → fused qkv_proj + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ] + + params_dict = dict(self.named_parameters()) + loaded_params: set[str] = set() + + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + + # Only load MTP-related weights, shared embeddings, and lm_head + if ( + "model.mtp" not in name + and "model.embed_tokens" not in name + and not name.startswith("lm_head") + ): + continue + + # Support fused qkv_proj checkpoint (Pro format). + # The checkpoint is stored pre-sharded for TP=8 as + # [Q_rank0, K_rank0, V_rank0, Q_rank1, ...], so splitting along + # dim 0 with chunk(tp_size) gives each rank its Q+K+V slice for + # both the FP8 weight and the block weight_scale_inv. This matches + # how the main model loads the same layout. + if "qkv_proj" in name: + if name in params_dict: + param = params_dict[name] + loaded_weight = loaded_weight.chunk(tp_size, dim=0)[tp_rank] + default_weight_loader(param, loaded_weight) + loaded_params.add(name) + continue + + # gate_proj/up_proj → gate_up_proj stacking (both formats); + # Flash: q_proj/k_proj/v_proj → qkv_proj merging. + stacked_matched = False + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + name_rewritten = name.replace(weight_name, param_name) + if ( + name_rewritten.endswith(".bias") + and name_rewritten not in params_dict + ): + continue + if name_rewritten not in params_dict: + continue + param = params_dict[name_rewritten] + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight, shard_id) + loaded_params.add(name_rewritten) + stacked_matched = True + break + + if stacked_matched: + continue + + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + continue + + param = params_dict[name] + # attention_sink_bias is head-parallel; slice by tp + if "attention_sink_bias" in name: + total_heads = loaded_weight.shape[0] + heads_per_rank = total_heads // tp_size + loaded_weight = loaded_weight.narrow( + 0, tp_rank * heads_per_rank, heads_per_rank + ) + + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight) + loaded_params.add(name) + + return loaded_params + + +class MiMoV2OmniMTP(MiMoV2MTP, SupportsMultiModal): + def embed_input_ids( + self, + input_ids: torch.Tensor, + multimodal_embeddings: MultiModalEmbeddings | None = None, + *, + is_multimodal: torch.Tensor | None = None, + ) -> torch.Tensor: + inputs_embeds = self._embed_text_input_ids( + input_ids, + self.model.embed_input_ids, + is_multimodal=is_multimodal, + ) + + if multimodal_embeddings is None or len(multimodal_embeddings) == 0: + return inputs_embeds + + is_multimodal = _require_is_multimodal(is_multimodal) + + inputs_embeds = _merge_multimodal_embeddings( + inputs_embeds=inputs_embeds, + multimodal_embeddings=multimodal_embeddings, + is_multimodal=is_multimodal, + ) + + return inputs_embeds diff --git a/vllm/model_executor/models/mimo_v2_omni.py b/vllm/model_executor/models/mimo_v2_omni.py new file mode 100644 index 00000000000..1cd2c6919a3 --- /dev/null +++ b/vllm/model_executor/models/mimo_v2_omni.py @@ -0,0 +1,1488 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import math +from collections.abc import Callable, Iterable, Mapping, Sequence +from functools import partial +from typing import Any + +import einops +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import BatchFeature, PretrainedConfig +from transformers.models.qwen2_vl.image_processing_qwen2_vl import smart_resize + +from vllm.config import VllmConfig +from vllm.config.multimodal import BaseDummyOptions +from vllm.distributed import parallel_state +from vllm.distributed import utils as dist_utils +from vllm.inputs import MultiModalDataDict +from vllm.model_executor.layers.activation import get_act_and_mul_fn +from vllm.model_executor.layers.attention import MMEncoderAttention +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ( + ColumnParallelLinear, + QKVParallelLinear, + RowParallelLinear, +) +from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.model_executor.layers.rotary_embedding import get_rope +from vllm.model_executor.layers.rotary_embedding.common import ApplyRotaryEmb +from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm.model_executor.models.vision import is_vit_use_data_parallel +from vllm.multimodal import MULTIMODAL_REGISTRY +from vllm.multimodal.inputs import MultiModalFieldConfig, MultiModalKwargsItems +from vllm.multimodal.parse import ImageSize, MultiModalDataItems +from vllm.multimodal.processing import ( + BaseDummyInputsBuilder, + BaseMultiModalProcessor, + BaseProcessingInfo, + PromptReplacement, + PromptUpdate, + PromptUpdateDetails, +) +from vllm.transformers_utils.configs.mimo_v2_omni import Mimo_VLVisionConfig +from vllm.transformers_utils.processors.mimo_v2_omni import ( + MiMoOmniProcessor, + VideoAudioInput, + _format_timestamp, +) + +from .interfaces import ( + MultiModalEmbeddings, + SupportsMultiModal, + SupportsPP, + SupportsQuant, +) +from .mimo_audio import MimoAudioEncoder +from .mimo_v2 import MiMoV2FlashForCausalLM +from .qwen2_5_vl import ( + Qwen2_5_VisionMLP, + Qwen2_5_VisionPatchEmbed, + Qwen2_5_VLImageEmbeddingInputs, + Qwen2_5_VLImageInputs, + Qwen2_5_VLImagePixelInputs, + Qwen2_5_VLVideoEmbeddingInputs, + Qwen2_5_VLVideoInputs, + Qwen2_5_VLVideoPixelInputs, +) +from .qwen2_vl import _create_qwen2vl_field_factory +from .utils import AutoWeightsLoader, IntermediateTensors, WeightsMapper, maybe_prefix + + +class MiMoVisionMLP(Qwen2_5_VisionMLP): + pass + + +class MiMoVisionPatchEmbed(Qwen2_5_VisionPatchEmbed): + pass + + +class MiMoVisionPatchMerger(nn.Module): + def __init__( + self, + d_model: int, + context_dim: int, + norm_layer: Callable[[int], nn.Module] | None = None, + spatial_merge_size: int = 2, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.hidden_size = context_dim * (spatial_merge_size**2) + if norm_layer is None: + norm_layer = partial(nn.LayerNorm, eps=1e-6) + self.ln_q = norm_layer(context_dim) + + self.mlp = nn.Sequential( + ColumnParallelLinear( + self.hidden_size, + self.hidden_size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.mlp.0", + return_bias=False, + disable_tp=use_data_parallel, + ), + nn.GELU(), + RowParallelLinear( + self.hidden_size, + d_model, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.mlp.2", + return_bias=False, + disable_tp=use_data_parallel, + ), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.ln_q(x) + x = x.view(-1, self.hidden_size) + out = self.mlp(x) + return out + + +class MiMoVisionAttention(nn.Module): + def __init__( + self, + embed_dim: int, + num_heads: int, + num_kv_heads: int, + qk_channels: int, + kv_channels: int, + use_sink: bool = False, + visual_token_window_size: int = 64, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.tp_size = ( + 1 + if use_data_parallel + else parallel_state.get_tensor_model_parallel_world_size() + ) + self.tp_rank = parallel_state.get_tensor_model_parallel_rank() + + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.qk_channels = qk_channels + self.kv_channels = kv_channels + self.embed_dim = embed_dim + + self.num_heads_per_partition = dist_utils.divide(num_heads, self.tp_size) + self.num_kv_heads_per_partition = dist_utils.divide(num_kv_heads, self.tp_size) + + # Attention scale uses the Q/K head dimension (qk_channels) + self.scale = qk_channels**-0.5 + + # QKV: Q is (num_heads * qk_channels), KV are (num_kv_heads * kv_channels) + self.qkv = QKVParallelLinear( + hidden_size=embed_dim, + head_size=qk_channels, + total_num_heads=num_heads, + total_num_kv_heads=num_kv_heads, + v_head_size=kv_channels, + bias=True, + quant_config=quant_config, + prefix=f"{prefix}.qkv", + disable_tp=use_data_parallel, + ) + + # Output projection: input is (num_heads * kv_channels) after attention + self.proj = RowParallelLinear( + input_size=num_heads * kv_channels, + output_size=embed_dim, + bias=True, + quant_config=quant_config, + prefix=f"{prefix}.proj", + disable_tp=use_data_parallel, + ) + + # For full attention (non-window blocks) + self.attn = MMEncoderAttention( + num_heads=self.num_heads_per_partition, + head_size=kv_channels, + scale=self.scale, + num_kv_heads=self.num_kv_heads_per_partition, + prefix=f"{prefix}.attn", + ) + + # Rotary embeddings applied separately to Q and K + self.apply_rotary_emb = ApplyRotaryEmb(enforce_enable=True) + + # Sink attention weights (loaded but not used in vLLM flash_attn) + # The checkpoint stores these only for non-full-attention blocks + self.use_sink = use_sink + if use_sink: + self.sinks = nn.Parameter( + torch.empty(num_heads), + requires_grad=False, + ) + else: + self.sinks = None + + self.visual_token_window_size = visual_token_window_size + + def _forward_window_attn( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + cu_seqlens: torch.Tensor, + max_seqlen: torch.Tensor, + ) -> torch.Tensor: + """Window attention via flash_attn_varlen_func with window_size.""" + from vllm.vllm_flash_attn import flash_attn_varlen_func + + w = self.visual_token_window_size + output = flash_attn_varlen_func( + q, + k, + v, + cu_seqlens_q=cu_seqlens, + cu_seqlens_k=cu_seqlens, + max_seqlen_q=max_seqlen, + max_seqlen_k=max_seqlen, + softmax_scale=self.scale, + causal=False, + window_size=[w, w], + ) + return output + + def forward( + self, + x: torch.Tensor, + cu_seqlens: torch.Tensor, + rotary_pos_emb_cos: torch.Tensor, + rotary_pos_emb_sin: torch.Tensor, + max_seqlen: torch.Tensor, + full_attn: bool = True, + ) -> torch.Tensor: + """ + Args: + x: [seq_len, batch=1, embed_dim] (seq-first convention) + cu_seqlens: cumulative sequence lengths [num_seqs+1], int32 + rotary_pos_emb_cos: [seq_len, qk_channels // 2] + rotary_pos_emb_sin: [seq_len, qk_channels // 2] + max_seqlen: maximum sequence length + full_attn: if True, full attention; if False, window attention + """ + # [seq_len, 1, embed_dim] -> QKV projection + qkv, _ = self.qkv(x) # [seq_len, 1, q_size + kv_size + kv_size] + seq_len, batch_size, _ = qkv.shape + + q_size = self.num_heads_per_partition * self.qk_channels + kv_size = self.num_kv_heads_per_partition * self.kv_channels + q, k, v = qkv.split([q_size, kv_size, kv_size], dim=-1) + + # Rearrange to [batch, seq, head, head_dim] for rotary application + q = einops.rearrange(q, "s b (h d) -> b s h d", h=self.num_heads_per_partition) + k = einops.rearrange( + k, "s b (h d) -> b s h d", h=self.num_kv_heads_per_partition + ) + v = einops.rearrange( + v, "s b (h d) -> b s h d", h=self.num_kv_heads_per_partition + ) + + # Apply rotary embeddings to Q and K independently (handles GQA) + if rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None: + q = self.apply_rotary_emb(q, rotary_pos_emb_cos, rotary_pos_emb_sin) + k = self.apply_rotary_emb(k, rotary_pos_emb_cos, rotary_pos_emb_sin) + + if full_attn: + # Full attention via MMEncoderAttention + # Flatten to [batch, seq, heads * head_dim] + q_flat = q.reshape(batch_size, seq_len, -1) + k_flat = k.reshape(batch_size, seq_len, -1) + v_flat = v.reshape(batch_size, seq_len, -1) + context_layer = self.attn( + query=q_flat, + key=k_flat, + value=v_flat, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + ) + # context_layer: [batch, seq, num_heads, head_dim] or [batch, seq, hidden] + # Ensure shape is [seq, batch, num_heads * kv_channels] + if context_layer.dim() == 4: + context_layer = einops.rearrange( + context_layer, "b s h d -> s b (h d)" + ).contiguous() + else: + context_layer = einops.rearrange( + context_layer, "b s d -> s b d" + ).contiguous() + else: + # Window attention via flash_attn_varlen_func with window_size + # Flatten batch dimension: [seq, head, head_dim] + q_varlen = einops.rearrange(q, "b s h d -> (b s) h d") + k_varlen = einops.rearrange(k, "b s h d -> (b s) h d") + v_varlen = einops.rearrange(v, "b s h d -> (b s) h d") + output = self._forward_window_attn( + q_varlen, k_varlen, v_varlen, cu_seqlens, max_seqlen + ) + # output: [total_tokens, num_heads, kv_channels] + context_layer = einops.rearrange( + output, "(b s) h d -> s b (h d)", b=batch_size + ).contiguous() + + output, _ = self.proj(context_layer) + return output + + +class MiMoVisionBlock(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + num_kv_heads: int, + qk_channels: int, + kv_channels: int, + mlp_hidden_dim: int, + act_fn: Callable[[torch.Tensor], torch.Tensor] = F.silu, + norm_eps: float = 1e-6, + use_sink: bool = False, + visual_token_window_size: int = 64, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + self.norm1 = RMSNorm(dim, eps=norm_eps) + self.norm2 = RMSNorm(dim, eps=norm_eps) + self.attn = MiMoVisionAttention( + embed_dim=dim, + num_heads=num_heads, + num_kv_heads=num_kv_heads, + qk_channels=qk_channels, + kv_channels=kv_channels, + use_sink=use_sink, + visual_token_window_size=visual_token_window_size, + quant_config=quant_config, + prefix=f"{prefix}.attn", + ) + self.mlp = MiMoVisionMLP( + in_features=dim, + hidden_features=mlp_hidden_dim, + act_fn=act_fn, + bias=True, + quant_config=quant_config, + prefix=f"{prefix}.mlp", + ) + + def forward( + self, + x: torch.Tensor, + cu_seqlens: torch.Tensor, + rotary_pos_emb_cos: torch.Tensor, + rotary_pos_emb_sin: torch.Tensor, + max_seqlen: torch.Tensor, + full_attn: bool = True, + ) -> torch.Tensor: + # x: [seq_len, batch=1, dim] + x_attn = self.attn( + self.norm1(x), + cu_seqlens=cu_seqlens, + rotary_pos_emb_cos=rotary_pos_emb_cos, + rotary_pos_emb_sin=rotary_pos_emb_sin, + max_seqlen=max_seqlen, + full_attn=full_attn, + ) + # Fused residual add + norm2 + x_norm, residual = self.norm2(x, residual=x_attn) + x = residual + self.mlp(x_norm) + return x + + +class MiMoVisionTransformer(nn.Module): + def __init__( + self, + vision_cfg: PretrainedConfig, + *, + norm_eps: float = 1e-6, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ): + super().__init__() + self.spatial_merge_size = vision_cfg.spatial_merge_size + self.spatial_merge_unit = self.spatial_merge_size**2 + self.fullatt_block_indexes = vision_cfg.fullatt_block_indexes + self.vit_window_attn_types = vision_cfg.vit_window_attn_types + self.visual_token_window_size = vision_cfg.visual_token_window_size + self.hidden_size = vision_cfg.hidden_size + self.num_heads = vision_cfg.num_heads + self.num_kv_heads = vision_cfg.num_key_value_heads + self.qk_channels = vision_cfg.qk_channels + self.kv_channels = vision_cfg.kv_channels + + self.patch_embed = MiMoVisionPatchEmbed( + patch_size=vision_cfg.patch_size, + temporal_patch_size=vision_cfg.temporal_patch_size, + in_channels=vision_cfg.in_channels, + hidden_size=vision_cfg.hidden_size, + ) + + norm_layer = partial(RMSNorm, eps=norm_eps) + + # Rotary embedding for 2D positions. + # With partial_rotary_factor=0.5 and head_size=qk_channels: + # rotary_dim = qk_channels // 2 + # get_cos_sin returns cos, sin each of shape [pos, rotary_dim // 2] + # After indexing with 2D pos_ids and flattening: + # result shape = [tokens, rotary_dim] = [tokens, qk_channels // 2] + # which is what ApplyRotaryEmb expects as cos/sin input. + self.rotary_pos_emb = get_rope( + head_size=vision_cfg.qk_channels, + max_position=8192, + is_neox_style=True, + rope_parameters={"partial_rotary_factor": 0.5}, + ) + + self.blocks = nn.ModuleList( + [ + MiMoVisionBlock( + dim=vision_cfg.hidden_size, + num_heads=vision_cfg.num_heads, + num_kv_heads=vision_cfg.num_key_value_heads, + qk_channels=vision_cfg.qk_channels, + kv_channels=vision_cfg.kv_channels, + mlp_hidden_dim=vision_cfg.intermediate_size, + act_fn=get_act_and_mul_fn(vision_cfg.hidden_act), + norm_eps=norm_eps, + use_sink=( + vision_cfg.use_sink + and i not in vision_cfg.fullatt_block_indexes + ), + visual_token_window_size=vision_cfg.visual_token_window_size, + quant_config=quant_config, + prefix=f"{prefix}.blocks.{i}", + ) + for i in range(vision_cfg.depth) + ] + ) + + self.merger = MiMoVisionPatchMerger( + d_model=vision_cfg.out_hidden_size, + context_dim=vision_cfg.hidden_size, + norm_layer=norm_layer, + spatial_merge_size=vision_cfg.spatial_merge_size, + quant_config=quant_config, + prefix=f"{prefix}.merger", + ) + + @property + def dtype(self) -> torch.dtype: + return self.patch_embed.proj.weight.dtype + + @property + def device(self) -> torch.device: + return self.patch_embed.proj.weight.device + + def apply_index(self, tensor: torch.Tensor, index: torch.Tensor) -> torch.Tensor: + """Reindex tensor at the spatial_merge_unit granularity.""" + tensor = tensor.unflatten(0, (-1, self.spatial_merge_unit)) + tensor = tensor[index] + tensor = tensor.flatten(0, 1) + return tensor + + def get_window_index_1d( + self, grid_thw: torch.Tensor, col: bool = True + ) -> torch.Tensor: + """Compute 1D window indices for col-based or row-based SWA reordering.""" + window_index: list[torch.Tensor] = [] + window_index_id = 0 + for grid_t, grid_h, grid_w in grid_thw: + llm_grid_h = grid_h // self.spatial_merge_size + llm_grid_w = grid_w // self.spatial_merge_size + index = torch.arange(grid_t * llm_grid_h * llm_grid_w).reshape( + grid_t, llm_grid_h, llm_grid_w + ) + index_new = index.transpose(1, 2).reshape(-1) if col else index.reshape(-1) + window_index.append(index_new + window_index_id) + window_index_id += int((grid_t * llm_grid_h * llm_grid_w).item()) + return torch.cat(window_index, dim=0) + + def rot_pos_emb(self, grid_thw: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Compute 2D rotary position embedding cos/sin for given grid sizes. + + Returns: + cos: [total_tokens, qk_channels // 2] + sin: [total_tokens, qk_channels // 2] + """ + cos_list, sin_list = [], [] + for i in range(grid_thw.size(0)): + t, h, w = int(grid_thw[i, 0]), int(grid_thw[i, 1]), int(grid_thw[i, 2]) + + # Build 2D position IDs with spatial_merge_size interleaving + hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w) + hpos_ids = ( + hpos_ids.reshape( + h // self.spatial_merge_size, + self.spatial_merge_size, + w // self.spatial_merge_size, + self.spatial_merge_size, + ) + .permute(0, 2, 1, 3) + .flatten() + ) + wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1) + wpos_ids = ( + wpos_ids.reshape( + h // self.spatial_merge_size, + self.spatial_merge_size, + w // self.spatial_merge_size, + self.spatial_merge_size, + ) + .permute(0, 2, 1, 3) + .flatten() + ) + pos_ids = torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1) + # pos_ids: [t*h*w, 2] + + max_grid_size = max(h, w) + # get_cos_sin returns cos, sin each of shape [max_grid_size, rotary_dim//2] + # where rotary_dim = qk_channels // 2 (from partial_rotary_factor=0.5) + cos, sin = self.rotary_pos_emb.get_cos_sin(max_grid_size) + + # [t*h*w, 2, rotary_dim//2] -> [t*h*w, rotary_dim] (= qk_channels // 2) + cos_img = cos[pos_ids].flatten(1) + sin_img = sin[pos_ids].flatten(1) + cos_list.append(cos_img) + sin_list.append(sin_img) + + return torch.cat(cos_list, dim=0), torch.cat(sin_list, dim=0) + + def forward(self, x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor: + """ + Args: + x: [total_tokens, C] pre-flattened patches + grid_thw: [num_images, 3] tensor of (t, h, w) for each image/video + Returns: + [merged_tokens, out_hidden_size] + """ + # Ensure grid_thw is a tensor + if not isinstance(grid_thw, torch.Tensor): + grid_thw = torch.tensor(grid_thw, dtype=torch.long) + + # Move to visual model device/dtype + x = x.to(device=self.device, dtype=self.dtype) + + # Patch embedding: [total_tokens, hidden_size] + x = self.patch_embed(x) + + # Compute 2D rotary positional embeddings + # cos, sin: [total_tokens, qk_channels // 2] + rotary_cos, rotary_sin = self.rot_pos_emb(grid_thw) + rotary_cos = rotary_cos.to(device=x.device) + rotary_sin = rotary_sin.to(device=x.device) + + # Compute cu_seqlens for flash_attn (per-image/video sequence lengths) + seqlens = torch.repeat_interleave( + grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0] + ) + cu_seqlens = torch.cat( + [ + torch.tensor([0], device=x.device, dtype=torch.int32), + seqlens.cumsum(dim=0).to(device=x.device, dtype=torch.int32), + ] + ) + max_seqlen = seqlens.max() + + # Precompute col-based window index for type=1 (col SWA) layers + window_index_1d_col = self.get_window_index_1d(grid_thw, col=True).to( + device=x.device + ) + reverse_window_index_1d_col = torch.argsort(window_index_1d_col) + + # Col-based rotary embeddings (reordered at spatial_merge_unit granularity). + # apply_index reorders groups of spatial_merge_unit tokens, just like x. + col_cos = self.apply_index(rotary_cos, window_index_1d_col) + col_sin = self.apply_index(rotary_sin, window_index_1d_col) + + # Add batch dimension: [total_tokens, 1, hidden_size] + x = x.unsqueeze(1) + + for i, blk in enumerate(self.blocks): + window_attn_type = self.vit_window_attn_types[i] + + # Reorder tokens to col-based layout when entering col-SWA region + if window_attn_type == 1 and ( + i == 0 or self.vit_window_attn_types[i - 1] != 1 + ): + x = self.apply_index(x, window_index_1d_col) + + # Restore row-based order when leaving col-SWA region + if ( + i > 0 + and window_attn_type != 1 + and self.vit_window_attn_types[i - 1] == 1 + ): + x = self.apply_index(x, reverse_window_index_1d_col) + + # Use col-based embeddings for col-SWA layers + cos_now = col_cos if window_attn_type == 1 else rotary_cos + sin_now = col_sin if window_attn_type == 1 else rotary_sin + + full_attn = i in self.fullatt_block_indexes + x = blk( + x, + cu_seqlens=cu_seqlens, + rotary_pos_emb_cos=cos_now, + rotary_pos_emb_sin=sin_now, + max_seqlen=max_seqlen, + full_attn=full_attn, + ) + + # Restore row-based order if last block was col-SWA + if self.vit_window_attn_types[-1] == 1: + x = self.apply_index(x, reverse_window_index_1d_col) + + # Remove batch dim and merge spatial tokens + # x: [total_tokens, 1, hidden_size] -> [total_tokens, hidden_size] + x = x.squeeze(1) + x = self.merger(x) + return x + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + stacked_params_mapping = [ + ("mlp.gate_up_proj", "mlp.gate_proj", 0), + ("mlp.gate_up_proj", "mlp.up_proj", 1), + ] + params_dict = dict(self.named_parameters(remove_duplicate=False)) + loaded_params: set[str] = set() + + for name, loaded_weight in weights: + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + name = name.replace(weight_name, param_name) + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + break + else: + param = params_dict[name] + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight) + loaded_params.add(name) + return loaded_params + + +class MiMoV2OmniProcessingInfo(BaseProcessingInfo): + def get_supported_mm_limits(self) -> Mapping[str, int | None]: + return {"audio": None, "image": None, "video": None} + + def get_hf_config(self): + config = self.ctx.get_hf_config() + if isinstance(config.vision_config, dict): + config.vision_config = Mimo_VLVisionConfig.from_dict(config.vision_config) + return config + + def get_hf_processor(self, **kwargs: object) -> MiMoOmniProcessor: + hf_config = self.get_hf_config() + tokenizer = self.get_tokenizer() + return MiMoOmniProcessor.from_hf_config(tokenizer, hf_config) + + def get_image_processor(self, **kwargs: object): + return self.get_hf_processor(**kwargs).image_processor + + def get_data_parser(self): + from vllm.multimodal.parse import MultiModalDataParser + + return MultiModalDataParser(target_sr=24000.0) + + def get_mm_max_tokens_per_item( + self, + seq_len: int, + mm_counts: Mapping[str, int], + ) -> Mapping[str, int]: + return { + "image": self.get_max_image_tokens(), + "video": self.get_max_video_tokens(seq_len, mm_counts), + } + + def _get_vision_info( + self, + *, + image_width: int, + image_height: int, + num_frames: int = 1, + do_resize: bool = True, + image_processor, + mm_kwargs: Mapping[str, object], + ) -> tuple[ImageSize, int]: + hf_config = self.get_hf_config() + vision_config = hf_config.vision_config + patch_size = vision_config.patch_size + merge_size = vision_config.spatial_merge_size + temporal_patch_size = vision_config.temporal_patch_size + tokens_per_second = vision_config.tokens_per_second + + mm_kwargs = self.ctx.get_merged_mm_kwargs(mm_kwargs) + size = image_processor.size + if override_size := mm_kwargs.get("size"): + size = size | override_size + if (override_min_pixels := mm_kwargs.get("min_pixels")) is not None: + size = size | {"shortest_edge": override_min_pixels} + if (override_max_pixels := mm_kwargs.get("max_pixels")) is not None: + size = size | {"longest_edge": override_max_pixels} + + if do_resize: + resized_height, resized_width = smart_resize( + height=image_height, + width=image_width, + factor=patch_size * merge_size, + min_pixels=size["shortest_edge"], + max_pixels=size["longest_edge"], + ) + preprocessed_size = ImageSize(width=resized_width, height=resized_height) + else: + preprocessed_size = ImageSize(width=image_width, height=image_height) + + # For video, MiMo resamples to tokens_per_second fps before temporal patching, + # effective tokens = num_frames * tokens_per_second / temporal_patch_size. + # For images (num_frames == 1) no resampling is applied. + if num_frames > 1: + effective_frames = num_frames * tokens_per_second + else: + effective_frames = num_frames + padded_num_frames = effective_frames + effective_frames % temporal_patch_size + grid_t = max(padded_num_frames // temporal_patch_size, 1) + grid_h = preprocessed_size.height // patch_size + grid_w = preprocessed_size.width // patch_size + num_patches = grid_t * grid_h * grid_w + num_vision_tokens = num_patches // (merge_size**2) + return preprocessed_size, num_vision_tokens + + def get_num_image_tokens( + self, + *, + image_width: int, + image_height: int, + image_processor, + mm_kwargs: Mapping[str, object], + ) -> int: + _, num_image_tokens = self._get_vision_info( + image_width=image_width, + image_height=image_height, + num_frames=1, + image_processor=image_processor, + mm_kwargs=mm_kwargs, + ) + return num_image_tokens + + def get_num_video_tokens( + self, + *, + image_width: int, + image_height: int, + num_frames: int, + image_processor, + mm_kwargs: Mapping[str, object], + ) -> int: + _, num_video_tokens = self._get_vision_info( + image_width=image_width, + image_height=image_height, + num_frames=num_frames, + image_processor=image_processor, + mm_kwargs=mm_kwargs, + ) + return num_video_tokens + + def get_image_size_with_most_features( + self, max_pixels: int | None = None + ) -> ImageSize: + hf_config = self.get_hf_config() + vision_config = hf_config.vision_config + patch_size = vision_config.patch_size + merge_size = vision_config.spatial_merge_size + + if max_pixels is None: + image_processor = self.get_image_processor() + mm_kwargs = self.ctx.get_merged_mm_kwargs({}) + size = image_processor.size + if override_size := mm_kwargs.get("size"): + size = size | override_size + if (override_min_pixels := mm_kwargs.get("min_pixels")) is not None: + size = size | {"shortest_edge": override_min_pixels} + if (override_max_pixels := mm_kwargs.get("max_pixels")) is not None: + size = size | {"longest_edge": override_max_pixels} + max_pixels = size["longest_edge"] + + unit = patch_size * merge_size + max_seq_len = max_pixels // (unit * unit) + + def closest_factor_pair(n: int) -> tuple[int, int]: + for d in range(math.isqrt(n), 0, -1): + if n % d == 0: + return d, n // d + return 1, n + + height_factor, width_factor = 1, max_seq_len + for seq_len in range(max_seq_len, 0, -1): + height_factor, width_factor = closest_factor_pair(seq_len) + if width_factor / height_factor <= 200: + break + + return ImageSize(width=unit * width_factor, height=unit * height_factor) + + def get_max_image_tokens(self) -> int: + image_processor = self.get_image_processor() + target_width, target_height = self.get_image_size_with_most_features() + return self.get_num_image_tokens( + image_width=target_width, + image_height=target_height, + image_processor=image_processor, + mm_kwargs={}, + ) + + def _get_max_video_frames(self, max_tokens: int, start_num_frames: int = 1) -> int: + image_processor = self.get_image_processor() + target_width, target_height = self.get_image_size_with_most_features() + num_frames = start_num_frames + while True: + next_num_frames = num_frames + 1 + next_max_tokens = self.get_num_video_tokens( + image_width=target_width, + image_height=target_height, + num_frames=next_num_frames, + image_processor=image_processor, + mm_kwargs={}, + ) + if next_max_tokens > max_tokens: + break + num_frames = next_num_frames + return num_frames + + def get_num_frames_with_most_features( + self, + seq_len: int, + mm_counts: Mapping[str, int], + max_frames_per_video: int = 14, + ) -> int: + max_videos = mm_counts.get("video", 0) + max_total_frames = self._get_max_video_frames(seq_len) + max_frames_per_video = min( + max_total_frames // max(max_videos, 1), max_frames_per_video + ) + return max(max_frames_per_video, 1) + + def get_max_video_tokens( + self, + seq_len: int, + mm_counts: Mapping[str, int], + ) -> int: + image_processor = self.get_image_processor() + target_width, target_height = self.get_image_size_with_most_features() + return self.get_num_video_tokens( + image_width=target_width, + image_height=target_height, + num_frames=self.get_num_frames_with_most_features(seq_len, mm_counts), + image_processor=image_processor, + mm_kwargs={}, + ) + + +class MiMoV2OmniMultiModalProcessor(BaseMultiModalProcessor[MiMoV2OmniProcessingInfo]): + """vLLM multimodal processor for MiMo-Omni (image + video). + + Key differences from Qwen2.5-VL: + - Videos use timestamp tokens between temporal grid positions. + - The HF processor expects ``(TCHW_tensor, timestamps_T_tensor)`` video + tuples rather than plain numpy arrays. + - ``video_start_times`` is tracked so prompt-update reconstruction can + regenerate the exact same timestamp token IDs. + """ + + # fps assumed for vllm-decoded video (numpy T,H,W,C arrays). + # The video loader samples ~32 frames; treat each frame as 1 s apart so + # MiMoVLProcessor sees 1 fps input and resamples internally. + _INPUT_FPS: float = 1.0 + + def _get_mm_fields_config( + self, + hf_inputs: BatchFeature, + hf_processor_mm_kwargs: Mapping[str, object], + ) -> Mapping[str, MultiModalFieldConfig]: + merge_size = self.info.get_hf_config().vision_config.spatial_merge_size + fields: dict[str, MultiModalFieldConfig] = dict( + **_create_qwen2vl_field_factory(merge_size)(hf_inputs), + second_per_grid_ts=MultiModalFieldConfig.batched("video"), + video_start_times=MultiModalFieldConfig.batched("video"), + audio_features=MultiModalFieldConfig.batched("audio"), + audio_token_lens=MultiModalFieldConfig.batched("audio"), + ) + # video_audio fields: only present when video_audio content was processed + if "video_audio_n_segs" in hf_inputs: + fields["video_audio_n_segs"] = MultiModalFieldConfig.batched("video") + # video_audio_seg_lens: list of per-video 1D tensors, batched("video") + if "video_audio_seg_lens" in hf_inputs: + fields["video_audio_seg_lens"] = MultiModalFieldConfig.batched("video") + if "va_audio_features" in hf_inputs: + fields["va_audio_features"] = MultiModalFieldConfig.batched("va_audio") + return fields + + def _call_hf_processor( + self, + prompt: str, + mm_data: Mapping[str, object], + mm_kwargs: Mapping[str, object], + tok_kwargs: Mapping[str, object], + ) -> BatchFeature: + """Convert numpy video arrays to (TCHW, timestamps) tuples for MiMo. + Also remap 'audios' → 'audio' since MiMoOmniProcessor.__call__ uses + the singular form. + """ + # Remap audios → audio (MiMoOmniProcessor uses singular param name) + if "audios" in mm_data: + mm_data = {**mm_data, "audio": mm_data["audios"]} + mm_data = {k: v for k, v in mm_data.items() if k != "audios"} + + # Handle video_audio items: convert video part to (TCHW, timestamps) tuple + if "video_audio" in mm_data: + va_converted: list[VideoAudioInput] = [] + for va_item in mm_data["video_audio"]: + if isinstance(va_item, VideoAudioInput): + vid = va_item.video + else: + # Expect (video_frames, audio_source) tuple + vid, audio_src = va_item + va_item = VideoAudioInput(video=vid, audio=audio_src) + vid = vid + # Convert video frames to (TCHW, timestamps) if needed + if ( + isinstance(vid, tuple) + and len(vid) == 2 + and isinstance(vid[0], torch.Tensor) + and isinstance(vid[1], torch.Tensor) + ): + va_converted.append(va_item) + else: + if isinstance(vid, np.ndarray): + frames = torch.from_numpy(vid) + elif isinstance(vid, torch.Tensor): + frames = vid + else: + frames = torch.tensor(np.array(vid)) + if frames.ndim == 4 and frames.shape[-1] in (1, 3, 4): + frames = frames.permute(0, 3, 1, 2).float() + else: + frames = frames.float() + T = frames.shape[0] + timestamps = torch.arange(T, dtype=torch.float32) / self._INPUT_FPS + va_converted.append( + VideoAudioInput( + video=(frames, timestamps), + audio=va_item.audio, + ) + ) + mm_data = {**mm_data, "video_audio": va_converted} + + if "videos" in mm_data: + converted: list[tuple[torch.Tensor, torch.Tensor]] = [] + for video in mm_data["videos"]: + if ( + isinstance(video, tuple) + and len(video) == 2 + and isinstance(video[0], torch.Tensor) + and isinstance(video[1], torch.Tensor) + ): + # already in MiMo format + converted.append(video) + else: + # numpy (T, H, W, C) or torch (T, H, W, C) / (T, C, H, W) + if isinstance(video, np.ndarray): + frames = torch.from_numpy(video) + elif isinstance(video, torch.Tensor): + frames = video + else: + frames = torch.tensor(np.array(video)) + + if frames.ndim == 4 and frames.shape[-1] in (1, 3, 4): + # THWC → TCHW + frames = frames.permute(0, 3, 1, 2).float() + else: + frames = frames.float() + + T = frames.shape[0] + timestamps = torch.arange(T, dtype=torch.float32) / self._INPUT_FPS + converted.append((frames, timestamps)) + + mm_data = {**mm_data, "videos": converted} + + return super()._call_hf_processor(prompt, mm_data, mm_kwargs, tok_kwargs) + + def _get_prompt_updates( + self, + mm_items: MultiModalDataItems, + hf_processor_mm_kwargs: Mapping[str, Any], + out_mm_kwargs: MultiModalKwargsItems, + ) -> Sequence[PromptUpdate]: + hf_processor = self.info.get_hf_processor(**hf_processor_mm_kwargs) + hf_config = self.info.get_hf_config() + tokenizer = self.info.get_tokenizer() + vocab = tokenizer.get_vocab() + + merge_size = hf_config.vision_config.spatial_merge_size + p = hf_processor.mimo_processor + + image_pad_id = vocab[hf_processor.image_token] + video_pad_id = vocab[hf_processor.video_token] + audio_pad_id = vocab.get("<|audio_pad|>") + vision_start_id = p.vision_start_token_id + vision_end_id = p.vision_end_token_id + video_start_id = p.video_start_token_id + video_end_id = p.video_end_token_id + audio_start_id = p.audio_start_token_id + audio_end_id = p.audio_end_token_id + + def get_image_replacement(item_idx: int) -> PromptUpdateDetails: + out_item = out_mm_kwargs["image"][item_idx] + grid_thw = out_item["image_grid_thw"].data + n_tokens = int(grid_thw.prod()) // merge_size**2 + return [image_pad_id] * n_tokens + + def get_video_replacement(item_idx: int) -> PromptUpdateDetails: + out_item = out_mm_kwargs["video"][item_idx] + grid_thw = out_item["video_grid_thw"].data + spt = float(out_item["second_per_grid_ts"].data) + start = float(out_item["video_start_times"].data) + + T, H, W = map(int, grid_thw) + n_per_grid = H * W // (merge_size * merge_size) + + # Check if this is a video_audio item + n_segs_field = out_item.get("video_audio_n_segs") + n_segs_val = int(n_segs_field.data) if n_segs_field is not None else 0 + va_seg_lens: list[int] | None = None + if n_segs_val > 0: + seg_lens_field = out_item.get("video_audio_seg_lens") + if seg_lens_field is not None: + va_seg_lens = seg_lens_field.data[:n_segs_val].tolist() + + full: list[int] = [video_start_id] + is_embed_mask: list[bool] = [False] + + if va_seg_lens is None: + # Regular video: timestamp + vision tokens per grid + for j in range(T): + ts_text = _format_timestamp(start + j * spt) + ts_ids = tokenizer.encode(ts_text, add_special_tokens=False) + full.extend(ts_ids) + is_embed_mask.extend([False] * len(ts_ids)) + full.append(vision_start_id) + is_embed_mask.append(False) + full.extend([video_pad_id] * n_per_grid) + is_embed_mask.extend([True] * n_per_grid) + full.append(vision_end_id) + is_embed_mask.append(False) + else: + # video_audio: interleaved vision+audio per group + n_groups = len(va_seg_lens) + frames_per_group = T // n_groups # 1 for il=0, T for il=-1 + for g in range(n_groups): + # Timestamp for first frame of this group + frame0 = g * frames_per_group + ts_text = _format_timestamp(start + frame0 * spt) + ts_ids = tokenizer.encode(ts_text, add_special_tokens=False) + full.extend(ts_ids) + is_embed_mask.extend([False] * len(ts_ids)) + # Vision tokens for all frames in this group + for f in range(frames_per_group): + full.append(vision_start_id) + is_embed_mask.append(False) + full.extend([video_pad_id] * n_per_grid) + is_embed_mask.extend([True] * n_per_grid) + full.append(vision_end_id) + is_embed_mask.append(False) + # Audio tokens for this group + seg_len = va_seg_lens[g] + full.append(audio_start_id) + is_embed_mask.append(False) + full.extend([audio_pad_id] * seg_len) + is_embed_mask.extend([True] * seg_len) + full.append(audio_end_id) + is_embed_mask.append(False) + + full.append(video_end_id) + is_embed_mask.append(False) + + embed_t = torch.tensor(is_embed_mask) + return PromptUpdateDetails( + full=full, + is_embed=lambda _tok, _seq: embed_t, + ) + + def get_audio_replacement(item_idx: int) -> PromptUpdateDetails: + out_item = out_mm_kwargs["audio"][item_idx] + tok_len = int(out_item["audio_token_lens"].data) + return [audio_pad_id] * tok_len + + updates: list[PromptUpdate] = [ + PromptReplacement( + modality="image", + target=[image_pad_id], + replacement=get_image_replacement, + ), + PromptReplacement( + modality="video", + target=[video_pad_id], + replacement=get_video_replacement, + ), + ] + if audio_pad_id is not None and audio_start_id is not None: + updates.append( + PromptReplacement( + modality="audio", + target=[audio_pad_id], + replacement=get_audio_replacement, + ) + ) + return updates + + +class MiMoV2OmniDummyInputsBuilder(BaseDummyInputsBuilder[MiMoV2OmniProcessingInfo]): + def get_dummy_text(self, mm_counts: Mapping[str, int]) -> str: + num_images = mm_counts.get("image", 0) + num_videos = mm_counts.get("video", 0) + num_audios = mm_counts.get("audio", 0) + image_ph = "<|vision_start|><|image_pad|><|vision_end|>" + video_ph = "<|vision_start|><|video_pad|><|vision_end|>" + audio_ph = "<|mimo_audio_start|><|audio_pad|><|mimo_audio_end|>" + return image_ph * num_images + video_ph * num_videos + audio_ph * num_audios + + def get_dummy_mm_data( + self, + seq_len: int, + mm_counts: Mapping[str, int], + mm_options: Mapping[str, BaseDummyOptions], + ) -> MultiModalDataDict: + num_images = mm_counts.get("image", 0) + num_videos = mm_counts.get("video", 0) + + target_width, target_height = self.info.get_image_size_with_most_features() + target_num_frames = self.info.get_num_frames_with_most_features( + seq_len, mm_counts + ) + + return { + "image": self._get_dummy_images( + width=target_width, + height=target_height, + num_images=num_images, + overrides=mm_options.get("image"), + ), + "video": self._get_dummy_videos( + width=target_width, + height=target_height, + num_frames=target_num_frames, + num_videos=num_videos, + overrides=mm_options.get("video"), + ), + } + + +@MULTIMODAL_REGISTRY.register_processor( + MiMoV2OmniMultiModalProcessor, + info=MiMoV2OmniProcessingInfo, + dummy_inputs=MiMoV2OmniDummyInputsBuilder, +) +class MiMoV2OmniForCausalLM(nn.Module, SupportsMultiModal, SupportsPP, SupportsQuant): + # To ensure correct weight loading and mapping. + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_prefix={ + # audio encoder + "speech_embeddings.": "audio_encoder.speech_embeddings.", + # mapping for new names in checkpoint saved after transformers v4.52 + "model.language_model.": "language_model.model.", + "model.visual.": "visual.", + # mapping for original checkpoint + "lm_head.": "language_model.lm_head.", + "model.": "language_model.model.", + } + ) + + @classmethod + def get_placeholder_str(cls, modality: str, i: int) -> str | None: + if modality.startswith("image"): + return "<|vision_start|><|image_pad|><|vision_end|>" + if modality.startswith("video"): + return "<|vision_start|><|video_pad|><|vision_end|>" + if modality.startswith("audio"): + return "<|mimo_audio_start|><|audio_pad|><|mimo_audio_end|>" + + raise ValueError(f"Unsupported modality: {modality}") + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + config = vllm_config.model_config.hf_config + self.config = config + # Omni ViT/Audio Encoder BF16 + vision_config = ( + Mimo_VLVisionConfig.from_dict(config.vision_config) + if isinstance(config.vision_config, dict) + else config.vision_config + ) + with self._mark_tower_model(vllm_config, {"image", "video"}): + self.visual = MiMoVisionTransformer( + vision_config, + norm_eps=getattr(vllm_config, "rms_norm_eps", 1e-6), + quant_config=None, + prefix=maybe_prefix(prefix, "visual"), + ) + audio_config = getattr(config, "audio_config", None) + model_path = vllm_config.model_config.model + if audio_config is not None: + with self._mark_tower_model(vllm_config, "audio"): + self.audio_encoder = MimoAudioEncoder( + audio_config, model_path=model_path + ) + else: + self.audio_encoder = None + with self._mark_language_model(vllm_config): + self.language_model = MiMoV2FlashForCausalLM( + vllm_config=vllm_config, + prefix=maybe_prefix(prefix, "language_model"), + ) + + self.make_empty_intermediate_tensors = ( + self.language_model.make_empty_intermediate_tensors + ) + + def _parse_and_validate_image_input( + self, **kwargs: object + ) -> Qwen2_5_VLImageInputs | None: + pixel_values = kwargs.pop("pixel_values", None) + image_embeds = kwargs.pop("image_embeds", None) + image_grid_thw = kwargs.pop("image_grid_thw", None) + + if pixel_values is None and image_embeds is None: + return None + + if pixel_values is not None: + return Qwen2_5_VLImagePixelInputs( + type="pixel_values", + pixel_values=pixel_values, + image_grid_thw=image_grid_thw, + ) + + if image_embeds is not None: + return Qwen2_5_VLImageEmbeddingInputs( + type="image_embeds", + image_embeds=image_embeds, + image_grid_thw=image_grid_thw, + ) + + def _parse_and_validate_video_input( + self, **kwargs: object + ) -> Qwen2_5_VLVideoInputs | None: + pixel_values_videos = kwargs.pop("pixel_values_videos", None) + video_embeds = kwargs.pop("video_embeds", None) + video_grid_thw = kwargs.pop("video_grid_thw", None) + second_per_grid_ts = kwargs.pop("second_per_grid_ts", None) + + if pixel_values_videos is None and video_embeds is None: + return None + + if pixel_values_videos is not None: + return Qwen2_5_VLVideoPixelInputs( + type="pixel_values_videos", + pixel_values_videos=pixel_values_videos, + video_grid_thw=video_grid_thw, + second_per_grid_ts=second_per_grid_ts, + ) + + if video_embeds is not None: + return Qwen2_5_VLVideoEmbeddingInputs( + type="video_embeds", + video_embeds=video_embeds, + video_grid_thw=video_grid_thw, + second_per_grid_ts=second_per_grid_ts, + ) + + def _process_image_input( + self, image_input: Qwen2_5_VLImageInputs + ) -> tuple[torch.Tensor, ...]: + grid_thw = image_input["image_grid_thw"] + assert grid_thw.ndim == 2 + grid_thw_list = grid_thw.tolist() + + if image_input["type"] == "image_embeds": + image_embeds = image_input["image_embeds"].type(self.visual.dtype) + else: + pixel_values = image_input["pixel_values"] + image_embeds = self.visual(pixel_values, grid_thw=grid_thw_list) + + # Split concatenated embeddings for each image item. + merge_size = self.visual.spatial_merge_size + sizes = (grid_thw.prod(-1) // merge_size // merge_size).tolist() + return image_embeds.split(sizes) + + def _process_video_input( + self, video_input: Qwen2_5_VLVideoInputs + ) -> tuple[torch.Tensor, ...]: + grid_thw = video_input["video_grid_thw"] + assert grid_thw.ndim == 2 + grid_thw_list = grid_thw.tolist() + + if video_input["type"] == "video_embeds": + video_embeds = video_input["video_embeds"].type(self.visual.dtype) + else: + pixel_values_videos = video_input["pixel_values_videos"] + video_embeds = self.visual(pixel_values_videos, grid_thw=grid_thw_list) + + # Split concatenated embeddings for each video item. + merge_size = self.visual.spatial_merge_size + sizes = (grid_thw.prod(-1) // merge_size // merge_size).tolist() + return video_embeds.split(sizes) + + def _parse_and_validate_audio_input(self, **kwargs: object) -> dict | None: + audio_features = kwargs.pop("audio_features", None) + audio_token_lens = kwargs.pop("audio_token_lens", None) + if audio_features is None: + return None + return { + "type": "audio", + "audio_features": audio_features, + "audio_token_lens": audio_token_lens, + } + + def _parse_and_validate_multimodal_inputs(self, **kwargs: object) -> dict: + mm_input_by_modality = {} + + # Preserve the order of modalities if there are multiple of them + # from the order of kwargs. + for input_key in kwargs: + if ( + input_key in ("pixel_values", "image_embeds") + and "image" not in mm_input_by_modality + ): + mm_input_by_modality["image"] = self._parse_and_validate_image_input( + **kwargs + ) + if ( + input_key in ("pixel_values_videos", "video_embeds") + and "video" not in mm_input_by_modality + ): + mm_input_by_modality["video"] = self._parse_and_validate_video_input( + **kwargs + ) + if input_key == "audio_features" and "audio" not in mm_input_by_modality: + mm_input_by_modality["audio"] = self._parse_and_validate_audio_input( + **kwargs + ) + return mm_input_by_modality + + def _process_audio_input(self, audio_input: dict) -> tuple[torch.Tensor, ...]: + mel_specs = audio_input["audio_features"] + if self.audio_encoder is None: + return () + # Normalize to List[2D-Tensor]. + # MultiModalBatchedField._reduce_data either wraps a single [T, 128] + # into [1, T, 128] via unsqueeze(0) or stacks N same-T items into + # [N, T, 128]. Indexing along dim-0 extracts the per-item [T, 128]. + if isinstance(mel_specs, torch.Tensor): + mel_specs = list(mel_specs) # [1,T,128] or [N,T,128] → [[T,128],...] + if not mel_specs: + return () + audio_embeds, item_token_lens = self.audio_encoder.get_audio_feature(mel_specs) + return tuple(audio_embeds.split(item_token_lens)) + + def embed_multimodal(self, **kwargs: object) -> MultiModalEmbeddings: + # Pop video_audio-specific fields before main mm parsing + video_audio_n_segs = kwargs.pop("video_audio_n_segs", None) + video_audio_seg_lens = kwargs.pop("video_audio_seg_lens", None) + va_audio_features = kwargs.pop("va_audio_features", None) + + mm_input_by_modality = self._parse_and_validate_multimodal_inputs(**kwargs) + if not mm_input_by_modality and va_audio_features is None: + return [] + + # The result multimodal_embeddings is tuple of tensors, with each + # tensor corresponding to a multimodal data item (image, video, or audio). + multimodal_embeddings: list[torch.Tensor] = [] + + # Pre-process va audio: one mel spec per va video → per-video audio embeddings + # keyed by va video index (0-based among va videos only) + va_audio_embs_list: list[tuple[torch.Tensor, ...]] = [] + if va_audio_features is not None and self.audio_encoder is not None: + mel_list = ( + list(va_audio_features) + if isinstance(va_audio_features, torch.Tensor) + else list(va_audio_features) + ) + for mel_spec in mel_list: + embs, tok_lens = self.audio_encoder.get_audio_feature([mel_spec]) + # tok_lens is a list/tensor with one entry (total tokens for this mel) + va_audio_embs_list.append(embs) # shape (total_tok, hidden) + + va_cursor = 0 # index into va_audio_embs_list + + # NOTE: Iterate in dict insertion order to preserve token sequence order. + for modality in mm_input_by_modality: + multimodal_input = mm_input_by_modality[modality] + if modality == "image": + multimodal_embeddings.extend( + self._process_image_input(multimodal_input) + ) + elif modality == "video": + video_embs_tuple = self._process_video_input(multimodal_input) + if video_audio_n_segs is None: + multimodal_embeddings.extend(video_embs_tuple) + else: + grid_thw = multimodal_input["video_grid_thw"] + for i, vid_embs in enumerate(video_embs_tuple): + n_segs = int(video_audio_n_segs[i]) + if n_segs == 0 or not va_audio_embs_list: + multimodal_embeddings.append(vid_embs) + else: + T = int(grid_thw[i][0]) + n_per_grid = vid_embs.shape[0] // T + frames = list(vid_embs.split(n_per_grid, dim=0)) + frames_per_group = T // n_segs + # Per-group audio token lengths for this va video + # video_audio_seg_lens is (num_videos, max_T); row i + # has valid values in [:n_segs], rest are zeros. + seg_lens = video_audio_seg_lens[i][:n_segs].tolist() + # Split full audio embs for this va video by group lengths + full_va_embs = va_audio_embs_list[va_cursor] + va_cursor += 1 + group_audio_embs = full_va_embs.split(seg_lens) + # Interleave: all vid frames in group, then audio for group + for g in range(n_segs): + for f in range(frames_per_group): + multimodal_embeddings.append( + frames[g * frames_per_group + f] + ) + multimodal_embeddings.append(group_audio_embs[g]) + elif modality == "audio": + multimodal_embeddings.extend( + self._process_audio_input(multimodal_input) + ) + return tuple(multimodal_embeddings) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + **kwargs: object, + ) -> torch.Tensor | IntermediateTensors: + """Run forward pass for Qwen2.5-VL. + + Args: + input_ids: Flattened (concatenated) input_ids corresponding to a + batch. + positions: Flattened (concatenated) position ids corresponding to a + batch. **NOTE**: If mrope is enabled (default setting for + Qwen2.5-VL opensource models), the shape will be `(3, seq_len)`, + otherwise it will be `(seq_len,). + """ + + if intermediate_tensors is not None: + inputs_embeds = None + + hidden_states = self.language_model.model( + input_ids=input_ids, + positions=positions, + intermediate_tensors=intermediate_tensors, + inputs_embeds=inputs_embeds, + ) + return hidden_states + + def compute_logits( + self, + hidden_states: torch.Tensor, + ) -> torch.Tensor | None: + return self.language_model.compute_logits(hidden_states) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + audio_loaded: set[str] = set() + + loader = AutoWeightsLoader(self, skip_prefixes=["audio_tokenizer."]) + auto_loaded = loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) + return audio_loaded | auto_loaded diff --git a/vllm/model_executor/models/registry.py b/vllm/model_executor/models/registry.py index 01f357a4993..5a3eb2edbb4 100644 --- a/vllm/model_executor/models/registry.py +++ b/vllm/model_executor/models/registry.py @@ -171,7 +171,8 @@ _TEXT_GENERATION_MODELS = { "MptForCausalLM": ("mpt", "MPTForCausalLM"), "MPTForCausalLM": ("mpt", "MPTForCausalLM"), "MiMoForCausalLM": ("mimo", "MiMoForCausalLM"), - "MiMoV2FlashForCausalLM": ("mimo_v2_flash", "MiMoV2FlashForCausalLM"), + "MiMoV2FlashForCausalLM": ("mimo_v2", "MiMoV2FlashForCausalLM"), + "MiMoV2ProForCausalLM": ("mimo_v2", "MiMoV2ProForCausalLM"), "NemotronForCausalLM": ("nemotron", "NemotronForCausalLM"), "NemotronHForCausalLM": ("nemotron_h", "NemotronHForCausalLM"), "NemotronHPuzzleForCausalLM": ("nemotron_h", "NemotronHForCausalLM"), @@ -468,6 +469,7 @@ _MULTIMODAL_MODELS = { ), "MantisForConditionalGeneration": ("llava", "MantisForConditionalGeneration"), "MiDashengLMModel": ("midashenglm", "MiDashengLMModel"), + "MiMoV2OmniForCausalLM": ("mimo_v2_omni", "MiMoV2OmniForCausalLM"), "MiniMaxVL01ForConditionalGeneration": ( "minimax_vl_01", "MiniMaxVL01ForConditionalGeneration", @@ -570,6 +572,8 @@ _MULTIMODAL_MODELS = { _SPECULATIVE_DECODING_MODELS = { "ExtractHiddenStatesModel": ("extract_hidden_states", "ExtractHiddenStatesModel"), "MiMoMTPModel": ("mimo_mtp", "MiMoMTP"), + "MiMoV2MTPModel": ("mimo_v2_mtp", "MiMoV2MTP"), + "MiMoV2OmniMTPModel": ("mimo_v2_mtp", "MiMoV2OmniMTP"), "EagleLlamaForCausalLM": ("llama_eagle", "EagleLlamaForCausalLM"), "EagleLlama4ForCausalLM": ("llama4_eagle", "EagleLlama4ForCausalLM"), "EagleMiniCPMForCausalLM": ("minicpm_eagle", "EagleMiniCPMForCausalLM"), diff --git a/vllm/transformers_utils/configs/mimo_v2_omni.py b/vllm/transformers_utils/configs/mimo_v2_omni.py new file mode 100644 index 00000000000..b87ca22a9a8 --- /dev/null +++ b/vllm/transformers_utils/configs/mimo_v2_omni.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + + +from transformers import PretrainedConfig + + +class Mimo_VLVisionConfig(PretrainedConfig): + model_type = "mimovl" + base_config_key = "vision_config" + + def __init__( + self, + depth=28, + hidden_size=1280, + hidden_act="silu", + intermediate_size=4608, + num_heads=32, + in_channels=3, + patch_size=16, + spatial_merge_size=2, + temporal_patch_size=2, + tokens_per_second=2, + window_size=128, + out_hidden_size=2048, + fullatt_block_indexes=None, + initializer_range=0.02, + kv_channels=64, # HACK + qk_channels=64, + num_query_groups=4, + num_key_value_heads=8, + vit_window_attn_types=None, + visual_token_window_size=64, + **kwargs, + ): + super().__init__(**kwargs) + + self.depth = depth + self.hidden_size = hidden_size + self.hidden_act = hidden_act + self.intermediate_size = intermediate_size + self.num_heads = num_heads + # Support GQA: if num_key_value_heads is not provided, + # default to num_heads (MHA) + if num_key_value_heads is None: + num_key_value_heads = num_heads + self.num_key_value_heads = num_key_value_heads + self.in_channels = in_channels + self.patch_size = patch_size + self.spatial_merge_size = spatial_merge_size + self.temporal_patch_size = temporal_patch_size + self.tokens_per_second = tokens_per_second + self.window_size = window_size + self.fullatt_block_indexes = ( + fullatt_block_indexes + if fullatt_block_indexes is not None + else [7, 15, 23, 31] + ) + self.out_hidden_size = out_hidden_size + self.initializer_range = initializer_range + self.kv_channels = kv_channels + self.qk_channels = qk_channels + self.num_query_groups = num_query_groups + self.vit_window_attn_types = vit_window_attn_types or [-1] * depth + self.visual_token_window_size = visual_token_window_size diff --git a/vllm/transformers_utils/model_arch_config_convertor.py b/vllm/transformers_utils/model_arch_config_convertor.py index 443223689a9..b3c912cf340 100644 --- a/vllm/transformers_utils/model_arch_config_convertor.py +++ b/vllm/transformers_utils/model_arch_config_convertor.py @@ -445,6 +445,37 @@ class MimoMTPModelArchConfigConvertor(ModelArchConfigConvertorBase): return getattr(self.hf_text_config, "num_nextn_predict_layers", 0) +def _strip_mimo_v2_attention_chunk_size( + hf_config: PretrainedConfig, hf_text_config: PretrainedConfig +) -> None: + # MiMo-V2-Flash's config.json sets `attention_chunk_size=128` but the + # architecture does not actually use chunked local attention. Leaving it + # set makes vLLM disable the hybrid KV cache manager + for cfg in (hf_text_config, hf_config): + if cfg is not None and hasattr(cfg, "attention_chunk_size"): + delattr(cfg, "attention_chunk_size") + + +class MimoV2ModelArchConfigConvertor(ModelArchConfigConvertorBase): + def __init__(self, hf_config: PretrainedConfig, hf_text_config: PretrainedConfig): + super().__init__(hf_config, hf_text_config) + _strip_mimo_v2_attention_chunk_size(hf_config, hf_text_config) + + +class MimoV2MTPModelArchConfigConvertor(ModelArchConfigConvertorBase): + def __init__(self, hf_config: PretrainedConfig, hf_text_config: PretrainedConfig): + super().__init__(hf_config, hf_text_config) + _strip_mimo_v2_attention_chunk_size(hf_config, hf_text_config) + + def get_num_hidden_layers(self) -> int: + n = getattr(self.hf_text_config, "num_nextn_predict_layers", None) + if n is not None: + return n + # Fall back to n_predict set by hf_config_override + n = getattr(self.hf_text_config, "n_predict", None) + return n if n is not None else 0 + + class GLM4MoeMTPModelArchConfigConvertor(ModelArchConfigConvertorBase): def get_num_hidden_layers(self) -> int: return getattr(self.hf_text_config, "num_nextn_predict_layers", 0) @@ -511,6 +542,10 @@ MODEL_ARCH_CONFIG_CONVERTORS = { "qwen3_next_mtp": Qwen3NextMTPModelArchConfigConvertor, "qwen3_5_mtp": Qwen3_5MTPModelArchConfigConvertor, "mimo_mtp": MimoMTPModelArchConfigConvertor, + "mimo_v2_pro": MimoV2ModelArchConfigConvertor, + "mimo_v2_flash": MimoV2ModelArchConfigConvertor, + "mimo_v2_mtp": MimoV2MTPModelArchConfigConvertor, + "mimo_v2_omni_mtp": MimoV2MTPModelArchConfigConvertor, "glm4_moe_mtp": GLM4MoeMTPModelArchConfigConvertor, "glm_ocr_mtp": GLM4MoeMTPModelArchConfigConvertor, "ernie_mtp": ErnieMTPModelArchConfigConvertor, diff --git a/vllm/transformers_utils/processors/__init__.py b/vllm/transformers_utils/processors/__init__.py index c1fe9eaf934..546a5c45329 100644 --- a/vllm/transformers_utils/processors/__init__.py +++ b/vllm/transformers_utils/processors/__init__.py @@ -27,6 +27,7 @@ __all__ = [ "IsaacProcessor", "KimiAudioProcessor", "KimiK25Processor", + "MiMoOmniProcessor", "MistralCommonPixtralProcessor", "MistralCommonVoxtralProcessor", "NanoNemotronVLProcessor", @@ -57,6 +58,7 @@ _CLASS_TO_MODULE: dict[str, str] = { "IsaacProcessor": "vllm.transformers_utils.processors.isaac", "KimiAudioProcessor": "vllm.transformers_utils.processors.kimi_audio", "KimiK25Processor": "vllm.transformers_utils.processors.kimi_k25", + "MiMoOmniProcessor": "vllm.transformers_utils.processors.mimo_v2_omni", "MistralCommonPixtralProcessor": "vllm.transformers_utils.processors.pixtral", "MistralCommonVoxtralProcessor": "vllm.transformers_utils.processors.voxtral", "NanoNemotronVLProcessor": "vllm.transformers_utils.processors.nano_nemotron_vl", diff --git a/vllm/transformers_utils/processors/mimo_v2_omni.py b/vllm/transformers_utils/processors/mimo_v2_omni.py new file mode 100644 index 00000000000..97df3184113 --- /dev/null +++ b/vllm/transformers_utils/processors/mimo_v2_omni.py @@ -0,0 +1,1285 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# mypy: ignore-errors +"""MiMo-Omni multimodal processor for vLLM. + +Ported from SGLang's MiMoV2OmniProcessor / MiMoVLProcessor implementations. +""" + +import contextlib +import copy +import io +import logging +import math +from collections import OrderedDict +from concurrent.futures import ThreadPoolExecutor, as_completed +from dataclasses import dataclass, field +from io import BytesIO +from typing import Any, Literal + +import numpy as np +import regex as re +import requests +import torch +import torch.nn.functional as F +from PIL import Image +from transformers import BatchFeature, TensorType +from transformers.processing_utils import ProcessorMixin + +try: + from torchcodec.decoders import AudioDecoder + + _HAS_TORCHCODEC = True +except ImportError: + AudioDecoder = None + _HAS_TORCHCODEC = False + +try: + import torchaudio + from torchaudio.transforms import MelSpectrogram as _MelSpectrogram + + _HAS_TORCHAUDIO = True +except ImportError: + torchaudio = None # type: ignore[assignment] + _MelSpectrogram = None # type: ignore[assignment,misc] + _HAS_TORCHAUDIO = False + +logger = logging.getLogger(__name__) + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +_PIXEL_MEAN = [123.675, 116.28, 103.53] +_PIXEL_STD = [58.395, 57.12, 57.375] +_mean_std_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] = {} + + +# --------------------------------------------------------------------------- +# Data classes +# --------------------------------------------------------------------------- + + +@dataclass +class ImageInput: + # PIL.Image | str (path/url/base64) | bytes | torch.Tensor (C,H,W) + image: Any + max_pixels: int | None = None + min_pixels: int | None = None + + +@dataclass +class VideoInput: + # tuple[frames_TCHW: torch.Tensor, timestamps_T: torch.Tensor] + video: Any + min_pixels: int | None = None + max_pixels: int | None = None + total_max_pixels: int | None = None + fps: float | None = None + num_frames: int | None = None + max_frames: int | None = None + min_frames: int | None = None + do_include_last_frame: bool | None = False + start_time: float | None = None + end_time: float | None = None + segment_type: Literal["individual", "partial"] = "individual" + + +@dataclass +class AudioInput: + # str (path/url/base64) | bytes | tuple[waveform_1D, sr] + # | np.ndarray | torch.Tensor (T,n_vq) + audio: Any + + +@dataclass +class VideoAudioInput: + video: Any # same as VideoInput.video + audio: Any # same as AudioInput.audio + min_pixels: int | None = None + max_pixels: int | None = None + total_max_pixels: int | None = None + fps: float | None = None + num_frames: int | None = None + max_frames: int | None = None + min_frames: int | None = None + do_include_last_frame: bool | None = False + start_time: float | None = None + end_time: float | None = None + segment_type: Literal["individual", "partial"] = "individual" + + +@dataclass +class Content: + type: Literal["text", "image", "video", "audio", "video_audio"] + content: Any + is_target: bool | None = None + + +@dataclass +class MiMoVLInputSample: + input_ids: torch.Tensor + labels: torch.Tensor | None + pixel_values: list[torch.Tensor] + pixel_values_videos: list[torch.Tensor] + image_thw_grids: list[torch.Tensor] + video_thw_grids: list[torch.Tensor] + audio_inputs: list[torch.Tensor] + second_per_grid_ts: list[float] = field(default_factory=list) + video_start_times: list[float] = field(default_factory=list) + audio_token_lens: list[int] = field(default_factory=list) + va_audio_inputs: list[torch.Tensor] = field(default_factory=list) + video_audio_n_segs: list[int] = field(default_factory=list) + video_audio_seg_lens: list[int] = field(default_factory=list) + position_ids: torch.Tensor | None = None + rope_deltas: torch.Tensor | None = None + extra: dict = field(default_factory=dict) + + +# --------------------------------------------------------------------------- +# Vision utilities +# --------------------------------------------------------------------------- + + +def _format_timestamp(ts: float) -> str: + return f"{int(ts // 60):02d}:{int(ts % 60):02d}" + + +def _smart_resize( + h: int, w: int, factor: int, min_px: int, max_px: int +) -> tuple[int, int]: + if min(h, w) < factor: + if h < w: + h, w = factor, int(w * factor / h) + else: + w, h = factor, int(h * factor / w) + elif max(h, w) / min(h, w) > 200: + raise ValueError(f"Aspect ratio > 200 not allowed: {h}x{w}") + h_bar = round(h / factor) * factor + w_bar = round(w / factor) * factor + if h_bar * w_bar > max_px: + beta = math.sqrt((h * w) / max_px) + h_bar = math.floor(h / beta / factor) * factor + w_bar = math.floor(w / beta / factor) * factor + elif h_bar * w_bar < min_px: + beta = math.sqrt(min_px / (h * w)) + h_bar = math.ceil(h * beta / factor) * factor + w_bar = math.ceil(w * beta / factor) * factor + return int(h_bar), int(w_bar) + + +def _to_rgb(img: Image.Image) -> Image.Image: + if img.mode == "RGBA": + bg = Image.new("RGB", img.size, (255, 255, 255)) + bg.paste(img, mask=img.split()[3]) + return bg + return img.convert("RGB") + + +def _standardize(images: torch.Tensor) -> torch.Tensor: + key = str(images.device) + if key not in _mean_std_cache: + mean = torch.tensor(_PIXEL_MEAN, device=images.device).view(1, -1, 1, 1) + std = torch.tensor(_PIXEL_STD, device=images.device).view(1, -1, 1, 1) + _mean_std_cache[key] = (mean, std) + mean, std = _mean_std_cache[key] + return (images - mean) / std + + +def _transform_batch( + frames: torch.Tensor, + factor: int, + min_px: int, + max_px: int, + device: torch.device | None = None, +) -> tuple[torch.Tensor, int, int]: + if device is not None: + frames = frames.to(device) + _, _, h, w = frames.shape + h_bar, w_bar = _smart_resize(h, w, factor, min_px, max_px) + resized = F.interpolate( + frames.float(), (h_bar, w_bar), mode="bilinear", align_corners=False + ) + return _standardize(resized), w_bar, h_bar + + +def _transform_single( + img: Any, + factor: int, + min_px: int, + max_px: int, + device: torch.device | None = None, +) -> tuple[torch.Tensor, int, int]: + if isinstance(img, torch.Tensor): + t = img.float() + _, h, w = t.shape + elif isinstance(img, Image.Image): + img = img.convert("RGB") + w, h = img.size + t = torch.from_numpy(np.array(img)).permute(2, 0, 1).float() + else: + raise TypeError(f"Expected Tensor or PIL.Image, got {type(img)}") + if device is not None: + t = t.to(device) + h_bar, w_bar = _smart_resize(h, w, factor, min_px, max_px) + out = F.interpolate( + t.unsqueeze(0), (h_bar, w_bar), mode="bilinear", align_corners=False + ) + return _standardize(out).squeeze(0), w_bar, h_bar + + +def _fetch_image(src: Any) -> Image.Image: + if isinstance(src, Image.Image): + return _to_rgb(src) + if isinstance(src, bytes): + return _to_rgb(copy.deepcopy(Image.open(BytesIO(src)))) + if isinstance(src, str): + if src.startswith(("http://", "https://")): + r = requests.get(src, timeout=30) + r.raise_for_status() + return _to_rgb(copy.deepcopy(Image.open(BytesIO(r.content)))) + if src.startswith("file://"): + return _to_rgb(Image.open(src[7:])) + if src.startswith("data:image"): + import pybase64 as _b64 + + _, b64 = src.split("base64,", 1) + return _to_rgb(copy.deepcopy(Image.open(BytesIO(_b64.b64decode(b64))))) + return _to_rgb(Image.open(src)) + raise ValueError(f"Unrecognized image source: {type(src)}") + + +# --------------------------------------------------------------------------- +# Core processor +# --------------------------------------------------------------------------- + + +class MiMoVLProcessor: + """Core MiMo-VL multimodal processor. + + Handles image/video/audio preprocessing and token sequence construction. + Ported from SGLang's MiMoVLProcessor. + """ + + def __init__( + self, + tokenizer: Any, + patch_size: int = 14, + merge_size: int = 2, + temporal_patch_size: int = 2, + temporal_compression_ratio: int = 1, + use_video_timestamps: bool = True, + video_audio_interleave_length: int = 0, + audio_kernel_size: int = 3, + audio_stride_size: int = 2, + audio_avg_pooler: int = 2, + audio_sampling_rate: int = 24000, + audio_nfft: int = 960, + audio_hop_length: int = 240, + audio_window_size: int = 960, + audio_fmin: float = 0.0, + audio_fmax: float | None = None, + audio_n_mels: int = 128, + audio_segment_size: int = 6000, + audio_channels: int = 8, + audio_group_size: int = 4, + audio_input_id_per_second: float = 25.0, + audio_zeroemb_idx: int = 4096, + image_min_pixels: int | None = None, + image_max_pixels: int | None = None, + video_min_pixels: int | None = None, + video_max_pixels: int | None = None, + video_total_max_pixels: int | None = None, + fps: float | None = None, + num_frames: int | None = None, + max_frames: int | None = None, + min_frames: int | None = None, + image_token_id: int | None = None, + video_token_id: int | None = None, + audio_token_id: int | None = None, + vision_start_token_id: int | None = None, + vision_end_token_id: int | None = None, + audio_start_token_id: int | None = None, + audio_end_token_id: int | None = None, + video_start_token_id: int | None = None, + video_end_token_id: int | None = None, + pad_token_id: int | None = None, + rope_type: str = "rope", + video_process_num_threads: int = 16, + device: Any | None = None, + **kwargs: Any, + ) -> None: + self.tokenizer = tokenizer + self.video_process_num_threads = video_process_num_threads + self.device = torch.device(device) if isinstance(device, str) else device + + self.rope_type = "rope" if rope_type == "1d" else rope_type + assert self.rope_type in ("rope", "mrope"), ( + f"Unknown rope_type: {self.rope_type}" + ) + + # video timestamps require 1-D rope + assert use_video_timestamps, "use_video_timestamps must be True" + assert self.rope_type == "rope", ( + "use_video_timestamps requires rope_type='rope'" + ) + self.use_video_timestamps = use_video_timestamps + self.video_audio_interleave_length = video_audio_interleave_length + + self.image_token_id = image_token_id + self.video_token_id = video_token_id + self.audio_token_id = audio_token_id + self.vision_start_token_id = vision_start_token_id + self.vision_end_token_id = vision_end_token_id + self.audio_start_token_id = audio_start_token_id + self.audio_end_token_id = audio_end_token_id + self.video_start_token_id = video_start_token_id + self.video_end_token_id = video_end_token_id + self.pad_token_id = pad_token_id + + self.patch_size = patch_size + self.merge_size = merge_size + self.temporal_patch_size = temporal_patch_size + self.temporal_compression_ratio = temporal_compression_ratio + + self.audio_sampling_rate = audio_sampling_rate + self.audio_nfft = audio_nfft + self.audio_hop_length = audio_hop_length + self.audio_window_size = audio_window_size + self.audio_fmin = audio_fmin + self.audio_fmax = audio_fmax + self.audio_n_mels = audio_n_mels + self.audio_segment_size = audio_segment_size + self.audio_kernel_size = audio_kernel_size + self.audio_stride_size = audio_stride_size + self.audio_avg_pooler = audio_avg_pooler + self.audio_channels = audio_channels + self.audio_group_size = audio_group_size + self.audio_input_id_per_second = audio_input_id_per_second + + self._mel_spec_kwargs = dict( + sample_rate=audio_sampling_rate, + n_fft=audio_nfft, + hop_length=audio_hop_length, + win_length=audio_window_size, + f_min=audio_fmin, + f_max=audio_fmax, + n_mels=audio_n_mels, + power=1.0, + center=True, + ) + self._mel_spectrogram: Any | None = None + self._resamplers: OrderedDict = OrderedDict() + self._resamplers_max = 16 + + if isinstance(audio_zeroemb_idx, int): + self.audio_zeroemb_idxs = torch.tensor( + [audio_zeroemb_idx] * audio_channels, dtype=torch.int32 + ) + else: + self.audio_zeroemb_idxs = torch.tensor(audio_zeroemb_idx, dtype=torch.int32) + + assert image_min_pixels is not None, "image_min_pixels must be set" + assert image_max_pixels is not None, "image_max_pixels must be set" + assert video_min_pixels is not None, "video_min_pixels must be set" + assert video_max_pixels is not None, "video_max_pixels must be set" + assert video_total_max_pixels is not None, "video_total_max_pixels must be set" + assert fps is not None or num_frames is not None, ( + "fps or num_frames must be set" + ) + + self._img_kw = {"min_pixels": image_min_pixels, "max_pixels": image_max_pixels} + self._vid_kw = { + "min_pixels": video_min_pixels, + "max_pixels": video_max_pixels, + "total_max_pixels": video_total_max_pixels, + "fps": fps, + "num_frames": num_frames, + "max_frames": max_frames, + "min_frames": min_frames, + } + + @property + def mel_spectrogram(self) -> Any: + if self._mel_spectrogram is None: + if _MelSpectrogram is None: + raise RuntimeError( + "torchaudio is required for audio. " + "Install with: pip install torchaudio" + ) + self._mel_spectrogram = _MelSpectrogram(**self._mel_spec_kwargs) + return self._mel_spectrogram + + def _resolve_img_kw(self, img: ImageInput) -> dict: + return { + "min_px": ( + img.min_pixels + if img.min_pixels is not None + else self._img_kw["min_pixels"] + ), + "max_px": ( + img.max_pixels + if img.max_pixels is not None + else self._img_kw["max_pixels"] + ), + } + + def _resolve_vid_kw(self, vid: VideoInput) -> dict: + kw: dict = {} + for k in ("min_pixels", "max_pixels", "total_max_pixels"): + kw[k] = getattr(vid, k) or self._vid_kw[k] + if vid.num_frames is not None: + kw["num_frames"] = vid.num_frames + elif vid.fps is not None: + kw["fps"] = vid.fps + if vid.max_frames is not None: + kw["max_frames"] = vid.max_frames + if vid.min_frames is not None: + kw["min_frames"] = vid.min_frames + elif self._vid_kw["num_frames"] is not None: + kw["num_frames"] = self._vid_kw["num_frames"] + elif self._vid_kw["fps"] is not None: + kw["fps"] = self._vid_kw["fps"] + if self._vid_kw["max_frames"] is not None: + kw["max_frames"] = self._vid_kw["max_frames"] + if self._vid_kw["min_frames"] is not None: + kw["min_frames"] = self._vid_kw["min_frames"] + else: + raise ValueError( + "No video sampling strategy specified (fps or num_frames)." + ) + return kw + + def preprocess_audio(self, audio: Any) -> tuple[torch.Tensor, int]: + """Decode audio bytes/path/tuple → (mel_spec (T, n_mels), token_len).""" + if isinstance(audio, tuple): + waveform, original_sr = audio + else: + if AudioDecoder is None: + raise RuntimeError( + "torchcodec is required for audio. " + "Install with: pip install torchcodec" + ) + if isinstance(audio, bytes): + file_obj: Any = io.BytesIO(audio) + elif isinstance(audio, str): + if audio.startswith("data:"): + import pybase64 as _b64 + + file_obj = io.BytesIO(_b64.b64decode(audio.split(",")[1])) + elif audio.startswith(("http://", "https://")): + r = requests.get(audio, timeout=30) + r.raise_for_status() + file_obj = io.BytesIO(r.content) + else: + file_obj = audio + else: + raise ValueError(f"Unsupported audio source type: {type(audio)}") + samples = AudioDecoder(file_obj).get_all_samples() + waveform = samples.data + original_sr = samples.sample_rate + + if original_sr != self.audio_sampling_rate: + if original_sr not in self._resamplers: + if len(self._resamplers) >= self._resamplers_max: + self._resamplers.popitem(last=False) + self._resamplers[original_sr] = torchaudio.transforms.Resample( + orig_freq=original_sr, new_freq=self.audio_sampling_rate + ) + self._resamplers.move_to_end(original_sr) + waveform = self._resamplers[original_sr](waveform) + + if waveform.ndim == 2: + waveform = waveform.mean(dim=0) + spec = self.mel_spectrogram(waveform[None, :]) + spec = torch.log(torch.clip(spec, min=1e-7)).squeeze().transpose(0, 1) + + n = spec.shape[0] + n = n + 3 - self.audio_kernel_size + n = (n + 2 - self.audio_kernel_size) // self.audio_stride_size + 1 + n = n // self.audio_avg_pooler + int(n % self.audio_avg_pooler != 0) + token_len = math.ceil(n / self.audio_group_size) + return spec, token_len + + def process_image(self, image: ImageInput) -> torch.Tensor: + kw = self._resolve_img_kw(image) + src = image.image + if isinstance(src, (str, bytes)): + src = _fetch_image(src) + tensor, _, _ = _transform_single( + src, + factor=self.patch_size * self.merge_size, + device=self.device, + **kw, + ) + return tensor + + def process_video( + self, video_input: VideoInput + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict]: + kw = self._resolve_vid_kw(video_input) + video = video_input.video + if not isinstance(video, tuple): + raise ValueError( + f"video must be a (frames_TCHW, timestamps_T) tuple, " + f"got {type(video)}. " + "Decode the video before calling the processor." + ) + frames, timestamps = video + + fps = ( + 1.0 + if len(timestamps) < 2 + else float(1.0 / (float(timestamps[1]) - float(timestamps[0]))) + ) + start = ( + video_input.start_time + if video_input.start_time is not None + else float(timestamps[0]) + ) + end = ( + video_input.end_time + if video_input.end_time is not None + else float(timestamps[-1]) + 1.0 / fps + ) + + if video_input.segment_type != "individual": + mask = (timestamps >= start) & (timestamps < end) + idxs = torch.where(mask)[0] + if len(idxs) == 0: + idxs = torch.where(timestamps <= start)[0][-1:] + frames, timestamps = frames[idxs], timestamps[idxs] + + tp = self.temporal_patch_size * self.temporal_compression_ratio + n = frames.shape[0] + total_px = kw["total_max_pixels"] + max_px = max( + kw["min_pixels"], min(total_px * tp // max(n, 1), kw["max_pixels"]) + ) + + if n % tp != 0: + pad = tp - n % tp + frames = torch.cat( + [frames, frames[-1:].repeat(pad, *([1] * (frames.ndim - 1)))], + dim=0, + ) + timestamps = torch.cat([timestamps, timestamps[-1:].repeat(pad)], dim=0) + + transformed, _, _ = _transform_batch( + frames, + factor=self.patch_size * self.merge_size, + min_px=kw["min_pixels"], + max_px=max_px, + device=self.device, + ) + patches, thw = self._flatten_visual(transformed, "video") + meta = { + "fps_sampled": fps, + "segment_start_time": start, + "segment_end_time": end, + } + return patches, thw, timestamps, meta + + def process_audio(self, audio: AudioInput) -> Any: + src = audio.audio + if isinstance(src, np.ndarray): + src = (torch.from_numpy(src).float(), self.audio_sampling_rate) + if isinstance(src, (str, bytes, tuple)): + return self.preprocess_audio(src) + # Pre-tokenized tensor (T, n_vq) + assert isinstance(src, torch.Tensor) and src.ndim == 2 + T = src.shape[0] + src = src[:, : self.audio_channels].to(torch.long) + pad_T = ( + (T + self.audio_group_size - 1) + // self.audio_group_size + * self.audio_group_size + ) + padding = ( + torch.zeros(pad_T - T, self.audio_channels, dtype=torch.long) + src[-1] + ) + src = torch.cat([src, padding], dim=0) + return src.reshape( + pad_T // self.audio_group_size, self.audio_group_size, self.audio_channels + ) + + def _flatten_visual( + self, visual: torch.Tensor, kind: str + ) -> tuple[torch.Tensor, torch.Tensor]: + if kind == "image": + h, w = visual.shape[-2:] + patches = visual.unsqueeze(0).repeat(self.temporal_patch_size, 1, 1, 1) + else: # video / video_audio + temporal_stride = self.temporal_compression_ratio * self.temporal_patch_size + assert visual.shape[0] % temporal_stride == 0 + patches = visual + h, w = patches.shape[-2:] + + C = patches.shape[1] + grid_t = patches.shape[0] // self.temporal_patch_size + grid_h, grid_w = h // self.patch_size, w // self.patch_size + + patches = ( + patches.contiguous() + .view( + grid_t, + self.temporal_patch_size, + C, + grid_h // self.merge_size, + self.merge_size, + self.patch_size, + grid_w // self.merge_size, + self.merge_size, + self.patch_size, + ) + .permute(0, 3, 6, 4, 7, 2, 1, 5, 8) + .contiguous() + .view( + grid_t * grid_h * grid_w, + C * self.temporal_patch_size * self.patch_size * self.patch_size, + ) + ) + thw = torch.tensor([grid_t, grid_h, grid_w], dtype=torch.int32) + return patches, thw + + def process( + self, contents: list[Content], verbose: bool = False + ) -> MiMoVLInputSample: + input_ids: list[int] = [] + labels: list[int] = [] + img_pv: list[torch.Tensor] = [] + img_grids: list[torch.Tensor] = [] + vid_pv: list[torch.Tensor] = [] + vid_grids: list[torch.Tensor] = [] + audio_inputs: list[torch.Tensor] = [] + is_audio_tokenized: list[bool] = [] + audio_token_lens: list[int] = [] + second_per_grid_ts: list[float] = [] + video_start_times: list[float] = [] + va_audio_inputs: list[torch.Tensor] = [] + video_audio_n_segs: list[int] = [] + video_audio_seg_lens: list[int] = [] + + # Pre-decode videos in parallel + vid_info = [ + (i, c.content, c.type == "video_audio") + for i, c in enumerate(contents) + if c.type in ("video", "video_audio") + ] + vid_results: dict[int, tuple] = {} + if vid_info: + n_t = min(self.video_process_num_threads, len(vid_info)) + if n_t > 1 and len(vid_info) > 1: + with ThreadPoolExecutor(max_workers=n_t) as ex: + fut_map = { + ex.submit(self.process_video, vi): idx + for idx, vi, _ in vid_info + } + for fut in as_completed(fut_map): + vid_results[fut_map[fut]] = fut.result() + else: + for idx, vi, _ in vid_info: + vid_results[idx] = self.process_video(vi) + + for ci, content in enumerate(contents): + _ids: list[int] = [] + _lbls: list[int] | None = None + + if content.type == "text": + _ids = ( + self.tokenizer.encode(content.content) + if isinstance(content.content, str) + else list(content.content) + ) + if content.is_target: + _lbls = _ids + + elif content.type == "image": + tensor = self.process_image(content.content) + patches, thw = self._flatten_visual(tensor, "image") + t, h, w = thw.tolist() + n_tok = (t * h * w) // (self.merge_size**2) + img_pv.append(patches) + img_grids.append(thw) + _ids = ( + [self.vision_start_token_id] + + [self.image_token_id] * n_tok + + [self.vision_end_token_id] + ) + + elif content.type == "video": + patches, thw, ts, meta = vid_results[ci] + t, h, w = thw.tolist() + n_per_grid = h * w // (self.merge_size**2) + vid_pv.append(patches) + vid_grids.append(thw) + second_per_grid_ts.append( + self.temporal_patch_size / meta["fps_sampled"] + ) + video_start_times.append(float(ts[0])) + video_audio_n_segs.append(0) + + stride = self.temporal_patch_size * self.temporal_compression_ratio + ts_texts = [_format_timestamp(float(x)) for x in ts[::stride]] + ts_ids_list = [self.tokenizer.encode(s) for s in ts_texts] + + _ids = [self.video_start_token_id] + for ts_ids in ts_ids_list: + _ids += ( + ts_ids + + [self.vision_start_token_id] + + [self.video_token_id] * n_per_grid + + [self.vision_end_token_id] + ) + _ids += [self.video_end_token_id] + + elif content.type == "audio": + processed = self.process_audio(content.content) + if isinstance(processed, tuple): + is_audio_tokenized.append(False) + spec, tok_len = processed + audio_inputs.append(spec) + else: + is_audio_tokenized.append(True) + tok_len = processed.shape[0] + audio_inputs.append(processed) + audio_token_lens.append(tok_len) + _ids = ( + [self.audio_start_token_id] + + [self.audio_token_id] * tok_len + + [self.audio_end_token_id] + ) + + elif content.type == "video_audio": + patches, thw, ts, meta = vid_results[ci] + second_per_grid_ts.append( + self.temporal_patch_size / meta["fps_sampled"] + ) + video_start_times.append(float(ts[0])) + processed_audio = self.process_audio(content.content) + tok_per_sec = self.audio_input_id_per_second / self.audio_group_size + + t, h, w = thw.tolist() + vid_pv.append(patches) + vid_grids.append(thw) + + if isinstance(processed_audio, tuple): + # Mel spec (not pre-tokenized): store in va_audio_inputs separately + spec, total_atok = processed_audio + va_audio_inputs.append(spec) + _va_is_tokenized = False + else: + # Pre-tokenized: not expected in vLLM, but handle defensively + total_atok = processed_audio.shape[0] + _va_is_tokenized = True + + n_per_grid = h * w // (self.merge_size**2) + stride = self.temporal_patch_size * self.temporal_compression_ratio + grid_ts = ts[::stride] + ts_texts = [_format_timestamp(float(x)) for x in grid_ts] + ts_ids_list = [self.tokenizer.encode(s) for s in ts_texts] + + units: list[tuple] = [] + for i in range(len(grid_ts)): + a_start = int(float(grid_ts[i]) * tok_per_sec) + a_end = ( + int(float(grid_ts[i + 1]) * tok_per_sec) + if i < len(grid_ts) - 1 + else int(meta["segment_end_time"] * tok_per_sec) + ) + seg_len = min(a_end, total_atok) - a_start + assert seg_len > 0, f"Zero-length audio segment at grid index {i}" + seg = ( + processed_audio[a_start : a_start + seg_len] + if _va_is_tokenized + else None + ) + units.append( + ( + float(grid_ts[i]), + ts_texts[i], + ts_ids_list[i], + n_per_grid, + seg_len, + seg, + ) + ) + + il = self.video_audio_interleave_length + if il == -1: + groups: list[list] = [list(enumerate(units))] + elif il == 0: + groups = [[(i, u)] for i, u in enumerate(units)] + else: + groups, cur, t_ptr = [], [], 0.0 + for i, u in enumerate(units): + while u[0] >= t_ptr + il: + if cur: + groups.append(cur) + cur = [] + t_ptr += il + cur.append((i, u)) + if cur: + groups.append(cur) + + # Track n_segs (= num groups) and per-group audio token counts + video_audio_n_segs.append(len(groups)) + for group in groups: + group_seg_len = sum(u[4] for _, u in group) + video_audio_seg_lens.append(group_seg_len) + + _ids = [self.video_start_token_id] + for group in groups: + _ids += group[0][1][2] # first-unit timestamp token ids + _vid_tok: list[int] = [] + _aud_tok: list[int] = [] + for _, u in group: + _, _, _, vid_n, seg_n, seg_audio = u + _vid_tok += ( + [self.vision_start_token_id] + + [self.video_token_id] * vid_n + + [self.vision_end_token_id] + ) + _aud_tok += [self.audio_token_id] * seg_n + if seg_audio is not None: + # Pre-tokenized per-frame segments (rare in vLLM) + audio_inputs.append(seg_audio) + _ids += ( + _vid_tok + + [self.audio_start_token_id] + + _aud_tok + + [self.audio_end_token_id] + ) + _ids += [self.video_end_token_id] + + input_ids.extend(_ids) + labels.extend( + _lbls if _lbls is not None else [self.pad_token_id] * len(_ids) + ) + + ids_t = torch.tensor(input_ids) + lbl_arr = np.roll(labels, shift=-1) + lbl_arr[-1] = self.pad_token_id + lbl_t = torch.tensor(lbl_arr) + + extra: dict = {} + if is_audio_tokenized: + assert all(is_audio_tokenized) or not any(is_audio_tokenized) + extra["is_audio_tokenized"] = is_audio_tokenized[0] + + position_ids = torch.arange(ids_t.shape[0]).expand(3, -1) + rope_deltas = torch.zeros((1, 1), dtype=torch.int32) + + return MiMoVLInputSample( + input_ids=ids_t, + labels=lbl_t, + pixel_values=img_pv, + pixel_values_videos=vid_pv, + image_thw_grids=img_grids, + video_thw_grids=vid_grids, + audio_inputs=audio_inputs, + second_per_grid_ts=second_per_grid_ts, + video_start_times=video_start_times, + audio_token_lens=audio_token_lens, + va_audio_inputs=va_audio_inputs, + video_audio_n_segs=video_audio_n_segs, + video_audio_seg_lens=video_audio_seg_lens, + position_ids=position_ids, + rope_deltas=rope_deltas, + extra=extra, + ) + + +# --------------------------------------------------------------------------- +# vLLM ProcessorMixin wrapper +# --------------------------------------------------------------------------- + + +class MiMoOmniProcessor(ProcessorMixin): + """HuggingFace-compatible ProcessorMixin wrapper for MiMo-Omni. + + Accepts PIL images, pre-decoded video tuples (frames_TCHW, timestamps_T), + and audio (file path / bytes / (waveform, sr) tuple / numpy array). + """ + + attributes = ["tokenizer"] + tokenizer_class = "AutoTokenizer" + + # Single or multi-pad placeholders produced by the chat template / prior expansion + _IMG_RE = re.compile(r"<\|vision_start\|>(?:<\|image_pad\|>)+<\|vision_end\|>") + _VID_RE = re.compile(r"<\|vision_start\|>(?:<\|video_pad\|>)+<\|vision_end\|>") + _AUD_RE = re.compile( + r"<\|mimo_audio_start\|>(?:<\|audio_pad\|>)+<\|mimo_audio_end\|>" + ) + + _MM_RE = re.compile( + r"(<\|vision_start\|>(?:<\|image_pad\|>)+<\|vision_end\|>" + r"|<\|vision_start\|>(?:<\|video_pad\|>)+<\|vision_end\|>" + r"|<\|mimo_audio_start\|>(?:<\|audio_pad\|>)+<\|mimo_audio_end\|>)" + ) + + def __init__( + self, + tokenizer: Any, + *, + patch_size: int = 14, + merge_size: int = 2, + temporal_patch_size: int = 2, + temporal_compression_ratio: int = 1, + image_min_pixels: int | None = None, + image_max_pixels: int | None = None, + video_min_pixels: int | None = None, + video_max_pixels: int | None = None, + video_total_max_pixels: int | None = None, + fps: float = 2.0, + num_frames: int | None = None, + max_frames: int = 256, + min_frames: int = 8, + video_audio_interleave_length: int = 0, + audio_sampling_rate: int = 24000, + audio_nfft: int = 960, + audio_hop_length: int = 240, + audio_window_size: int = 960, + audio_fmin: float = 0.0, + audio_fmax: float | None = None, + audio_n_mels: int = 128, + audio_segment_size: int = 6000, + audio_kernel_size: int = 3, + audio_stride_size: int = 2, + audio_avg_pooler: int = 2, + audio_channels: int = 8, + audio_group_size: int = 4, + audio_input_id_per_second: float = 25.0, + audio_zeroemb_idx: int = 4096, + image_token_id: int | None = None, + video_token_id: int | None = None, + audio_token_id: int | None = None, + vision_start_token_id: int | None = None, + vision_end_token_id: int | None = None, + audio_start_token_id: int | None = None, + audio_end_token_id: int | None = None, + video_start_token_id: int | None = None, + video_end_token_id: int | None = None, + rope_type: str = "rope", + ) -> None: + self.tokenizer = tokenizer + + unit = patch_size * merge_size + self.mimo_processor = MiMoVLProcessor( + tokenizer=tokenizer, + patch_size=patch_size, + merge_size=merge_size, + temporal_patch_size=temporal_patch_size, + temporal_compression_ratio=temporal_compression_ratio, + use_video_timestamps=True, + video_audio_interleave_length=video_audio_interleave_length, + audio_sampling_rate=audio_sampling_rate, + audio_nfft=audio_nfft, + audio_hop_length=audio_hop_length, + audio_window_size=audio_window_size, + audio_fmin=audio_fmin, + audio_fmax=audio_fmax, + audio_n_mels=audio_n_mels, + audio_segment_size=audio_segment_size, + audio_kernel_size=audio_kernel_size, + audio_stride_size=audio_stride_size, + audio_avg_pooler=audio_avg_pooler, + audio_channels=audio_channels, + audio_group_size=audio_group_size, + audio_input_id_per_second=audio_input_id_per_second, + audio_zeroemb_idx=audio_zeroemb_idx, + image_min_pixels=image_min_pixels or (4 * unit * unit), + image_max_pixels=image_max_pixels or (4096 * unit * unit), + video_min_pixels=video_min_pixels or (4 * unit * unit), + video_max_pixels=video_max_pixels or (4096 * unit * unit), + video_total_max_pixels=video_total_max_pixels or (16384 * unit * unit), + fps=fps, + num_frames=num_frames, + max_frames=max_frames, + min_frames=min_frames, + image_token_id=image_token_id, + video_token_id=video_token_id, + audio_token_id=audio_token_id, + vision_start_token_id=vision_start_token_id, + vision_end_token_id=vision_end_token_id, + audio_start_token_id=audio_start_token_id, + audio_end_token_id=audio_end_token_id, + video_start_token_id=video_start_token_id, + video_end_token_id=video_end_token_id, + pad_token_id=tokenizer.pad_token_id, + rope_type=rope_type, + ) + + @classmethod + def from_hf_config(cls, tokenizer: Any, hf_config: Any) -> "MiMoOmniProcessor": + """Convenience factory: instantiate directly from an HF model config object.""" + vc = hf_config.vision_config + if isinstance(vc, dict): + patch_size = vc.get("patch_size", 14) + merge_size = vc.get("spatial_merge_size", 2) + temporal_patch_size = vc.get("temporal_patch_size", 2) + else: + patch_size = getattr(vc, "patch_size", 14) + merge_size = getattr(vc, "spatial_merge_size", 2) + temporal_patch_size = getattr(vc, "temporal_patch_size", 2) + + pc: dict = getattr(hf_config, "processor_config", {}) or {} + ac = getattr(hf_config, "audio_config", None) + audio_sr: int | None = pc.get("audio_sampling_rate") + if audio_sr is None and ac is not None: + if isinstance(ac, dict): + audio_sr = ac.get("sampling_rate") or ac.get("sample_rate") + else: + audio_sr = getattr(ac, "sampling_rate", None) or getattr( + ac, "sample_rate", None + ) + + rope_type = "rope" + rs = getattr(hf_config, "rope_scaling", None) + if rs and rs.get("type") == "default" and rs.get("mrope_section") is not None: + rope_type = "mrope" + + unit = patch_size * merge_size + return cls( + tokenizer, + patch_size=patch_size, + merge_size=merge_size, + temporal_patch_size=temporal_patch_size, + image_min_pixels=pc.get("image_min_pixels") or (4 * unit * unit), + image_max_pixels=pc.get("image_max_pixels") or (4096 * unit * unit), + video_min_pixels=pc.get("video_min_pixels") or (4 * unit * unit), + video_max_pixels=pc.get("video_max_pixels") or (4096 * unit * unit), + video_total_max_pixels=( + pc.get("video_total_max_pixels") or (16384 * unit * unit) + ), + fps=pc.get("fps") or 2.0, + num_frames=pc.get("num_frames"), + max_frames=pc.get("max_frames") or 256, + min_frames=pc.get("min_frames") or 8, + video_audio_interleave_length=pc.get("video_audio_interleave_length", 0), + audio_sampling_rate=audio_sr or 24000, + image_token_id=pc.get("image_token_id"), + video_token_id=pc.get("video_token_id"), + audio_token_id=pc.get("audio_token_id"), + vision_start_token_id=pc.get("vision_start_token_id"), + vision_end_token_id=pc.get("vision_end_token_id"), + audio_start_token_id=pc.get("audio_start_token_id"), + audio_end_token_id=pc.get("audio_end_token_id"), + video_start_token_id=pc.get("video_start_token_id"), + video_end_token_id=pc.get("video_end_token_id"), + rope_type=rope_type, + ) + + @property + def image_token(self) -> str: + """Token string used as image placeholder (for vLLM integration).""" + return "<|image_pad|>" + + @property + def video_token(self) -> str: + """Token string used as video placeholder (for vLLM integration).""" + return "<|video_pad|>" + + @property + def image_processor(self) -> Any: + """Minimal image-processor-like object for vLLM processing-info compat.""" + p = self.mimo_processor + + class _ImageProcessor: + merge_size = p.merge_size + size = { + "shortest_edge": p._img_kw["min_pixels"], + "longest_edge": p._img_kw["max_pixels"], + } + + return _ImageProcessor() + + def _modality(self, token: str) -> str: + if self._IMG_RE.fullmatch(token): + return "image" + if self._VID_RE.fullmatch(token): + return "video" + if self._AUD_RE.fullmatch(token): + return "audio" + return "unknown" + + def __call__( + self, + text: str | list[str] | None = None, + images: Any = None, + videos: Any = None, + audio: Any = None, + video_audio: Any = None, + return_tensors: str | TensorType | None = None, + **kwargs: Any, + ) -> BatchFeature: + """Process multimodal inputs into model-ready tensors. + + Args: + text: Prompt string(s) containing multimodal placeholders + ``<|vision_start|><|image_pad|><|vision_end|>``, + ``<|vision_start|><|video_pad|><|vision_end|>``, or + ``<|mimo_audio_start|><|audio_pad|><|mimo_audio_end|>``. + images: PIL.Image or list[PIL.Image]. + videos: list of ``(frames_TCHW: torch.Tensor, timestamps_T: torch.Tensor)`` + tuples (pre-decoded). + audio: list of ``str`` (path/url/base64), ``bytes``, + ``(waveform_1D, sample_rate)`` tuples, or ``np.ndarray``. + return_tensors: Passed to :class:`BatchFeature`. + + Returns: + :class:`BatchFeature` with keys: + - ``input_ids`` + - ``pixel_values`` + ``image_grid_thw`` + - ``pixel_values_videos`` + ``video_grid_thw`` + ``second_per_grid_ts`` + - ``audio_features`` + """ + if isinstance(text, list): + text = text[0] if len(text) == 1 else "\n".join(text) + + imgs: list = ( + ([images] if isinstance(images, Image.Image) else list(images)) + if images is not None + else [] + ) + vids: list = list(videos) if videos is not None else [] + auds: list = list(audio) if audio is not None else [] + va_items: list = list(video_audio) if video_audio is not None else [] + + # If audio exists but text has no audio placeholder, prepend one + _aud_placeholder = "<|mimo_audio_start|><|audio_pad|><|mimo_audio_end|>" + if auds and text is not None and not self._AUD_RE.search(text): + text = _aud_placeholder + text + + # Build Content list + contents: list[Content] = [] + + if text and (imgs or vids or auds or va_items): + parts = self._MM_RE.split(text) + img_it = iter(imgs) + vid_it = iter(vids) + aud_it = iter(auds) + va_it = iter(va_items) + for part in parts: + if self._MM_RE.fullmatch(part): + mod = self._modality(part) + if mod == "image": + with contextlib.suppress(StopIteration): + contents.append( + Content( + type="image", + content=ImageInput(image=next(img_it)), + ) + ) + elif mod == "video": + # Try regular video first, fall back to video_audio + vid_item = None + vid_type = "video" + with contextlib.suppress(StopIteration): + vid_item = next(vid_it) + if vid_item is None: + with contextlib.suppress(StopIteration): + vid_item = next(va_it) + vid_type = "video_audio" + if vid_item is not None: + if vid_type == "video": + contents.append( + Content( + type="video", + content=VideoInput(video=vid_item), + ) + ) + else: + contents.append( + Content( + type="video_audio", + content=vid_item, + ) + ) + elif mod == "audio": + with contextlib.suppress(StopIteration): + contents.append( + Content( + type="audio", + content=AudioInput(audio=next(aud_it)), + ) + ) + elif part: + contents.append(Content(type="text", content=part)) + elif text: + contents.append(Content(type="text", content=text)) + else: + for img in imgs: + contents.append(Content(type="image", content=ImageInput(image=img))) + for vid in vids: + contents.append(Content(type="video", content=VideoInput(video=vid))) + for aud in auds: + contents.append(Content(type="audio", content=AudioInput(audio=aud))) + for va_item in va_items: + contents.append(Content(type="video_audio", content=va_item)) + + if not contents: + ids = self.tokenizer(text or "", return_tensors=return_tensors)["input_ids"] + return BatchFeature(data={"input_ids": ids}, tensor_type=return_tensors) + + sample = self.mimo_processor.process(contents, verbose=False) + + # vLLM expects input_ids to have a batch dimension [1, seq_len]. + data: dict = {"input_ids": sample.input_ids.unsqueeze(0)} + + if sample.pixel_values: + data["pixel_values"] = torch.cat(sample.pixel_values, dim=0) + data["image_grid_thw"] = torch.stack(sample.image_thw_grids) + + if sample.pixel_values_videos: + data["pixel_values_videos"] = torch.cat(sample.pixel_values_videos, dim=0) + data["video_grid_thw"] = torch.stack(sample.video_thw_grids) + if sample.second_per_grid_ts: + data["second_per_grid_ts"] = torch.tensor( + sample.second_per_grid_ts, dtype=torch.float32 + ) + if sample.video_start_times: + data["video_start_times"] = torch.tensor( + sample.video_start_times, dtype=torch.float32 + ) + if sample.video_audio_n_segs: + data["video_audio_n_segs"] = torch.tensor( + sample.video_audio_n_segs, dtype=torch.long + ) + # video_audio_seg_lens: 2D padded tensor (num_videos, max_T). + # Row i has the per-group audio token lengths for video i + # (zeros for regular videos; valid values for video_audio videos). + n_segs_list = sample.video_audio_n_segs + max_segs = max(n_segs_list) if n_segs_list else 0 + if max_segs > 0: + seg_lens_2d = torch.zeros(len(n_segs_list), max_segs, dtype=torch.long) + flat_cursor = 0 + for vi, n in enumerate(n_segs_list): + if n > 0: + seg_lens_2d[vi, :n] = torch.tensor( + sample.video_audio_seg_lens[flat_cursor : flat_cursor + n], + dtype=torch.long, + ) + flat_cursor += n + data["video_audio_seg_lens"] = seg_lens_2d + + # audio_features is a list of variable-length mel-spec tensors; pop it + # before BatchFeature conversion to avoid "batched tensors of the same + # length" errors, then re-attach it after. + audio_features = None + if sample.audio_inputs: + audio_features = sample.audio_inputs + if "is_audio_tokenized" in sample.extra: + data["is_audio_tokenized"] = sample.extra["is_audio_tokenized"] + if sample.audio_token_lens: + data["audio_token_lens"] = torch.tensor( + sample.audio_token_lens, dtype=torch.long + ) + + bf = BatchFeature(data=data, tensor_type=return_tensors) + if audio_features is not None: + bf["audio_features"] = audio_features + # va_audio_features: list of mel-spec tensors (one per video_audio item) + if sample.va_audio_inputs: + bf["va_audio_features"] = sample.va_audio_inputs + return bf diff --git a/vllm/v1/core/kv_cache_utils.py b/vllm/v1/core/kv_cache_utils.py index 879cd0928c1..3e0e7fcb8c5 100644 --- a/vllm/v1/core/kv_cache_utils.py +++ b/vllm/v1/core/kv_cache_utils.py @@ -1381,7 +1381,9 @@ def unify_hybrid_kv_cache_specs(kv_cache_spec: dict[str, KVCacheSpec]): block_size=spec.block_size, num_kv_heads=spec.num_kv_heads, head_size=spec.head_size, + head_size_v=spec.head_size_v, dtype=spec.dtype, + kv_quant_mode=spec.kv_quant_mode, sliding_window=spec.sliding_window, page_size_padded=spec.page_size_padded, ) diff --git a/vllm/v1/spec_decode/llm_base_proposer.py b/vllm/v1/spec_decode/llm_base_proposer.py index 94e09b209cf..36662d76f86 100644 --- a/vllm/v1/spec_decode/llm_base_proposer.py +++ b/vllm/v1/spec_decode/llm_base_proposer.py @@ -1344,6 +1344,7 @@ class SpecDecodeBaseProposer: "Exaone4_5_ForConditionalGeneration", "GlmOcrForConditionalGeneration", "HunYuanVLForConditionalGeneration", + "MiMoV2OmniForCausalLM", "Qwen2_5_VLForConditionalGeneration", "Qwen3_5ForConditionalGeneration", "Qwen3_5MoeForConditionalGeneration",