model: support qwen3-asr (#22073)

Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
Aditya Sharma
2026-04-07 13:27:05 +08:00
committed by GitHub
co-authored by Xinyuan Tong
parent a757c1e3fb
commit f6e85676b5
10 changed files with 571 additions and 11 deletions
+2
View File
@@ -6,6 +6,8 @@ This benchmark evaluates the performance and accuracy (Word Error Rate - WER) of
- `openai/whisper-large-v3`
- `openai/whisper-large-v3-turbo`
- `Qwen/Qwen3-ASR-1.7B`
- `Qwen/Qwen3-ASR-0.6B`
## Setup
@@ -51,9 +51,38 @@ in the GitHub search bar.
| **Ernie4.5-VL** | `baidu/ERNIE-4.5-VL-28B-A3B-PT` | Baidu's vision-language models(28B,424B). Support image and video comprehension, and also support thinking. | |
| **JetVLM** | | JetVLM is an vision-language model designed for high-performance multimodal understanding and generation tasks built upon Jet-Nemotron. | Coming soon |
| **Step3-VL** (10B) | `stepfun-ai/Step3-VL-10B` | StepFun's lightweight open-source 10B parameter VLM for multimodal intelligence, excelling in visual perception, complex reasoning, and human alignment. | |
| **Qwen3-ASR** (0.6B, 1.7B) | `Qwen/Qwen3-ASR-1.7B` | Alibaba's automatic speech recognition models supporting 52 languages. Served via the `/v1/audio/transcriptions` endpoint. | |
| **Qwen3-Omni** | `Qwen/Qwen3-Omni-30B-A3B-Instruct` | Alibaba's omni-modal MoE model. Currently supports the **Thinker** component (multimodal understanding for text, images, audio, and video), while the **Talker** component (audio generation) is not yet supported. | |
| **LFM2-VL** | `LiquidAI/LFM2.5-VL-1.6B` | Liquid AI's vision-language model combining a SigLip2 vision encoder (NaFlex variable-resolution) with the LFM2 hybrid attention + short convolution language model. Supports multi-image inputs. | |
## Audio Transcription
SGLang supports audio-only ASR models via the OpenAI-compatible `/v1/audio/transcriptions` endpoint. Upload an audio file and receive a transcription.
### Launch Command
```shell
sglang serve \
--model-path Qwen/Qwen3-ASR-1.7B \
--served-model-name qwen3-asr \
--trust-remote-code \
--host 0.0.0.0 --port 30000
```
### Example Request
```bash
curl http://localhost:30000/v1/audio/transcriptions \
-F file=@audio.wav \
-F model=qwen3-asr \
-F response_format=verbose_json
```
| Model Family | Example Identifier | Notes |
|--------------|--------------------|-------|
| **Whisper** | `openai/whisper-large-v3` | OpenAI's speech recognition model. |
| **Qwen3-ASR** (0.6B, 1.7B) | `Qwen/Qwen3-ASR-1.7B` | Use `--trust-remote-code`. Supports 52 languages. |
## Video Input Support
SGLang supports video input for Vision-Language Models (VLMs), enabling temporal reasoning tasks such as video question answering, captioning, and holistic scene understanding. Video clips are decoded, key frames are sampled, and the resulting tensors are batched together with the text prompt, allowing multimodal inference to integrate visual and linguistic context.
+2
View File
@@ -23,6 +23,7 @@ from sglang.srt.configs.nano_nemotron_vl import NemotronH_Nano_VL_V2_Config
from sglang.srt.configs.nemotron_h import NemotronHConfig
from sglang.srt.configs.olmo3 import Olmo3Config
from sglang.srt.configs.qwen3_5 import Qwen3_5Config, Qwen3_5MoeConfig
from sglang.srt.configs.qwen3_asr import Qwen3ASRConfig
from sglang.srt.configs.qwen3_next import Qwen3NextConfig
from sglang.srt.configs.step3_vl import (
Step3TextConfig,
@@ -63,4 +64,5 @@ __all__ = [
"JetNemotronConfig",
"JetVLMConfig",
"Step3p5Config",
"Qwen3ASRConfig",
]
+12 -2
View File
@@ -196,8 +196,16 @@ class ModelConfig:
self.is_image_understandable_model = enable_multimodal and hasattr(
self.hf_config, "vision_config"
)
self.is_audio_understandable_model = enable_multimodal and hasattr(
self.hf_config, "audio_config"
# Models expose audio_config at different nesting levels:
# - top-level audio_config: e.g. Qwen2Audio
# - thinker_config.audio_config: Qwen3-Omni, Qwen3-ASR (nested thinker arch)
# - is_audio_model(): Whisper, Qwen3-ASR (architecture-based fallback)\
# TODO: Handle this more robustly by standardizing the config structure in the future
self.is_audio_understandable_model = enable_multimodal and (
hasattr(self.hf_config, "audio_config")
or hasattr(getattr(self.hf_config, "thinker_config", None), "audio_config")
or is_audio_model(self.hf_config.architectures)
)
self.is_multimodal_chunked_prefill_supported = (
@@ -1332,6 +1340,7 @@ multimodal_model_archs = [
"Qwen3VLMoeForConditionalGeneration",
"Qwen3_5ForConditionalGeneration",
"Qwen3_5MoeForConditionalGeneration",
"Qwen3ASRForConditionalGeneration",
"Qwen3OmniMoeForConditionalGeneration",
"KimiVLForConditionalGeneration",
"InternVLChatModel",
@@ -1380,6 +1389,7 @@ def is_multimodal_model(model_architectures: List[str]):
def is_audio_model(model_architectures: List[str]):
models = [
"WhisperForConditionalGeneration",
"Qwen3ASRForConditionalGeneration",
]
return any(model in model_architectures for model in models)
+172
View File
@@ -0,0 +1,172 @@
import torch
from transformers import (
AutoConfig,
AutoFeatureExtractor,
AutoTokenizer,
PretrainedConfig,
ProcessorMixin,
)
from sglang.srt.configs.qwen3_omni import Qwen3OmniMoeAudioEncoderConfig
from sglang.srt.multimodal.customized_mm_processor_utils import (
register_customized_processor,
)
from sglang.utils import logger
class Qwen3ASRThinkerConfig(PretrainedConfig):
model_type = "qwen3_asr_thinker"
sub_configs = {
"audio_config": Qwen3OmniMoeAudioEncoderConfig,
}
def __init__(
self,
audio_config=None,
text_config=None,
audio_token_id=151676,
audio_start_token_id=151669,
audio_end_token_id=151670,
**kwargs,
):
super().__init__(**kwargs)
if isinstance(audio_config, dict):
audio_config = Qwen3OmniMoeAudioEncoderConfig(**audio_config)
elif audio_config is None:
audio_config = Qwen3OmniMoeAudioEncoderConfig()
self.audio_config = audio_config
if isinstance(text_config, dict):
from transformers.models.qwen3.configuration_qwen3 import (
Qwen3Config as HFQwen3Config,
)
text_config = HFQwen3Config(**text_config)
elif text_config is None:
raise ValueError(
"Qwen3ASRThinkerConfig requires a text_config dict with "
"model parameters (hidden_size, num_attention_heads, etc.). "
"Got None."
)
self.text_config = text_config
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
class Qwen3ASRConfig(PretrainedConfig):
model_type = "qwen3_asr"
sub_configs = {
"thinker_config": Qwen3ASRThinkerConfig,
}
def __init__(self, thinker_config=None, **kwargs):
super().__init__(**kwargs)
if thinker_config is None:
thinker_config = {}
logger.info(
"thinker_config is None. "
"Initializing Qwen3-ASR thinker with default values"
)
if isinstance(thinker_config, dict):
self.thinker_config = Qwen3ASRThinkerConfig(**thinker_config)
else:
self.thinker_config = thinker_config
def get_text_config(self, decoder=False) -> PretrainedConfig:
return self.thinker_config.text_config
class Qwen3ASRProcessor(ProcessorMixin):
"""Minimal composite processor: WhisperFeatureExtractor + Qwen2Tokenizer.
AutoProcessor.from_pretrained() for Qwen3-ASR returns just a tokenizer,
but SGLang's multimodal pipeline needs a processor that handles audio.
"""
attributes = ["feature_extractor", "tokenizer"]
feature_extractor_class = "WhisperFeatureExtractor"
tokenizer_class = "AutoTokenizer"
def __init__(self, feature_extractor=None, tokenizer=None, **kwargs):
super().__init__(feature_extractor=feature_extractor, tokenizer=tokenizer)
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
trust_remote_code = kwargs.pop("trust_remote_code", True)
feature_extractor = AutoFeatureExtractor.from_pretrained(
pretrained_model_name_or_path,
trust_remote_code=trust_remote_code,
)
tokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path,
trust_remote_code=trust_remote_code,
)
return cls(feature_extractor=feature_extractor, tokenizer=tokenizer)
def _get_feat_extract_output_lengths(self, input_lengths):
if not isinstance(input_lengths, torch.Tensor):
input_lengths = torch.tensor(input_lengths)
input_lengths_leave = input_lengths % 100
feat_lengths = (input_lengths_leave - 1) // 2 + 1
return ((feat_lengths - 1) // 2 + 1 - 1) // 2 + 1 + (input_lengths // 100) * 13
def __call__(self, text=None, audio=None, audio_kwargs=None, **kwargs):
inputs = {}
if audio is not None:
audio_kwargs = audio_kwargs or {}
audio_inputs = self.feature_extractor(
audio,
sampling_rate=self.feature_extractor.sampling_rate,
return_attention_mask=True,
return_tensors=kwargs.get("return_tensors"),
**audio_kwargs,
)
inputs["input_features"] = audio_inputs["input_features"]
if "attention_mask" in audio_inputs:
inputs["feature_attention_mask"] = audio_inputs["attention_mask"]
if text is not None:
text_inputs = self.tokenizer(
text,
return_tensors=kwargs.get("return_tensors"),
padding=kwargs.get("padding", False),
)
input_ids = text_inputs["input_ids"]
# Expand the single <|audio_pad|> placeholder in the prompt to N
# copies, where N is the audio encoder's output length for this clip.
# Without this, the model only sees 1 audio token for hundreds of
# feature frames and can't align audio embeddings with token positions.
if audio is not None and "feature_attention_mask" in inputs:
audio_pad_id = self.tokenizer.convert_tokens_to_ids("<|audio_pad|>")
feat_lengths = inputs["feature_attention_mask"].sum(dim=-1)
audio_token_counts = self._get_feat_extract_output_lengths(feat_lengths)
expanded = []
for seq_idx in range(input_ids.shape[0]):
ids = input_ids[seq_idx].tolist()
audio_idx = 0
new_ids = []
for tid in ids:
if tid == audio_pad_id and audio_idx < len(audio_token_counts):
n = int(audio_token_counts[audio_idx].item())
new_ids.extend([audio_pad_id] * n)
audio_idx += 1
else:
new_ids.append(tid)
expanded.append(new_ids)
max_len = max(len(s) for s in expanded)
pad_id = self.tokenizer.pad_token_id or 0
padded = [s + [pad_id] * (max_len - len(s)) for s in expanded]
input_ids = torch.tensor(padded, dtype=torch.long)
inputs["input_ids"] = input_ids
return inputs
AutoConfig.register("qwen3_asr", Qwen3ASRConfig)
AutoConfig.register("qwen3_asr_thinker", Qwen3ASRThinkerConfig)
register_customized_processor(Qwen3ASRProcessor)(Qwen3ASRConfig)
@@ -476,8 +476,8 @@ class MMEncoder:
if self.model_type in ["qwen2_audio", "qwen2_5_omni"]:
input_length = (feature_lens - 1) // 2 + 1
return (input_length - 2) // 2 + 1
# qwen3_omni_moe
elif self.model_type == "qwen3_omni_moe":
# qwen3_asr / qwen3_omni_moe (same audio encoder architecture)
elif self.model_type in ["qwen3_asr", "qwen3_omni_moe"]:
input_lengths_leave = feature_lens % 100
feat_lengths = (input_lengths_leave - 1) // 2 + 1
output_lengths = (
@@ -50,12 +50,22 @@ logger = logging.getLogger(__name__)
TIMESTAMP_BASE_TOKEN_ID = 50365 # <|0.00|>
TIMESTAMP_BASE_OFFSET = 0.02 # Each token step = 0.02 seconds
_QWEN3_ASR_TEXT_TAG = "<asr_text>"
def _detect_model_family(model_config) -> str:
archs = getattr(getattr(model_config, "hf_config", None), "architectures", []) or []
if "Qwen3ASRForConditionalGeneration" in archs:
return "qwen3_asr"
return "whisper"
class OpenAIServingTranscription(OpenAIServingBase):
"""Handler for /v1/audio/transcriptions requests"""
def __init__(self, tokenizer_manager: TokenizerManager):
super().__init__(tokenizer_manager)
self._model_family = _detect_model_family(tokenizer_manager.model_config)
def _request_id_prefix(self) -> str:
return "trsc-"
@@ -71,6 +81,27 @@ class OpenAIServingTranscription(OpenAIServingBase):
raw_request: Request = None,
) -> tuple[GenerateReqInput, TranscriptionRequest]:
"""Convert transcription request to internal format."""
if self._model_family == "qwen3_asr":
prompt = (
"<|im_start|>user\n"
"<|audio_start|><|audio_pad|><|audio_end|>"
"<|im_end|>\n"
"<|im_start|>assistant\n"
)
sampling_params = {
"temperature": request.temperature,
"max_new_tokens": 1024,
}
adapted_request = GenerateReqInput(
text=prompt,
audio_data=request.audio_data,
sampling_params=sampling_params,
stream=request.stream,
modalities=["audio"],
routing_key=self.extract_routing_key(raw_request),
)
return adapted_request, request
# Build sampling params - include language for WhisperProcessor
sampling_params = {
"temperature": request.temperature,
@@ -232,6 +263,8 @@ class OpenAIServingTranscription(OpenAIServingBase):
return self.create_error_response(str(e))
text = ret.get("text", "")
if self._model_family == "qwen3_asr":
text = _postprocess_qwen3_asr(text)
usage = TranscriptionUsage(seconds=int(math.ceil(request.audio_duration_s)))
# Build response based on format
@@ -239,15 +272,22 @@ class OpenAIServingTranscription(OpenAIServingBase):
return Response(content=text, media_type="text/plain")
if request.response_format == "verbose_json":
output_ids = ret.get("output_ids", [])
tokenizer = self.tokenizer_manager.tokenizer
parsed_text, segments = self._parse_segments(output_ids, tokenizer)
if self._model_family == "whisper":
output_ids = ret.get("output_ids", [])
tokenizer = self.tokenizer_manager.tokenizer
parsed_text, segments = self._parse_segments(output_ids, tokenizer)
return TranscriptionVerboseResponse(
language=request.language or "en",
duration=round(request.audio_duration_s, 2),
text=parsed_text or text,
segments=segments,
usage=usage,
)
return TranscriptionVerboseResponse(
language=request.language or "en",
language=request.language,
duration=round(request.audio_duration_s, 2),
text=parsed_text or text,
segments=segments,
text=text,
segments=[],
usage=usage,
)
@@ -324,3 +364,13 @@ class OpenAIServingTranscription(OpenAIServingBase):
yield f"data: {error}\n\n"
yield "data: [DONE]\n\n"
# TODO (adityavaid): refactor model-specific postprocessing into a plugin/adapter mechanism.
def _postprocess_qwen3_asr(text: str) -> str:
if not text:
return ""
if _QWEN3_ASR_TEXT_TAG in text:
_, text_part = text.rsplit(_QWEN3_ASR_TEXT_TAG, 1)
return text_part.strip()
return text.strip()
+199
View File
@@ -0,0 +1,199 @@
"""Qwen3-ASR model compatible with HuggingFace weights"""
import logging
from typing import Any, Iterable, List, Optional, Tuple
import torch
import torch.nn as nn
from sglang.srt.configs.qwen3_asr import Qwen3ASRConfig
from sglang.srt.configs.qwen3_omni import Qwen3OmniMoeAudioEncoderConfig
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.managers.mm_utils import (
MultiModalityDataPaddingPatternMultimodalTokens,
general_mm_embed_routine,
)
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalInputs,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.qwen3 import Qwen3ForCausalLM
from sglang.srt.models.qwen3_omni_moe import Qwen3OmniMoeAudioEncoder
from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__)
class Qwen3ASRForConditionalGeneration(nn.Module):
default_bitsandbytes_target_modules = [
".gate_proj.",
".down_proj.",
".up_proj.",
".q_proj.",
".k_proj.",
".v_proj.",
".o_proj.",
]
bitsandbytes_stacked_params_mapping = {
"q_proj": ("qkv_proj", 0),
"k_proj": ("qkv_proj", 1),
"v_proj": ("qkv_proj", 2),
"gate_proj": ("gate_up_proj", 0),
"up_proj": ("gate_up_proj", 1),
}
def __init__(
self,
config: Qwen3ASRConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
thinker_config = config.thinker_config
if getattr(thinker_config, "audio_config", None) is None:
thinker_config.audio_config = Qwen3OmniMoeAudioEncoderConfig()
self.audio_tower = Qwen3OmniMoeAudioEncoder(thinker_config.audio_config)
self.language_model = Qwen3ForCausalLM(
thinker_config.text_config,
quant_config,
prefix=add_prefix("language_model", prefix),
)
self.pattern = MultiModalityDataPaddingPatternMultimodalTokens()
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
return self.pattern.pad_input_tokens(input_ids, mm_inputs)
def get_audio_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
device = next(self.audio_tower.parameters()).device
input_features = (
torch.cat([item.feature for item in items])
.type(self.audio_tower.dtype)
.to(device)
)
has_mask = all(
getattr(item, "feature_attention_mask", None) is not None for item in items
)
if has_mask:
feature_attention_mask = (
torch.cat([item.feature_attention_mask for item in items], dim=0)
.type(torch.long)
.to(device)
)
audio_feature_lengths = torch.sum(feature_attention_mask, dim=1)
input_features = input_features.permute(0, 2, 1)[
feature_attention_mask.bool()
].permute(1, 0)
else:
audio_feature_lengths = torch.tensor(
[input_features.shape[-1]] * input_features.shape[0],
dtype=torch.long,
device=device,
)
input_features = input_features.permute(0, 2, 1).reshape(
-1, input_features.shape[1]
)
audio_outputs = self.audio_tower(
input_features,
feature_lens=audio_feature_lengths,
)
return audio_outputs.last_hidden_state
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
**kwargs: Any,
) -> torch.Tensor:
hidden_states = general_mm_embed_routine(
input_ids=input_ids,
forward_batch=forward_batch,
language_model=self.language_model,
data_embedding_funcs={
Modality.AUDIO: self.get_audio_feature,
},
positions=positions,
)
return hidden_states
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
llm_stacked_params = [
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
("gate_up_proj", "gate_proj", 0),
("gate_up_proj", "up_proj", 1),
]
# Audio tower has separate q/k/v in checkpoint → stack into qkv_proj
audio_stacked_params = [
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
]
params_dict = dict(self.named_parameters(remove_duplicate=False))
for name, loaded_weight in weights:
if "rotary_emb.inv_freq" in name:
continue
if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name:
continue
if (
getattr(
self.config.thinker_config.text_config, "tie_word_embeddings", False
)
and "lm_head.weight" in name
):
continue
if "talker" in name or "code2wav" in name:
continue
if name.startswith("thinker.audio_tower."):
name = name.replace("thinker.audio_tower.", "audio_tower.", 1)
elif name.startswith("thinker.lm_head."):
name = name.replace("thinker.lm_head.", "language_model.lm_head.", 1)
elif name.startswith("thinker.model."):
name = name.replace("thinker.model.", "language_model.model.", 1)
is_audio = "audio_tower" in name
# Audio tower: remap out_proj → proj for VisionAttention
if is_audio and "out_proj" in name:
name = name.replace("out_proj", "proj")
stacked_params = audio_stacked_params if is_audio else llm_stacked_params
for param_name, weight_name, shard_id in stacked_params:
if weight_name not in name:
continue
name_tmp = name.replace(weight_name, param_name)
if name_tmp.endswith(".bias") and name_tmp not in params_dict:
continue
if name_tmp not in params_dict:
continue
param = params_dict[name_tmp]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
EntryClass = Qwen3ASRForConditionalGeneration
@@ -389,6 +389,7 @@ class BaseMultimodalProcessor(ABC):
"Gemma4Processor",
"GlmAsrProcessor",
"Qwen2AudioProcessor",
"Qwen3ASRProcessor",
"Qwen3OmniMoeProcessor",
}:
# Note(Xinyuan): for gemma3n, ref: https://github.com/huggingface/transformers/blob/ccf2ca162e33f381e454cdb74bf4b41a51ab976d/src/transformers/models/gemma3n/processing_gemma3n.py#L107
@@ -0,0 +1,95 @@
import re
from typing import Union
import torch
from sglang.srt.managers.schedule_batch import Modality, MultimodalProcessorOutput
from sglang.srt.models.qwen3_asr import Qwen3ASRForConditionalGeneration
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
MultimodalSpecialTokens,
)
_DEFAULT_ASR_PROMPT = (
"<|im_start|>user\n"
"<|audio_start|><|audio_pad|><|audio_end|>"
"<|im_end|>\n"
"<|im_start|>assistant\n"
)
class Qwen3ASRMultimodalProcessor(BaseMultimodalProcessor):
models = [Qwen3ASRForConditionalGeneration]
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
super().__init__(hf_config, server_args, _processor, *args, **kwargs)
self.AUDIO_TOKEN = "<|audio_start|><|audio_pad|><|audio_end|>"
self.AUDIO_TOKEN_REGEX = re.compile(
r"<\|audio_start\|>(?:<\|audio_pad\|>)+<\|audio_end\|>"
)
tokenizer = self._processor.tokenizer
self.audio_start_id = tokenizer.convert_tokens_to_ids("<|audio_start|>")
self.audio_token_id = tokenizer.convert_tokens_to_ids("<|audio_pad|>")
self.audio_end_id = tokenizer.convert_tokens_to_ids("<|audio_end|>")
self.mm_tokens = MultimodalSpecialTokens(
audio_token=self.AUDIO_TOKEN,
audio_token_regex=self.AUDIO_TOKEN_REGEX,
audio_token_id=self.audio_token_id,
).build(_processor)
self.ATTR_NAME_TO_MODALITY.update({"feature_attention_mask": Modality.AUDIO})
def _build_transcription_prompt(self, input_text: Union[str, list]) -> str:
if isinstance(input_text, list):
input_text = self._tokenizer.decode(input_text)
if not input_text or not input_text.strip():
return _DEFAULT_ASR_PROMPT
return input_text
def compute_mrope_positions(self, input_ids, mm_items):
if isinstance(input_ids, list):
seq_len = len(input_ids)
else:
seq_len = input_ids.shape[-1] if input_ids.dim() > 1 else input_ids.shape[0]
positions = torch.arange(seq_len, dtype=torch.long)
mrope_positions = positions.unsqueeze(0).expand(3, -1).clone()
return mrope_positions, torch.tensor([0], dtype=torch.long)
async def process_mm_data_async(
self,
audio_data=None,
input_text=None,
request_obj=None,
**kwargs,
):
if not audio_data:
return None
prompt = self._build_transcription_prompt(input_text)
base_output = self.load_mm_data(
prompt=prompt,
audio_data=audio_data,
multimodal_tokens=self.mm_tokens,
)
if base_output is None:
return None
mm_items, input_ids, ret = self.process_and_combine_mm_data(
base_output, self.mm_tokens
)
mrope_positions, mrope_position_delta = self.compute_mrope_positions(
input_ids, mm_items
)
return MultimodalProcessorOutput(
mm_items=mm_items,
input_ids=input_ids.tolist(),
audio_start_id=self.audio_start_id,
audio_token_id=self.audio_token_id,
audio_end_id=self.audio_end_id,
mrope_positions=mrope_positions,
mrope_position_delta=mrope_position_delta,
)