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",