Support MiMo v2 ASR (#26278)

Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
Co-authored-by: yanyihan <yanyihan@xiaomi.com>
Co-authored-by: zhuqingchao <zhuqingchao@xiaomi.com>
This commit is contained in:
Qingchao Zhu
2026-06-11 13:43:13 -07:00
committed by GitHub
co-authored by github-actions[bot] yanyihan zhuqingchao
parent 24c5d76f74
commit ec0eb6cce8
9 changed files with 990 additions and 433 deletions
@@ -1548,6 +1548,7 @@ multimodal_model_archs = [
"Lfm2VlForConditionalGeneration",
"LightOnOCRForConditionalGeneration",
*MIMO_V2_MULTIMODAL_ARCHS,
"MiMoV2ASRForCausalLM",
"MiniCPMO",
"MiniCPMV",
"Mistral3ForConditionalGeneration",
@@ -7,6 +7,9 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import ( # noqa:
)
# Import built-in adapters so they self-register via @register_transcription_adapter.
from sglang.srt.entrypoints.openai.transcription_adapters.mimo_v2_asr import ( # noqa: F401
MiMoV2ASRAdapter,
)
from sglang.srt.entrypoints.openai.transcription_adapters.qwen3_asr import ( # noqa: F401
Qwen3ASRAdapter,
)
@@ -20,4 +23,5 @@ __all__ = [
"resolve_adapter",
"WhisperAdapter",
"Qwen3ASRAdapter",
"MiMoV2ASRAdapter",
]
@@ -0,0 +1,46 @@
from __future__ import annotations
from sglang.srt.entrypoints.openai.protocol import (
TranscriptionRequest,
TranscriptionUsage,
TranscriptionVerboseResponse,
)
from sglang.srt.entrypoints.openai.transcription_adapters.base import (
TranscriptionAdapter,
register_transcription_adapter,
)
@register_transcription_adapter("MiMoV2ASR")
class MiMoV2ASRAdapter(TranscriptionAdapter):
"""Adapter for MiMo-V2-ASR.
The multimodal processor (``MiMoV2ASRProcessor``) prepends the audio
placeholder ``<|sosp|><|empty|>...<|eosp|>`` when ``input_text`` lacks
one, so the request text can stay empty and the adapter only has to
supply sampling params and the verbose-response shape.
"""
def build_sampling_params(self, request: TranscriptionRequest) -> dict:
return {
"temperature": request.temperature,
"max_new_tokens": 448,
}
def build_verbose_response(
self,
request: TranscriptionRequest,
text: str,
ret: dict,
tokenizer,
usage: TranscriptionUsage,
) -> TranscriptionVerboseResponse:
# MiMo-V2-ASR does not emit timestamp tokens; segments stay empty
# until a forced-aligner path is added.
return TranscriptionVerboseResponse(
language=request.language or "auto",
duration=round(request.audio_duration_s, 2),
text=text,
segments=[],
usage=usage,
)
+98 -91
View File
@@ -8,7 +8,7 @@ import os
import typing as tp
from dataclasses import dataclass
from functools import wraps
from typing import List, Optional, Tuple
from typing import Optional, Tuple
import torch
import torch.nn as nn
@@ -1141,72 +1141,89 @@ class MiMoV2AudioConfig:
return config
class MiMoAudioEncoder(nn.Module):
config: MiMoAudioEncoderConfig
class AudioEncoderMixin:
"""LM model mixin that adds MiMo audio encoder components.
def __init__(self, config):
super().__init__()
Components are attached as top-level attributes (no ``audio_encoder.``
prefix), matching the checkpoint state_dict layout. Inner naming
variations are normalized via ``AUDIO_WEIGHT_REMAP``.
Hot config fields are cached as direct ``self.audio_*`` attributes at
build time so helper methods can stay short and uniform (no
``self.audio_config.foo`` indirection inside hot paths).
Subclasses call ``self.build_audio_encoder(audio_config)`` from their
``__init__`` after the language model is constructed; the mixin's
``get_audio_feature`` then handles audio item batching end-to-end.
"""
AUDIO_WEIGHT_REMAP: tuple[tuple[str, str], ...] = (
("audio_projection", "projection"),
("speech_group_downcast", "projection"),
("audio_input_local_transformer", "input_local_transformer"),
)
def build_audio_encoder(self, config) -> None:
if not isinstance(config, MiMoV2AudioConfig):
config_dict = (
vars(config) if hasattr(config, "__dict__") else config.__dict__
)
config = MiMoV2AudioConfig(**config_dict)
self.config = config
self.server_args = get_global_server_args()
self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder
self.speech_empty_ids = self.parsed_speech_empty_ids()
cfg_dict = vars(config) if hasattr(config, "__dict__") else config.__dict__
config = MiMoV2AudioConfig(**cfg_dict)
self.audio_channels = config.audio_channels
self.audio_group_size = config.group_size
self.audio_segment_size = config.audio_segment_size
speech_vocab_size = self._parse_maybe_list(
self.config.speech_vocab_size, self.config.audio_channels
self.audio_input_local_dim = config.input_local_dim
self.audio_input_full_attention = config.input_full_attention
self.audio_out_hidden_size = config.out_hidden_size
speech_vocab_size = config._parse_maybe_list(
config.speech_vocab_size, self.audio_channels
)
speech_empty_ids = config._parse_maybe_list(
config.speech_zeroemb_idx, self.audio_channels
)
input_local_config = Qwen2Config(
hidden_size=self.config.input_local_dim,
num_hidden_layers=self.config.input_local_layers,
num_attention_heads=self.config.input_local_attn_heads,
num_key_value_heads=self.config.input_local_attn_heads,
intermediate_size=self.config.input_local_intermediate_size,
attention_dropout=self.config.input_local_hidden_dropout,
rope_theta=self.config.rope_theta,
partial_rotary_factor=self.config.partial_rotary_factor,
hidden_size=self.audio_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,
)
input_local_config.head_dim = self.config.input_local_head_dim
input_local_config.head_dim = config.input_local_head_dim
self.input_local_transformer = Qwen2Model(input_local_config)
if not self.config.add_post_norm:
if not config.add_post_norm:
self.input_local_transformer.norm = nn.Identity()
self.speech_embeddings = nn.ModuleList(
[
nn.Embedding(
speech_vocab_size[i],
self.config.input_local_dim,
padding_idx=self.speech_empty_ids[i],
self.audio_input_local_dim,
padding_idx=speech_empty_ids[i],
)
for i in range(self.config.audio_channels)
for i in range(self.audio_channels)
]
)
if self.config.projection_layers == 1:
if config.projection_layers == 1:
self.projection = nn.Linear(
self.config.input_local_dim * self.config.group_size,
self.config.out_hidden_size,
self.audio_input_local_dim * self.audio_group_size,
self.audio_out_hidden_size,
bias=False,
)
elif self.config.projection_layers == 2:
elif config.projection_layers == 2:
self.projection = AudioProjection(
self.config.input_local_dim * self.config.group_size,
self.config.input_local_dim * self.config.group_size * 4,
self.config.out_hidden_size,
self.audio_input_local_dim * self.audio_group_size,
self.audio_input_local_dim * self.audio_group_size * 4,
self.audio_out_hidden_size,
)
else:
raise ValueError(
f"Invalid projection layers: {self.config.projection_layers}"
)
raise ValueError(f"Invalid projection layers: {config.projection_layers}")
model_path = self.server_args.model_path
model_path = get_global_server_args().model_path
if not os.path.isdir(model_path):
from huggingface_hub import snapshot_download
@@ -1216,13 +1233,16 @@ class MiMoAudioEncoder(nn.Module):
)
audio_tokenizer_path = os.path.join(model_path, "audio_tokenizer")
dev = torch.device(f"cuda:{torch.cuda.current_device()}")
self.audio_tokenizer = self._load_audio_tokenizer(audio_tokenizer_path, dev)
self.audio_tokenizer = self._load_mimo_audio_tokenizer(
audio_tokenizer_path, dev
)
@staticmethod
def _load_audio_tokenizer(path: str, device: torch.device) -> MiMoAudioTokenizer:
def _load_mimo_audio_tokenizer(
path: str, device: torch.device
) -> MiMoAudioTokenizer:
"""Load MiMoAudioTokenizer manually to avoid new-transformers compat issues."""
import json
import os
from safetensors.torch import load_file
@@ -1231,7 +1251,6 @@ class MiMoAudioEncoder(nn.Module):
config_dict = json.load(f)
config = MiMoAudioTokenizer.config_class(**config_dict)
model = MiMoAudioTokenizer(config)
# Load weights from safetensors or pytorch bin
safetensors_path = os.path.join(path, "model.safetensors")
bin_path = os.path.join(path, "pytorch_model.bin")
if os.path.exists(safetensors_path):
@@ -1243,43 +1262,34 @@ class MiMoAudioEncoder(nn.Module):
f"No model weights found in {path} "
"(expected model.safetensors or pytorch_model.bin)"
)
# strict=False: upstream ckpt also carries decoder/vocoder weights
# that this encoder-only MiMoAudioTokenizer doesn't materialize.
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 parsed_speech_empty_ids(self):
return self._parse_maybe_list(
self.config.speech_zeroemb_idx, self.config.audio_channels
)
def _parse_maybe_list(self, value: str | int, length: int) -> List[int]:
if isinstance(value, str) and "-" in value:
return [int(s) for s in value.split("-")]
return [int(value)] * length
# adapted from mimo-audio
def apply_input_local_transformer(self, speech_embeddings: torch.Tensor):
output = self.input_local_transformer(
def apply_input_local_transformer(
self, speech_embeddings: torch.Tensor
) -> torch.Tensor:
return self.input_local_transformer(
inputs_embeds=speech_embeddings,
return_dict=True,
is_causal=not self.config.input_full_attention, # for SDPA
)
return output.last_hidden_state # [T//group_size, group_size, input_local_dim]
is_causal=not self.audio_input_full_attention, # for SDPA
).last_hidden_state # [T//group_size, group_size, input_local_dim]
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),
embeds = torch.zeros(
(audio_codes.shape[0], self.audio_group_size, self.audio_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
for i in range(self.audio_channels):
embeds.add_(self.speech_embeddings[i](audio_codes[:, :, i]))
return embeds
def process_audio(self, audio):
def pad_audio_codes(self, audio: torch.Tensor) -> torch.Tensor:
T = audio.shape[0]
audio = audio[:, : self.audio_channels]
padded_T = (
@@ -1299,17 +1309,19 @@ class MiMoAudioEncoder(nn.Module):
+ audio[-1, :],
],
dim=0,
) # pad using the last embedding
padded_audio = padded_audio.reshape(
)
return padded_audio.reshape(
padded_T // self.audio_group_size,
self.audio_group_size,
self.audio_channels,
)
return padded_audio
def get_audio_feature(self, items) -> torch.Tensor:
# items: already audio-only MultimodalDataItem list from caller.
# Each item.feature is either one mel tensor or a list of mel tensors (e.g. long audio split into chunks).
"""Compute audio features for a list of audio MultimodalDataItem.
Each item.feature is either a mel tensor or a list of mel tensors
(long audio split into chunks).
"""
all_mels = []
for item in items:
f = item.feature
@@ -1321,9 +1333,8 @@ class MiMoAudioEncoder(nn.Module):
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
0, self.audio_out_hidden_size, device=device, dtype=dtype
)
# Batch tokenize: one encode_batch call for all mels
device = next(self.audio_tokenizer.encoder.parameters()).device
code_list = tokenize_audio_batch(
all_mels,
@@ -1331,20 +1342,16 @@ class MiMoAudioEncoder(nn.Module):
segment_size=self.audio_segment_size,
device=device,
)
codecs_to_concat = []
for codecs in code_list:
padded_codes = self.process_audio(
codecs
) # [T//group_size, group_size, audio_channels]
codecs_to_concat.append(padded_codes)
audio_codes = torch.cat(
codecs_to_concat, dim=0
) # [T//group_size, group_size, audio_channels]
audio_codes = torch.cat([self.pad_audio_codes(c) for c in code_list], dim=0)
embeds = self.apply_input_local_transformer(
self.apply_speech_embeddings(audio_codes)
)
return self.projection(embeds.reshape(embeds.shape[0], -1))
_audio_embeddings = self.apply_speech_embeddings(audio_codes)
audio_embeds = self.apply_input_local_transformer(
_audio_embeddings
) # [T//group_size, group_size, input_local_dim]
B = audio_embeds.shape[0]
audio_embeds = self.projection(audio_embeds.reshape(B, -1))
return audio_embeds
@classmethod
def remap_audio_weight_name(cls, name: str) -> str:
"""Normalize inner audio weight name variations to canonical form."""
for src, dst in cls.AUDIO_WEIGHT_REMAP:
if src in name:
return name.replace(src, dst)
return name
+24 -66
View File
@@ -78,7 +78,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader,
kv_cache_scales_loader,
)
from sglang.srt.models.mimo_audio import MiMoAudioEncoder, MiMoAudioEncoderConfig
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
@@ -993,7 +993,7 @@ class MiMoV2Model(nn.Module):
)
class MiMoV2ForCausalLM(nn.Module):
class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
# BitandBytes specific attributes
default_bitsandbytes_target_modules = [
".gate_proj.",
@@ -1015,8 +1015,10 @@ class MiMoV2ForCausalLM(nn.Module):
# Prefixes for weight routing in encoder_only/language_only modes
_LANGUAGE_WEIGHT_PREFIXES = ("model.", "lm_head.")
_VISION_AUDIO_WEIGHT_PREFIXES = ("visual.", "vision_model.", "audio_")
_VISION_AUDIO_WEIGHT_SUBSTRING = "speech_embeddings"
_VISION_WEIGHT_PREFIXES = ("visual.", "vision_model.")
# ``audio_`` already covers ``audio_encoder.`` so a single prefix is enough.
_AUDIO_WEIGHT_PREFIXES = ("audio_",)
_AUDIO_WEIGHT_SUBSTRING = "speech_embeddings"
def __init__(
self,
@@ -1070,8 +1072,7 @@ class MiMoV2ForCausalLM(nn.Module):
quant_config=None,
prefix=add_prefix("visual", prefix),
)
self.audio_config = MiMoAudioEncoderConfig(**audio_config)
self.audio_encoder = MiMoAudioEncoder(self.audio_config)
self.build_audio_encoder(MiMoAudioEncoderConfig(**audio_config))
self._routed_experts_weights_of_layer = LazyValue(
lambda: (
@@ -1126,9 +1127,6 @@ class MiMoV2ForCausalLM(nn.Module):
assert video_grid_thw.dim() == 2, video_grid_thw.dim()
return self.visual(pixel_values, grid_thw=video_grid_thw)
def get_audio_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
return self.audio_encoder.get_audio_feature(items)
@torch.inference_mode()
def encode_video_audio(self, mm_inputs: Dict) -> Optional[torch.Tensor]:
# EPD-side hook: encode audio tracks pulled from videos and trim to the
@@ -1276,14 +1274,14 @@ class MiMoV2ForCausalLM(nn.Module):
params_dict = dict(self.named_parameters())
skipped_mtp_weights = False
def _is_vision_audio_weight(name):
return (
name.startswith(self._VISION_AUDIO_WEIGHT_PREFIXES)
or self._VISION_AUDIO_WEIGHT_SUBSTRING in name
for name, loaded_weight in weights:
is_vision_weight = name.startswith(self._VISION_WEIGHT_PREFIXES)
is_audio_weight = (
name.startswith(self._AUDIO_WEIGHT_PREFIXES)
or self._AUDIO_WEIGHT_SUBSTRING in name
)
for name, loaded_weight in weights:
if not self._is_multimodal and _is_vision_audio_weight(name):
if not self._is_multimodal and (is_vision_weight or is_audio_weight):
continue
if self.config.encoder_only and name.startswith(
@@ -1291,61 +1289,21 @@ class MiMoV2ForCausalLM(nn.Module):
):
continue
if self._is_multimodal and "audio" in name:
if "projection" in name:
if (
"audio_encoder.audio_projection" in name
and "audio_encoder.projection" not in name
):
name = name.replace(
"audio_encoder.audio_projection", "audio_encoder.projection"
)
elif (
"audio_projection" in name
and "audio_encoder.projection" not in name
):
name = name.replace(
"audio_projection", "audio_encoder.projection"
)
param = params_dict[name]
weight_loader = getattr(
param, "weight_loader", default_weight_loader
if self._is_multimodal and is_audio_weight:
if name.startswith("audio_encoder."):
name = name[len("audio_encoder.") :]
name = self.remap_audio_weight_name(name)
if name not in params_dict:
logger.warning(
f"Audio param {name} not found in params_dict, skipping"
)
weight_loader(param, loaded_weight)
continue
if "input_local_transformer" in name:
if (
"audio_input_local_transformer" in name
and "audio_encoder.input_local_transformer" not in name
):
name = name.replace(
"audio_input_local_transformer",
"audio_encoder.input_local_transformer",
)
if name not in params_dict:
logger.warning(
f"Parameter {name} not found in params_dict, skipping"
)
continue
param = params_dict[name]
weight_loader = getattr(
param, "weight_loader", default_weight_loader
)
weight_loader(param, loaded_weight)
continue
if self._is_multimodal and "speech_embeddings" in name:
if (
"speech_embeddings" in name
and "audio_encoder.speech_embeddings" not in name
):
name = name.replace(
"speech_embeddings", "audio_encoder.speech_embeddings"
)
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight[: param.shape[0], :])
if self._AUDIO_WEIGHT_SUBSTRING in name:
weight_loader(param, loaded_weight[: param.shape[0], :])
else:
weight_loader(param, loaded_weight)
continue
if self._is_multimodal and "visual" in name:
+162
View File
@@ -0,0 +1,162 @@
"""MiMo-V2-ASR model.
Reuses the LM scaffold of ``MiMoForCausalLM`` and adds audio encoder
components via ``AudioEncoderMixin``. The encoder modules are attached as
top-level attributes (no ``audio_encoder.`` prefix) so the checkpoint
state_dict aligns 1:1 with ``self.named_parameters()``.
"""
import logging
from typing import Any, Iterable, List, Optional, Tuple
import torch
from sglang.srt.managers.mm_utils import (
MultiModalityDataPaddingPatternMultimodalTokens,
general_mm_embed_routine,
)
from sglang.srt.managers.schedule_batch import MultimodalInputs
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models import mimo_audio as _mimo_audio_module
from sglang.srt.models.mimo import MiMoForCausalLM
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
logger = logging.getLogger(__name__)
def _maybe_override_audio_attn_for_blackwell() -> None:
"""Swap mimo_audio.flash_attn_varlen_func to upstream FA2 on GPUs that
sgl-kernel's FA3 doesn't support.
sgl-kernel FA3 only covers sm80/86/89/90 — on Blackwell consumer cards
(sm_120 / RTX 50xx) its varlen kernel raises NotImplementedError. ASR is
small enough to be deployed on those GPUs, so when FA3 isn't supported
we replace the module-level reference with upstream flash-attn (FA2),
which works on sm_120. No-op on supported GPUs (FA3 stays).
MiMo-V2 (the heavy multimodal model) is only deployed on H100/A100, so
this override never triggers in its hot path.
"""
try:
from sgl_kernel.flash_attn import is_fa3_supported
except ImportError:
return
if is_fa3_supported():
return
try:
from flash_attn import flash_attn_varlen_func
except ImportError as e:
raise RuntimeError(
"MiMo-V2-ASR audio encoder needs upstream flash-attn on this GPU "
"(sgl-kernel FA3 doesn't support sm_120). Install with "
"`pip install flash-attn --no-build-isolation`."
) from e
_mimo_audio_module.flash_attn_varlen_func = flash_attn_varlen_func
MiMoV2ASRConfig = Any
# Top-level audio sub-module name prefixes (after AUDIO_WEIGHT_REMAP). Loaded
# directly by default_weight_loader because the LM branch's qkv/gate-up fused
# stacked-params mapping doesn't apply to the vanilla HF Qwen2Model used
# inside the audio encoder.
_AUDIO_NAME_PREFIXES: Tuple[str, ...] = (
"projection.",
"input_local_transformer.",
"speech_embeddings.",
)
# Training-only weights present in checkpoint but not used at inference.
# Checked AFTER the audio-prefix load path so substring matching here is
# safe: legitimate audio weights (``input_local_transformer.*``) are
# already consumed by ``_AUDIO_NAME_PREFIXES`` above.
_SKIP_NAME_SUBSTRINGS: Tuple[str, ...] = (
"hidden_states_downcast",
"local_transformer",
)
class MiMoV2ASRForCausalLM(MiMoForCausalLM, AudioEncoderMixin):
def __init__(
self,
config: MiMoV2ASRConfig,
quant_config=None,
prefix: str = "",
) -> None:
_maybe_override_audio_attn_for_blackwell()
super().__init__(config, quant_config=quant_config, prefix=prefix)
self.build_audio_encoder(MiMoAudioEncoderConfig(**config.audio_config))
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
return pattern.pad_input_tokens(input_ids, mm_inputs)
def get_input_embeddings(self):
if getattr(self.config, "encoder_only", False):
return None
return self.model.embed_tokens
@torch.no_grad()
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
input_embeds: torch.Tensor = None,
get_embedding: bool = False,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> torch.Tensor:
if getattr(self.config, "encoder_only", False):
raise NotImplementedError(
"forward() is not supported in encoder_only mode. "
"Use get_audio_feature() instead."
)
hidden_states = general_mm_embed_routine(
input_ids=input_ids,
forward_batch=forward_batch,
language_model=self.model,
multimodal_model=self,
positions=positions,
pp_proxy_tensors=pp_proxy_tensors,
)
if not get_embedding:
return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch
)
return self.pooler(hidden_states, forward_batch)
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
params_dict = dict(self.named_parameters())
deferred: List[Tuple[str, torch.Tensor]] = []
for name, loaded_weight in weights:
if name.startswith("audio_encoder."):
name = name[len("audio_encoder.") :]
name = self.remap_audio_weight_name(name)
if name.startswith(_AUDIO_NAME_PREFIXES):
if name not in params_dict:
logger.warning(
f"Audio param {name} not found in params_dict, skipping"
)
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
if name.startswith("speech_embeddings."):
weight_loader(param, loaded_weight[: param.shape[0], :])
else:
weight_loader(param, loaded_weight)
continue
if any(s in name for s in _SKIP_NAME_SUBSTRINGS):
continue
deferred.append((name, loaded_weight))
super().load_weights(iter(deferred))
EntryClass = MiMoV2ASRForCausalLM
@@ -0,0 +1,318 @@
"""Stateful audio preprocessing pipeline shared by MiMo multimodal and ASR processors."""
import io
import math
import os
import time
from collections import OrderedDict
from dataclasses import dataclass
from typing import Optional
import numpy as np
import pybase64
import requests
import torch
from sglang.utils import logger
try:
from torchcodec.decoders import AudioDecoder
except ImportError:
logger.warning(
"torchcodec is not installed; audio inputs will fail at request time"
)
AudioDecoder = None
try:
import torchaudio
from torchaudio.transforms import MelSpectrogram
except ImportError:
logger.warning(
"torchaudio is not installed; audio inputs will fail at request time"
)
torchaudio = None
MelSpectrogram = None
@dataclass
class AudioInput:
"""
if audio is str or bytes, only load it as mel spectrogram.
if audio is tuple, it is (waveform, original_sr)
if audio is torch.Tensor, it is tokenized input ids with shape (T, n_vq+).
if audio is np.ndarray, it is a pre-loaded waveform (1D, already resampled).
"""
audio: str | bytes | tuple | torch.Tensor | np.ndarray
def __post_init__(self):
if not isinstance(self.audio, (str, bytes, tuple, torch.Tensor, np.ndarray)):
raise ValueError(
f"audio must be a str, bytes, tuple, torch.Tensor, or np.ndarray, but got {type(self.audio)}"
)
if isinstance(self.audio, tuple):
if (
len(self.audio) != 2
or not isinstance(self.audio[0], torch.Tensor)
or not isinstance(self.audio[1], (int, float))
):
raise ValueError(
f"audio must be a tuple of (waveform-T, original_sr-int/float), but got {len(self.audio)} elements and {type(self.audio[0])} and {type(self.audio[1])}"
)
if self.audio[0].ndim != 1:
raise ValueError(
f"waveform must be a 1D tensor, but got {self.audio[0].ndim}D tensor"
)
if self.audio[1] <= 0:
raise ValueError(
f"original_sr must be a positive number, but got {self.audio[1]}"
)
if isinstance(self.audio, torch.Tensor) and self.audio.ndim != 2:
raise ValueError(
f"audio must be a 2D tensor, but got {self.audio.ndim}D tensor"
)
class MiMoAudioPipeline:
"""Stateful audio preprocessing pipeline.
Composable: held by both MiMoProcessor (multimodal) and MiMoV2ASRProcessor.
Owns the mel spectrogram, resampler cache, http session, and the special
token ids for ``<|sosp|> <|empty|>* <|eosp|>`` placeholders.
"""
def __init__(
self,
*,
audio_token_id: int,
audio_start_token_id: int,
audio_end_token_id: int,
audio_kernel_size: int = 3,
audio_stride_size: int = 2,
audio_avg_pooler: int = 2,
audio_group_size: int = 4,
audio_channels: int = 8,
audio_sampling_rate: int = 24000,
audio_nfft: int = 960,
audio_hop_length: int = 240,
audio_window_size: int = 960,
audio_fmin: int = 0,
audio_fmax: Optional[int] = None,
audio_n_mels: int = 128,
audio_input_id_per_second: int = 25,
max_resamplers: int = 16,
) -> None:
self.audio_token_id = audio_token_id
self.audio_start_token_id = audio_start_token_id
self.audio_end_token_id = audio_end_token_id
self.audio_kernel_size = audio_kernel_size
self.audio_stride_size = audio_stride_size
self.audio_avg_pooler = audio_avg_pooler
self.audio_group_size = audio_group_size
self.audio_channels = audio_channels
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_input_id_per_second = audio_input_id_per_second
self.mel_spectrogram_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 = None
self._resamplers: "OrderedDict[int, torchaudio.transforms.Resample]" = (
OrderedDict()
)
self._resamplers_max = max_resamplers
self.http_session = requests.Session()
@property
def audio_token_per_second(self) -> float:
return self.audio_input_id_per_second / self.audio_group_size
@staticmethod
def _ensure_audio_dependencies() -> None:
if torchaudio is None or MelSpectrogram is None:
raise RuntimeError(
"torchaudio is required for audio inputs; install torchaudio"
)
@property
def mel_spectrogram(self):
self._ensure_audio_dependencies()
if self._mel_spectrogram is None:
self._mel_spectrogram = MelSpectrogram(**self.mel_spectrogram_kwargs)
return self._mel_spectrogram
def compute_audio_token_len(self, mel_len: int) -> int:
n = mel_len + 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)
return math.ceil(n / self.audio_group_size)
def preprocess_audio(self, audio):
"""Load audio source → log-mel spectrogram + token length.
Input: filename string, bytes, or tuple of (waveform, original_sr).
Output: (mel-spectrogram tensor [T, n_mels], audio_token_len int).
"""
self._ensure_audio_dependencies()
assert isinstance(
audio, (str, bytes, tuple)
), f"audio must be a str, bytes or tuple, but got {type(audio)}"
if isinstance(audio, tuple):
waveform, original_sr = audio
else:
if isinstance(audio, bytes):
file = io.BytesIO(audio)
elif isinstance(audio, str):
if audio.startswith("data:"):
file = io.BytesIO(
pybase64.b64decode(audio.split(",")[1], validate=True)
)
elif audio.startswith("http://") or audio.startswith("https://"):
dl_start = time.perf_counter()
timeout = int(os.getenv("REQUEST_TIMEOUT", "5"))
try:
response = self.http_session.get(
audio, stream=True, timeout=timeout
)
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
if dl_elapsed_ms > 1000.0:
content_len = len(response.content)
logger.warning(
f"Slow audio download: {dl_elapsed_ms:.2f}ms, "
f"size={content_len / 1024:.1f}KB, url={audio}"
)
file = io.BytesIO(response.content)
response.close()
except Exception as e:
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
logger.error(
f"Failed to download audio: {dl_elapsed_ms:.2f}ms, "
f"error={type(e).__name__}: {e}, url={audio}"
)
raise
else:
file = audio
if AudioDecoder is None:
raise RuntimeError(
"torchcodec is required for audio decoding; install with `pip install torchcodec`."
)
try:
samples = AudioDecoder(file).get_all_samples()
except RuntimeError as e:
audio_source = (
audio
if isinstance(audio, str)
and (audio.startswith("http://") or audio.startswith("https://"))
else "<bytes or base64>"
)
logger.error(f"Failed to decode audio: {e}, source={audio_source}")
raise ValueError(
f"Invalid audio format: source={audio_source}, detail={e}"
) from e
waveform = samples.data
original_sr = samples.sample_rate
if original_sr != self.audio_sampling_rate:
if original_sr in self._resamplers:
self._resamplers.move_to_end(original_sr)
else:
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
)
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()
spec = spec.transpose(0, 1)
audio_token_len = self.compute_audio_token_len(spec.shape[0])
return spec, audio_token_len
def process_audio(self, audio_input: AudioInput):
"""Dispatch on the underlying audio payload.
- str/bytes/tuple/np.ndarray waveform → returns (mel-spec, token_len) tuple
- 2D tensor of pre-tokenized audio codes → returns padded codes tensor
shaped [T//group, group, channels]
"""
audio = audio_input.audio
if isinstance(audio, np.ndarray):
waveform = torch.from_numpy(audio).float()
audio = (waveform, self.audio_sampling_rate)
if isinstance(audio, (str, bytes, tuple)):
return self.preprocess_audio(audio)
assert (
audio.shape[1] >= self.audio_channels
), f"audio must have at least {self.audio_channels} channels, but got {audio.shape[1]}"
T = audio.shape[0]
audio = audio[:, : self.audio_channels].to(torch.long)
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.long)
+ 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 build_audio_placeholder_ids(self, audio_token_len: int) -> list[int]:
return (
[self.audio_start_token_id]
+ [self.audio_token_id] * audio_token_len
+ [self.audio_end_token_id]
)
def process_audio_input(self, audio_input: AudioInput) -> dict:
"""Run process_audio and produce the placeholder input_ids.
Replaces the duplicated _process_audio_content bodies in both processors.
Returns dict with input_ids, audio_input (mel or codes), and is_tokenized.
"""
processed = self.process_audio(audio_input)
if isinstance(processed, tuple):
is_tokenized = False
audio_spec, audio_token_len = processed
payload = audio_spec
else:
is_tokenized = True
audio_token_len = processed.shape[0]
payload = processed
return {
"input_ids": self.build_audio_placeholder_ids(audio_token_len),
"audio_input": payload,
"audio_token_len": audio_token_len,
"is_tokenized": is_tokenized,
}
@@ -3,21 +3,16 @@
import asyncio
import base64
import copy
import io
import json
import math
import os
import re
import subprocess
import time
from collections import OrderedDict
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field
from io import BytesIO
from typing import List, Literal, Optional, Union
import numpy as np
import pybase64
import requests
import torch
import torch.nn.functional as F
@@ -39,20 +34,14 @@ from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
MultimodalSpecialTokens,
)
from sglang.srt.multimodal.processors.mimo_audio import (
AudioInput,
MiMoAudioPipeline,
)
from sglang.srt.multimodal.processors.qwen_vl import smart_nframes
from sglang.srt.utils import ImageData, VideoData
from sglang.utils import logger
try:
import torchaudio
from torchaudio.transforms import MelSpectrogram
except ImportError:
logger.warning(
"torchaudio is not installed; audio inputs will fail at request time"
)
torchaudio = None
MelSpectrogram = None
@dataclass
class ImageInput:
@@ -112,45 +101,6 @@ class VideoInput:
)
@dataclass
class AudioInput:
"""
if audio is str or bytes, only load it as mel spectrogram.
if audio is tuple, it is (waveform, original_sr)
if audio is torch.Tensor, it is tokenized input ids with shape (T, n_vq+).
if audio is np.ndarray, it is a pre-loaded waveform (1D, already resampled).
"""
audio: str | bytes | tuple | torch.Tensor | np.ndarray
def __post_init__(self):
if not isinstance(self.audio, (str, bytes, tuple, torch.Tensor, np.ndarray)):
raise ValueError(
f"audio must be a str, bytes, tuple, torch.Tensor, or np.ndarray, but got {type(self.audio)}"
)
if isinstance(self.audio, tuple):
if (
len(self.audio) != 2
or not isinstance(self.audio[0], torch.Tensor)
or not isinstance(self.audio[1], (int, float))
):
raise ValueError(
f"audio must be a tuple of (waveform-T, original_sr-int/float), but got {len(self.audio)} elements and {type(self.audio[0])} and {type(self.audio[1])}"
)
if self.audio[0].ndim != 1:
raise ValueError(
f"waveform must be a 1D tensor, but got {self.audio[0].ndim}D tensor"
)
if self.audio[1] <= 0:
raise ValueError(
f"original_sr must be a positive number, but got {self.audio[1]}"
)
if isinstance(self.audio, torch.Tensor) and self.audio.ndim != 2:
raise ValueError(
f"audio must be a 2D tensor, but got {self.audio.ndim}D tensor"
)
@dataclass
class VideoAudioInput:
video: str | bytes | tuple[torch.Tensor, torch.Tensor]
@@ -334,11 +284,9 @@ class MiMoProcessor:
audio_fmin=0,
audio_fmax=None,
audio_n_mels=128,
audio_segment_size=6000,
audio_channels=8,
audio_group_size=4,
audio_input_id_per_second=25,
audio_zeroemb_idx=4096,
image_min_pixels=None,
image_max_pixels=None,
video_min_pixels=None,
@@ -395,11 +343,8 @@ class MiMoProcessor:
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
@@ -411,60 +356,24 @@ class MiMoProcessor:
self.video_tokens_per_second = video_tokens_per_second
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.mel_spectrogram_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.audio_pipeline = MiMoAudioPipeline(
audio_token_id=audio_token_id,
audio_start_token_id=audio_start_token_id,
audio_end_token_id=audio_end_token_id,
audio_kernel_size=audio_kernel_size,
audio_stride_size=audio_stride_size,
audio_avg_pooler=audio_avg_pooler,
audio_group_size=audio_group_size,
audio_channels=audio_channels,
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_input_id_per_second=audio_input_id_per_second,
)
self._mel_spectrogram = None
self._resamplers = OrderedDict()
self._resamplers_max = 16
self.audio_channels = audio_channels
self.audio_group_size = audio_group_size
self.audio_input_id_per_second = audio_input_id_per_second
if isinstance(audio_zeroemb_idx, int):
self.audio_zeroemb_idxs = torch.tensor(
[audio_zeroemb_idx] * self.audio_channels, dtype=torch.int32
)
elif isinstance(audio_zeroemb_idx, list):
if len(audio_zeroemb_idx) == 1:
self.audio_zeroemb_idxs = torch.tensor(
audio_zeroemb_idx * self.audio_channels, dtype=torch.int32
)
elif len(audio_zeroemb_idx) == self.audio_channels:
self.audio_zeroemb_idxs = torch.tensor(
audio_zeroemb_idx, dtype=torch.int32
)
else:
raise ValueError(
f"audio_zeroemb_idx must be a list of 1 or {self.audio_channels} integers, but got {len(audio_zeroemb_idx)}"
)
else:
raise ValueError(
f"audio_zeroemb_idx must be an integer or a list of {self.audio_channels} integers, but got {type(audio_zeroemb_idx)}"
)
assert image_min_pixels is not None
assert image_max_pixels is not None
assert video_min_pixels is not None
@@ -487,10 +396,18 @@ class MiMoProcessor:
"min_frames": min_frames,
}
self.http_session = requests.Session()
for k in kwargs:
logger.info(f"[Warning] Ignored unknown parameter {k} for MiMoProcessor")
def __getattr__(self, name):
# Delegate audio_pipeline fields so callers can use self.audio_token_id
# etc. directly. Only triggers when normal attribute lookup fails;
# __dict__.get avoids recursion before audio_pipeline is assigned.
pipeline = self.__dict__.get("audio_pipeline")
if pipeline is not None and hasattr(pipeline, name):
return getattr(pipeline, name)
raise AttributeError(name)
@classmethod
def from_hf_config(cls, hf_config, mm_config=None, **overrides):
# Params must come from hf_config.processor_config so E and D agree;
@@ -576,13 +493,13 @@ class MiMoProcessor:
return _ffprobe_has_audio(path_or_data, stdin=None, label=path_or_data)
if isinstance(path_or_data, bytes):
source = io.BytesIO(path_or_data)
source = BytesIO(path_or_data)
elif (
isinstance(path_or_data, str)
and path_or_data.startswith("data:")
and ";base64," in path_or_data
):
source = io.BytesIO(base64.b64decode(path_or_data.split(";base64,")[1]))
source = BytesIO(base64.b64decode(path_or_data.split(";base64,")[1]))
else:
source = path_or_data # local path or file://
try:
@@ -667,7 +584,9 @@ class MiMoProcessor:
all_timestamps.extend(aligned_ts[::step].tolist())
if self.has_audio_track(video_blob):
audio_spec, audio_token_len = self.preprocess_audio(video_blob)
audio_spec, audio_token_len = self.audio_pipeline.preprocess_audio(
video_blob
)
units = self._build_video_audio_units(
grid,
aligned_ts,
@@ -704,7 +623,7 @@ class MiMoProcessor:
for audio in mm_data:
if isinstance(audio, np.ndarray):
audio = (torch.from_numpy(audio).float(), self.audio_sampling_rate)
spec, token_len = self.preprocess_audio(audio)
spec, token_len = self.audio_pipeline.preprocess_audio(audio)
all_specs.append(spec)
all_lens.append(token_len)
return {
@@ -714,20 +633,6 @@ class MiMoProcessor:
raise ValueError(f"Unsupported modality for EPD preprocessing: {modality}")
@property
def mel_spectrogram(self):
self._ensure_audio_dependencies()
if self._mel_spectrogram is None:
self._mel_spectrogram = MelSpectrogram(**self.mel_spectrogram_kwargs)
return self._mel_spectrogram
@staticmethod
def _ensure_audio_dependencies():
if torchaudio is None or MelSpectrogram is None:
raise RuntimeError(
"torchaudio is required for audio inputs; install torchaudio"
)
def prepare_image_kwargs(self, image: ImageInput):
kwargs = {}
for k in ["min_pixels", "max_pixels"]:
@@ -764,95 +669,6 @@ class MiMoProcessor:
raise ValueError("Video sampling strategy not specified")
return kwargs
def preprocess_audio(self, audio: str | bytes):
self._ensure_audio_dependencies()
"""
- Input: audio filename string, bytes, or tuple of (waveform, original_sr)
- Output:
- mel spectrogram: torch.Tensor (T, n_mels)
- number of tokens: int
"""
assert isinstance(
audio, (str, bytes, tuple)
), f"audio must be a str, bytes or tuple, but got {type(audio)}"
if isinstance(audio, tuple):
waveform, original_sr = audio
else:
if isinstance(audio, bytes):
file = io.BytesIO(audio)
elif isinstance(audio, str):
if audio.startswith("data:"):
file = io.BytesIO(
pybase64.b64decode(audio.split(",")[1], validate=True)
)
elif audio.startswith("http://") or audio.startswith("https://"):
dl_start = time.perf_counter()
timeout = int(os.getenv("REQUEST_TIMEOUT", "5"))
try:
response = self.http_session.get(
audio, stream=True, timeout=timeout
)
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
if dl_elapsed_ms > 1000.0:
content_len = len(response.content)
logger.warning(
f"Slow audio download: {dl_elapsed_ms:.2f}ms, "
f"size={content_len / 1024:.1f}KB, url={audio}"
)
file = io.BytesIO(response.content)
response.close()
except Exception as e:
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
logger.error(
f"Failed to download audio: {dl_elapsed_ms:.2f}ms, "
f"error={type(e).__name__}: {e}, url={audio}"
)
raise
else:
file = audio
try:
samples = AudioDecoder(file).get_all_samples()
except RuntimeError as e:
audio_source = (
audio
if isinstance(audio, str)
and (audio.startswith("http://") or audio.startswith("https://"))
else "<bytes or base64>"
)
logger.error(f"Failed to decode audio: {e}, source={audio_source}")
raise ValueError(
f"Invalid audio format: source={audio_source}, detail={e}"
) from e
waveform = samples.data
original_sr = samples.sample_rate
if original_sr != self.audio_sampling_rate:
if original_sr in self._resamplers:
self._resamplers.move_to_end(original_sr)
else:
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
)
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()
spec = spec.transpose(0, 1)
audio_token_len = spec.shape[0] + 3 - self.audio_kernel_size
audio_token_len = (
audio_token_len + 2 - self.audio_kernel_size
) // self.audio_stride_size + 1
audio_token_len = audio_token_len // self.audio_avg_pooler + int(
audio_token_len % self.audio_avg_pooler != 0
)
audio_token_len = math.ceil(audio_token_len / self.audio_group_size)
return spec, audio_token_len
def process_image(self, image: ImageInput):
kwargs = self.prepare_image_kwargs(image)
image = image.image
@@ -1014,40 +830,6 @@ class MiMoProcessor:
)
return visual_patches, thw_grid, aligned_timestamps, video_meta
def process_audio(self, audio: AudioInput):
audio = audio.audio
if isinstance(audio, np.ndarray):
waveform = torch.from_numpy(audio).float()
audio = (waveform, self.audio_sampling_rate)
if isinstance(audio, (str, bytes, tuple)):
audio_spec, audio_token_len = self.preprocess_audio(audio)
return audio_spec, audio_token_len
assert (
audio.shape[1] >= self.audio_channels
), f"audio must have at least {self.audio_channels} channels, but got {audio.shape[1]}"
T = audio.shape[0]
audio = audio[:, : self.audio_channels].to(torch.long)
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.long)
+ 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 _process_videos_parallel(self, contents):
video_contents_info = []
for idx, content in enumerate(contents):
@@ -1171,29 +953,17 @@ class MiMoProcessor:
}
def _process_audio_content(self, content, verbose):
processed_audio = self.process_audio(content.content)
if isinstance(processed_audio, tuple):
is_tokenized = False
audio_spec, audio_token_len = processed_audio
audio_input = audio_spec
else:
is_tokenized = True
audio_token_len = processed_audio.shape[0]
audio_input = processed_audio
_input_ids = (
[self.audio_start_token_id]
+ [self.audio_token_id] * audio_token_len
+ [self.audio_end_token_id]
)
result = self.audio_pipeline.process_audio_input(content.content)
verbose_str = ""
if verbose:
verbose_str = f"Audio (is_tokenized={is_tokenized}): [<audio_start> {audio_token_len}*<audio> <audio_end>]\n"
verbose_str = (
f"Audio (is_tokenized={result['is_tokenized']}): "
f"[<audio_start> {result['audio_token_len']}*<audio> <audio_end>]\n"
)
return {
"input_ids": _input_ids,
"audio_input": audio_input,
"is_tokenized": is_tokenized,
"input_ids": result["input_ids"],
"audio_input": result["audio_input"],
"is_tokenized": result["is_tokenized"],
"verbose": verbose_str,
}
@@ -1219,7 +989,7 @@ class MiMoProcessor:
grid_t_timestamps = timestamps[
:: self.temporal_patch_size * self.temporal_compression_ratio
]
audio_token_per_second = self.audio_input_id_per_second / self.audio_group_size
audio_token_per_second = self.audio_token_per_second
units = []
for i in range(len(grid_t_timestamps)):
@@ -1363,7 +1133,7 @@ class MiMoProcessor:
self, content_idx, content, video_results, verbose
):
visual_patches, thw_grid, timestamps, video_meta = video_results[content_idx]
processed_audio = self.process_audio(content.content)
processed_audio = self.audio_pipeline.process_audio(content.content)
if isinstance(processed_audio, tuple):
assert (
@@ -2208,6 +1978,11 @@ class MiMoV2Processor(BaseMultimodalProcessor):
audio_source = raw_audio_item.get("url", loaded_audio)
elif isinstance(raw_audio_item, (str, bytes, torch.Tensor)):
audio_source = raw_audio_item
else:
raise ValueError(
f"unsupported audio item: loaded={type(loaded_audio).__name__}, "
f"raw={type(raw_audio_item).__name__}"
)
contents.append(
Content(
@@ -0,0 +1,286 @@
"""MiMo-V2-ASR multimodal processor.
Audio preprocessing is delegated to :class:`MiMoAudioPipeline`; this
processor only handles the special-token contract and content interleaving.
"""
import asyncio
import re
from dataclasses import dataclass
from typing import List, Literal, Union
import numpy as np
import torch
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalProcessorOutput,
)
from sglang.srt.models.mimo_v2_asr import MiMoV2ASRForCausalLM
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
MultimodalSpecialTokens,
)
from sglang.srt.multimodal.processors.mimo_audio import (
AudioInput,
MiMoAudioPipeline,
)
from sglang.utils import logger
TextInput = str | list[int]
@dataclass
class _Content:
type: Literal["text", "audio"]
content: TextInput | AudioInput
class MiMoV2ASRProcessor(BaseMultimodalProcessor):
"""ASR-only MiMo processor.
Wires three special tokens into the input id stream around each audio
span: ``<|sosp|> <|empty|>* <|eosp|>``. The actual mel/codec preparation
is owned by :class:`MiMoAudioPipeline`, which is shared with the
multimodal MiMo-V2 processor.
"""
models = [MiMoV2ASRForCausalLM]
AUDIO_PAD_TOKEN = "<|empty|>"
AUDIO_START_TOKEN = "<|sosp|>"
AUDIO_END_TOKEN = "<|eosp|>"
AUDIO_REGEX = re.compile(r"<\|sosp\|>(?:<\|empty\|>)+<\|eosp\|>")
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
super().__init__(hf_config, server_args, _processor, *args, **kwargs)
self.tokenizer = _processor
self.audio_pipeline = MiMoAudioPipeline(
audio_token_id=self._resolve_special_token_id(self.AUDIO_PAD_TOKEN),
audio_start_token_id=self._resolve_special_token_id(self.AUDIO_START_TOKEN),
audio_end_token_id=self._resolve_special_token_id(self.AUDIO_END_TOKEN),
audio_sampling_rate=24000,
)
self.mm_tokens = MultimodalSpecialTokens(
audio_token=f"{self.AUDIO_START_TOKEN}{self.AUDIO_PAD_TOKEN}{self.AUDIO_END_TOKEN}",
audio_token_id=self.audio_token_id,
audio_token_regex=self.AUDIO_REGEX,
).build(_processor)
def __getattr__(self, name):
# Delegate audio_pipeline fields so callers can use self.audio_token_id
# etc. directly. Only triggers when normal attribute lookup fails;
# __dict__.get avoids recursion before audio_pipeline is assigned.
pipeline = self.__dict__.get("audio_pipeline")
if pipeline is not None and hasattr(pipeline, name):
return getattr(pipeline, name)
raise AttributeError(name)
def _resolve_special_token_id(self, name: str) -> int:
tid = self.tokenizer.convert_tokens_to_ids(name)
if tid is None or tid == self.tokenizer.unk_token_id:
raise ValueError(
f"tokenizer missing required special token {name!r}; "
"checkpoint vocab does not match MiMo-V2-ASR"
)
return int(tid)
def _process_contents(self, contents: List[_Content]):
"""Run pipeline + tokenizer over an interleaved content list.
Returns ``(input_ids: Tensor[L], audio_inputs: list[Tensor],
position_ids: Tensor[3,L], rope_deltas: Tensor[1,1])``.
"""
input_ids: List[int] = []
audio_inputs: List[torch.Tensor] = []
for content in contents:
if content.type == "text":
if isinstance(content.content, str):
input_ids.extend(self.tokenizer.encode(content.content))
else:
input_ids.extend(content.content)
elif content.type == "audio":
result = self.audio_pipeline.process_audio_input(content.content)
audio_inputs.append(result["audio_input"])
input_ids.extend(result["input_ids"])
ids = torch.as_tensor(input_ids)
position_ids = torch.arange(ids.shape[0]).expand(3, -1)
rope_deltas = torch.zeros((1, 1), dtype=torch.int32)
return ids, audio_inputs, position_ids, rope_deltas
def process_mm_data(
self, input_text, images=None, videos=None, audios=None, **kwargs
) -> dict:
if audios and not self.AUDIO_REGEX.search(input_text or ""):
input_text = f"{self.mm_tokens.audio_token}{input_text or ''}"
processed_audios: List[Union[tuple, torch.Tensor]] = []
if audios:
for audio in audios:
if isinstance(audio, np.ndarray):
audio_tensor = torch.from_numpy(audio).float()
elif isinstance(audio, torch.Tensor):
audio_tensor = audio.float()
else:
processed_audios.append(audio)
continue
if audio_tensor.ndim == 1:
processed_audios.append(
(audio_tensor.cpu().contiguous(), self.audio_sampling_rate)
)
else:
processed_audios.append(audio_tensor.cpu().contiguous())
contents: List[_Content] = []
if input_text and processed_audios:
multimodal_tokens_pattern = self.mm_tokens.get_combined_regex()
text_parts = re.split(multimodal_tokens_pattern, input_text)
audio_iter = iter(processed_audios)
for text_part in text_parts:
if multimodal_tokens_pattern.match(text_part):
modality = self.mm_tokens.get_modality_of_token(text_part)
if modality == Modality.AUDIO:
try:
audio = next(audio_iter)
contents.append(
_Content(type="audio", content=AudioInput(audio=audio))
)
except StopIteration:
pass
else:
if text_part:
contents.append(_Content(type="text", content=text_part))
else:
contents.extend(
_Content(type="audio", content=AudioInput(audio=audio))
for audio in processed_audios
)
if not contents:
ids = self.tokenizer(
input_text or "",
return_tensors="pt",
add_special_tokens=True,
).input_ids
return {"input_ids": ids}
input_ids, audio_inputs, position_ids, rope_deltas = self._process_contents(
contents
)
ret: dict = {
"input_ids": input_ids,
"mrope_positions": position_ids,
"mrope_position_delta": rope_deltas,
}
if audio_inputs:
ret["audio_features"] = audio_inputs
return ret
async def process_mm_data_async(
self,
image_data: List[Union[str, bytes]],
audio_data: List[Union[str, bytes]],
input_text,
request_obj,
*args,
**kwargs,
):
if audio_data is None:
audio_data = getattr(request_obj, "audio_data", [])
if not audio_data:
return None
if not self.AUDIO_REGEX.search(input_text):
input_text = f"{self.mm_tokens.audio_token}{input_text}"
base_output = await self.load_mm_data(
prompt=input_text,
image_data=[],
video_data=[],
audio_data=audio_data,
multimodal_tokens=self.mm_tokens,
audio_sample_rate=self.audio_sampling_rate,
)
multimodal_tokens_pattern = self.mm_tokens.get_combined_regex()
raw_audio_data = audio_data or []
loaded_audio_iter = iter(base_output.audios)
raw_audio_iter = iter(raw_audio_data)
text_parts = re.split(multimodal_tokens_pattern, base_output.input_text)
contents: List[_Content] = []
for text_part in text_parts:
if multimodal_tokens_pattern.match(text_part):
modality = self.mm_tokens.get_modality_of_token(text_part)
assert modality is not None
if modality == Modality.AUDIO:
loaded_audio = next(loaded_audio_iter)
raw_audio_item = next(raw_audio_iter)
if isinstance(loaded_audio, np.ndarray):
audio_source = loaded_audio
elif isinstance(raw_audio_item, dict):
audio_source = raw_audio_item.get("url", loaded_audio)
elif isinstance(raw_audio_item, (str, bytes, torch.Tensor)):
audio_source = raw_audio_item
else:
raise ValueError(
f"unsupported audio item: loaded={type(loaded_audio).__name__}, "
f"raw={type(raw_audio_item).__name__}"
)
contents.append(
_Content(
type="audio",
content=AudioInput(audio=audio_source),
)
)
else:
if text_part:
contents.append(_Content(type="text", content=text_part))
loop = asyncio.get_running_loop()
try:
input_ids, audio_inputs, position_ids, rope_deltas = (
await loop.run_in_executor(
self.io_executor,
lambda: self._process_contents(contents),
)
)
except RuntimeError as e:
logger.error(f"MiMo ASR processor failed in process_mm_data_async: {e}")
raise ValueError(f"Multimodal data is corrupted or cannot be decoded: {e}")
input_ids_flat = input_ids.flatten()
if audio_inputs:
mm_items = [
MultimodalDataItem(
modality=Modality.AUDIO,
feature=audio_inputs,
offsets=self.get_mm_items_offset(
input_ids=input_ids_flat,
mm_token_id=self.audio_token_id,
),
)
]
else:
mm_items = []
return MultimodalProcessorOutput(
mm_items=mm_items,
input_ids=input_ids_flat.tolist(),
audio_token_id=self.audio_token_id,
audio_start_id=self.audio_start_token_id,
audio_end_id=self.audio_end_token_id,
mrope_positions=position_ids,
mrope_position_delta=rope_deltas,
)