From c255483be71c5c8544cec966ca083482ff8a7e5c Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Wed, 19 Aug 2026 00:51:32 +0800 Subject: [PATCH 01/17] add single-request step-wise-execution, added decoder graph(not exercised yet) Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../models/omnivoice/pipeline_omnivoice.py | 709 ++++++++++++++---- .../entrypoints/openai/serving_speech.py | 9 + .../models/omnivoice/omnivoice_decoder.py | 102 ++- .../models/omnivoice/omnivoice_generator.py | 10 +- 4 files changed, 676 insertions(+), 154 deletions(-) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 62ac027ca8b..e6b55af65c6 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -12,13 +12,16 @@ from __future__ import annotations import json +import math import os +import random import re -from collections.abc import Iterable -from typing import ClassVar +from collections.abc import Iterable, Sequence +from typing import Any, ClassVar import numpy as np import torch +import torch.nn.functional as F from tokenizers import Tokenizer as HFTokenizer from torch import nn from vllm.logger import init_logger @@ -26,10 +29,16 @@ from vllm_omni.diffusion.data import DiffusionOutput, OmniDiffusionConfig from vllm_omni.diffusion.distributed.utils import get_local_device from vllm_omni.diffusion.models.interface import SupportAudioOutput +from vllm_omni.diffusion.worker.input_batch import InputBatch from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch +from vllm_omni.diffusion.worker.utils import StepRequestState from vllm_omni.model_executor.models.omnivoice.duration import RuleDurationEstimator from vllm_omni.model_executor.models.omnivoice.omnivoice_decoder import OmniVoiceDecoder -from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import OmniVoiceGenerator +from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( + OmniVoiceGenerator, + _get_time_steps, + _gumbel_sample, +) from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig from vllm_omni.utils.speaker_cache import get_speaker_cache @@ -132,6 +141,8 @@ class OmniVoicePipeline(nn.Module, SupportAudioOutput): """ support_audio_output: ClassVar[bool] = True + supports_request_batch: ClassVar[bool] = True + supports_step_execution: ClassVar[bool] = True def __init__(self, *, od_config: OmniDiffusionConfig, prefix: str = ""): super().__init__() @@ -205,162 +216,578 @@ def _encode_ref_audio(self, audio_signal: torch.Tensor, sr: int) -> torch.Tensor tokens = tokens.squeeze(0) # [8, T_ref] return tokens + def prepare_encode(self, states: list[StepRequestState]) -> DiffusionRequestBatch: + ref_audio = None + ref_text = None + lang = "None" + instruct = "None" + voice_name = None + device = self.device + num_cb = self.config.num_audio_codebook + mask_id = self.config.audio_mask_id + batch_target_len: list[int] = [] + batch_input_ids: list[torch.Tensor] = [] + batch_audio_mask: list[torch.Tensor] = [] + batch_attn_mask: list[torch.Tensor] = [] + if isinstance(states, StepRequestState): + states = [states] + for state in states: + prompt = state.prompt if state.prompt else "" + extra = state.sampling.extra_args or {} + seed = extra.get("seed", None) + if isinstance(prompt, dict): + # Top-level keys (used by serving_speech.py /v1/audio/speech path) + text = prompt.get("input") or prompt.get("text") or prompt.get("prompt") + ref_audio = prompt.get("ref_audio") + ref_text = prompt.get("ref_text") + voice_name = prompt.get("voice_name") + lang = prompt.get("lang") + instruct = prompt.get("instruct") + # OmniTextPrompt format (used by offline Omni.generate path): + # ref_audio comes via multi_modal_data["audio"] and the rest via + # mm_processor_kwargs. Fall back to those when top-level keys are + # absent so both invocation styles work. + mm_data = prompt.get("multi_modal_data") or {} + mm_kwargs = prompt.get("mm_processor_kwargs") or {} + if ref_audio is None: + audio_field = mm_data.get("audio") + # Standard multimodal shape allows a list of audios; OmniVoice + # voice cloning conditions on a single reference clip, so + # unwrap a length-1 list and reject multi-reference prompts up + # front (otherwise a list would later crash inside + # ``_encode_ref_audio`` when it calls ``audio.dim()``). + if isinstance(audio_field, list): + if len(audio_field) == 1: + audio_field = audio_field[0] + elif len(audio_field) > 1: + return DiffusionOutput( + error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" # noqa: E501 + ) + else: + audio_field = None + if audio_field is not None: + if isinstance(audio_field, tuple) and len(audio_field) == 2: + ref_audio = audio_field + else: + sr = mm_kwargs.get("sample_rate") or self.sample_rate + ref_audio = (audio_field, int(sr)) + if ref_text is None: + ref_text = mm_kwargs.get("ref_text") + if lang is None: + lang = mm_kwargs.get("lang") + if instruct is None: + instruct = mm_kwargs.get("instruct") + + if not text: + return DiffusionOutput(error="Empty text prompt") + lang = lang or "None" + instruct = instruct or "None" + else: + text = str(prompt) + if not text: + return DiffusionOutput(error="Empty text prompt") + + # Estimate target duration + target_len = self.duration_estimator.estimate_duration(text, "Nice to meet you.", 25) + target_len = max(1, int(target_len)) + batch_target_len.append(target_len) + + # Build text prompt with control tokens + style_text = f"<|denoise|><|lang_start|>{lang}<|lang_end|><|instruct_start|>{instruct}<|instruct_end|>" + full_text = _combine_text(ref_text=ref_text, text=text) + wrapped_text = f"<|text_start|>{full_text}<|text_end|>" + style_tokens = self.tokenizer.encode(style_text).ids + text_tokens = _tokenize_with_nonverbal_tags(wrapped_text, self.tokenizer) + encoding_ids = style_tokens + text_tokens + text_tokens = torch.tensor(encoding_ids, dtype=torch.long, device=device) + text_len = text_tokens.shape[0] + + # Encode reference audio tokens if provided (with voice caching) + ref_audio_tokens = None + if ref_audio is not None: + if self.audio_tokenizer is None: + raise RuntimeError( + "Voice cloning requires transformers>=5.3.0. Try: uv pip install 'transformers>=5.3.0'" + ) + # Check speaker cache first + _cache_key = None + if voice_name: + _cache_key = self._speaker_cache.make_cache_key( + voice_name, + model_type="omnivoice", + created_at=int(prompt.get("voice_created_at") or 0), + ) + cached = self._speaker_cache.get(_cache_key) + if cached is not None: + ref_audio_tokens = cached["ref_audio_tokens"].to(device) + _cache_key = None # hit → don't store again + logger.debug("Speaker cache HIT for OmniVoice speaker '%s'", voice_name) + + if ref_audio_tokens is None: + audio_signal, sr = ref_audio + if isinstance(audio_signal, np.ndarray): + audio_signal = torch.from_numpy(audio_signal).float() + ref_audio_tokens = self._encode_ref_audio(audio_signal, int(sr)).to(device) + + # Store in cache for next request + if _cache_key is not None: + self._speaker_cache.put(_cache_key, {"ref_audio_tokens": ref_audio_tokens.cpu()}) + logger.debug("Speaker cache STORE for OmniVoice speaker '%s'", voice_name) + + # Build conditional + unconditional batches [2, 8, max_len] + text_ids = text_tokens.unsqueeze(0).repeat(num_cb, 1) + target_ids = torch.full((num_cb, target_len), mask_id, dtype=torch.long, device=device) + + if ref_audio_tokens is not None: + cond_ids = torch.cat([text_ids, ref_audio_tokens, target_ids], dim=1) + else: + cond_ids = torch.cat([text_ids, target_ids], dim=1) + + cond_len = cond_ids.shape[1] + uncond_ids = target_ids.clone() + uncond_len = target_len + max_len = max(cond_len, uncond_len) + if uncond_len < max_len: + pad = torch.full( + (num_cb, max_len - uncond_len), + mask_id, + dtype=torch.long, + device=device, + ) + uncond_ids = torch.cat([uncond_ids, pad], dim=1) + input_ids = torch.stack([cond_ids, uncond_ids]) + batch_input_ids.append(input_ids) + + audio_mask = torch.zeros(2, max_len, dtype=torch.bool, device=device) + audio_mask[0, text_len:cond_len] = True + audio_mask[1, :uncond_len] = True + batch_audio_mask.append(audio_mask) + + attn_mask = torch.zeros(2, 1, max_len, max_len, dtype=torch.bool, device=device) + attn_mask[0, :, :cond_len, :cond_len] = True + attn_mask[1, :, :uncond_len, :uncond_len] = True + batch_attn_mask.append(attn_mask) + + if len(batch_input_ids) > 1: + max_input_ids_len = max([ids.shape[-1] for ids in batch_input_ids]) + num_padded_len = [max_input_ids_len - ids.shape[-1] for ids in batch_input_ids] + for i, ids in enumerate(batch_input_ids): + batch_input_ids[i] = torch.cat( + [ + ids, + torch.full( + (ids.shape[0], ids.shape[1], num_padded_len[i]), mask_id, dtype=torch.long, device=device + ), + ], + dim=-1, + ) + for i, mask in enumerate(batch_audio_mask): + batch_audio_mask[i] = torch.cat( + [mask, torch.full((mask.shape[0], num_padded_len[i]), False, dtype=torch.bool, device=device)], + dim=-1, + ) + for i, mask in enumerate(batch_attn_mask): + n = num_padded_len[i] + batch_attn_mask[i] = F.pad(mask, (0, n, 0, n)) + + target_lens = batch_target_len + B = len(target_lens) + device = input_ids.device + max_target_len = max(target_lens) + mask_id = self.config.audio_mask_id + num_codebooks = self.config.num_audio_codebook + if seed is None: + seed = random.randint(0, 2**63 - 1) + num_step = self.num_step + t_shift = self.t_shift + + # Initialize all target tokens as [MASK] + positions = torch.arange(max_target_len, device=device).unsqueeze(0) + valid_target_mask = positions < torch.tensor(target_lens, device=device).unsqueeze(1) + tokens = torch.zeros((B, num_codebooks, max_target_len), dtype=torch.long, device=device) + tokens.masked_fill_(valid_target_mask.unsqueeze(1), mask_id) + + timesteps = _get_time_steps(0.0, 1.0, num_step + 1, t_shift) + + # Compute unmasking schedule + schedules = [] + for t_len in target_lens: + total_mask = t_len * num_codebooks + rem = total_mask + sched = [] + for step in range(num_step): + num = ( + rem + if step == num_step - 1 + else min( + math.ceil(total_mask * (timesteps[step + 1] - timesteps[step])), + rem, + ) + ) + sched.append(int(num)) + rem -= int(num) + schedules.append(sched) + schedules = torch.tensor(schedules, dtype=torch.long, device=device) + + layer_ids = torch.arange(num_codebooks, device=device).view(1, -1, 1) + generator = torch.Generator(device=device).manual_seed(seed) + + for i in range(B): + states[i].latents = batch_input_ids[i] + states[i].timesteps = schedules[i] + states[i].guidance = self.guidance_scale + states[i].extra["schedules"] = schedules + states[i].extra["layer_ids"] = layer_ids + states[i].extra["generator"] = generator + states[i].extra["t_shift"] = t_shift + states[i].extra["target_len"] = target_lens[i] + states[i].extra["audio_mask"] = batch_audio_mask[i] + states[i].extra["attn_mask"] = batch_attn_mask[i] + states[i].extra["tokens"] = tokens[i] + + use_cuda_graph = self.generator._cuda_graph_fwd is not None and input_ids.is_cuda + if not use_cuda_graph: + # Eager-path-only constants (the cuda-graph captures its own). + text_embeds_cached = self.text_embedding(input_ids[:, 0, :]) + audio_mask_3d = audio_mask.unsqueeze(-1) + self._ensure_rope(input_ids.shape[-1], device) + target_dtype = text_embeds_cached.dtype + cos = self._rope_cos.to(device=device, dtype=target_dtype) + sin = self._rope_sin.to(device=device, dtype=target_dtype) + + for i in range(B): + states[i].extra["text_embeds"] = text_embeds_cached + states[i].extra["audio_mask_3d"] = audio_mask_3d + states[i].extra["cos"] = cos + states[i].extra["sin"] = sin + + def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestState] | None = None, **kwargs: Any): + use_cuda_graph = self.generator._cuda_graph_fwd is not None + input_ids = input_batch.latents + layer_ids = states[0].extra["layer_ids"] + generator = states[0].extra["generator"] + schedules = states[0].extra["schedules"] + + batch_audio_mask: list[torch.Tesor] = [] + batch_attn_mask: list[torch.Tensor] = [] + batch_target_len: list[int] = [] + batch_tokens: list[torch.Tensor] = [] + + steps: list[int] = [] + + for state in states: + batch_audio_mask.append(state.extra.get("audio_mask", None)) + batch_attn_mask.append(state.extra.get("attn_mask", None)) + guidance_scale = state.extra.get("guidance", self.guidance_scale) + batch_target_len.append(state.extra["target_len"]) + batch_tokens.append(state.extra["tokens"]) + steps.append(state.step_index) + batch_tokens = torch.stack(batch_tokens, dim=0) + if batch_tokens.dim() == 4: + batch_tokens = batch_tokens.squeeze(0) + if len(batch_audio_mask) == 1: + audio_mask = batch_audio_mask[0] + batch_attn_mask = batch_attn_mask[0] + else: + audio_mask = torch.stack(batch_audio_mask, dim=0) + batch_attn_mask = torch.stack(batch_attn_mask, dim=0) + + B = len(batch_target_len) + target_lens = batch_target_len + + text_embeds_cached = states[0].extra.get("text_embeds", None) + audio_mask_3d = states[0].extra.get("audio_mask_3d", None) + cos = states[0].extra.get("cos", None) + sin = states[0].extra.get("sin", None) + + mask_id = self.config.audio_mask_id + position_temperature = self.position_temperature + class_temperature = self.class_temperature + layer_penalty_factor = self.layer_penalty_factor + + c_lens = batch_attn_mask[:B, 0, 0].sum(dim=-1).tolist() + + # Materialize the SDPA float mask once so the captured graph (and eager path) skip per-layer conversion. + sdpa_attn_mask = torch.zeros_like(batch_attn_mask, dtype=torch.float32).masked_fill_( + ~batch_attn_mask, float("-inf") + ) + if use_cuda_graph: + # Float mask skips per-layer conversion; fp32 cast deferred to the per-item slices below. + batch_logits = self.generator._cuda_graph_fwd(input_ids, audio_mask, sdpa_attn_mask) + else: + # Eager fallback reuses hoisted constants (text embeds, sdpa mask, cos/sin). + inputs_embeds = self.generator._prepare_embeddings( + input_ids, audio_mask, text_embeds=text_embeds_cached, audio_mask_3d=audio_mask_3d + ) + hidden_states = self.generator._transformer_forward(inputs_embeds, sdpa_attn_mask, cos=cos, sin=sin) + # fp32 cast deferred to the per-item slices below. + batch_logits = self.generator._get_logits(hidden_states) + # batch_logits: [2*B, 8, S, 1025] + + for i in range(B): + k = schedules[i][steps[i]] + if k <= 0: + continue + + c_len = c_lens[i] + t_len = target_lens[i] + + # Extract logits for target region; upcast only the slices we actually consume. + c_logits = batch_logits[i : i + 1, :, c_len - t_len : c_len, :].to(torch.float32) + u_logits = batch_logits[B + i : B + i + 1, :, :t_len, :].to(torch.float32) + + # Classifier-free guidance. Fuse the chain: the two inner + # log_softmax normalizers are per-position scalars that the final + # shift-invariant log_softmax cancels, so guide on the raw logits + # with a single softmax: log_softmax((1+s)*c - s*u). Exact. + if guidance_scale != 0: + log_probs = F.log_softmax( + (1.0 + guidance_scale) * c_logits - guidance_scale * u_logits, + dim=-1, + ) + else: + log_probs = F.log_softmax(c_logits, dim=-1) + + # Prevent predicting [MASK] + log_probs[..., mask_id] = -float("inf") + + # Token prediction + if class_temperature > 0.0: + pred_tokens = _gumbel_sample(log_probs, class_temperature, generator).argmax(dim=-1) + else: + pred_tokens = log_probs.argmax(dim=-1) # [1, 8, T] + + # Confidence scores + scores = log_probs.max(dim=-1)[0] # [1, 8, T] + + # Layer penalty (earlier codebooks get higher priority) + scores = scores - (layer_ids * layer_penalty_factor) + + # Gumbel noise for position selection + if position_temperature > 0.0: + scores = _gumbel_sample(scores, position_temperature, generator) + + # Mask out already unmasked positions + sample_tokens = batch_tokens[i : i + 1, :, :t_len] + scores.masked_fill_(sample_tokens != mask_id, -float("inf")) + + # Select top-k positions to unmask. .flatten() on this non-contiguous view already copies. + _, topk_idx = torch.topk(scores.flatten(), k) + flat_tokens = sample_tokens.flatten() + flat_tokens[topk_idx] = pred_tokens.flatten()[topk_idx] + sample_tokens.copy_(flat_tokens.view_as(sample_tokens)) + states[i].extra["tokens"] = sample_tokens + + # Mirror update into both cond and uncond input_ids halves for the next step. + input_ids[i, :, c_len - t_len : c_len] = sample_tokens.squeeze(0) + input_ids[B + i, :, :t_len] = sample_tokens.squeeze(0) + + return input_ids + + def step_scheduler(self, state: StepRequestState, noise_pred: torch.Tensor, **kwargs: Any): + state.latents = noise_pred + state.step_index += 1 + + def post_decode(self, state: StepRequestState, **kwargs: Any): + tokens = state.extra.get("tokens", None) + if tokens.dim() == 2: + tokens = tokens.unsqueeze(0) + audio = self.decoder(tokens) + return DiffusionOutput(output=audio) + @torch.inference_mode() - def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: """Generate speech audio from text, optionally with voice cloning. Accepts either a plain text prompt or a structured dict: {"text": "...", "ref_audio": (samples, sr), "ref_text": "...", "lang": "...", "instruct": "..."} """ - prompt = req.prompts[0] if req.prompts else "" ref_audio = None ref_text = None lang = "None" instruct = "None" - extra = req.sampling_params.extra_args or {} - seed = extra.get("seed", None) - voice_name = None - if isinstance(prompt, dict): - # Top-level keys (used by serving_speech.py /v1/audio/speech path) - text = prompt.get("input") or prompt.get("text") or prompt.get("prompt") - ref_audio = prompt.get("ref_audio") - ref_text = prompt.get("ref_text") - voice_name = prompt.get("voice_name") - lang = prompt.get("lang") - instruct = prompt.get("instruct") - # OmniTextPrompt format (used by offline Omni.generate path): - # ref_audio comes via multi_modal_data["audio"] and the rest via - # mm_processor_kwargs. Fall back to those when top-level keys are - # absent so both invocation styles work. - mm_data = prompt.get("multi_modal_data") or {} - mm_kwargs = prompt.get("mm_processor_kwargs") or {} - if ref_audio is None: - audio_field = mm_data.get("audio") - # Standard multimodal shape allows a list of audios; OmniVoice - # voice cloning conditions on a single reference clip, so - # unwrap a length-1 list and reject multi-reference prompts up - # front (otherwise a list would later crash inside - # ``_encode_ref_audio`` when it calls ``audio.dim()``). - if isinstance(audio_field, list): - if len(audio_field) == 1: - audio_field = audio_field[0] - elif len(audio_field) > 1: - return DiffusionOutput( - error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" # noqa: E501 - ) - else: - audio_field = None - if audio_field is not None: - if isinstance(audio_field, tuple) and len(audio_field) == 2: - ref_audio = audio_field - else: - sr = mm_kwargs.get("sample_rate") or self.sample_rate - ref_audio = (audio_field, int(sr)) - if ref_text is None: - ref_text = mm_kwargs.get("ref_text") - if lang is None: - lang = mm_kwargs.get("lang") - if instruct is None: - instruct = mm_kwargs.get("instruct") - - if not text: - return DiffusionOutput(error="Empty text prompt") - lang = lang or "None" - instruct = instruct or "None" - else: - text = str(prompt) - if not text: - return DiffusionOutput(error="Empty text prompt") - + outputs: list[DiffusionOutput] = [] device = self.device num_cb = self.config.num_audio_codebook mask_id = self.config.audio_mask_id - - # Estimate target duration - target_len = self.duration_estimator.estimate_duration(text, "Nice to meet you.", 25) - target_len = max(1, int(target_len)) - - # Build text prompt with control tokens - style_text = f"<|denoise|><|lang_start|>{lang}<|lang_end|><|instruct_start|>{instruct}<|instruct_end|>" - full_text = _combine_text(ref_text=ref_text, text=text) - wrapped_text = f"<|text_start|>{full_text}<|text_end|>" - style_tokens = self.tokenizer.encode(style_text).ids - text_tokens = _tokenize_with_nonverbal_tags(wrapped_text, self.tokenizer) - encoding_ids = style_tokens + text_tokens - text_tokens = torch.tensor(encoding_ids, dtype=torch.long, device=device) - text_len = text_tokens.shape[0] - - # Encode reference audio tokens if provided (with voice caching) - ref_audio_tokens = None - if ref_audio is not None: - if self.audio_tokenizer is None: - raise RuntimeError( - "Voice cloning requires transformers>=5.3.0. Try: uv pip install 'transformers>=5.3.0'" + batch_target_len: list[int] = [] + batch_input_ids: list[torch.Tensor] = [] + batch_audio_mask: list[torch.Tensor] = [] + batch_attn_mask: list[torch.Tensor] = [] + for request in req.requests: + prompt = request.prompt if request.prompt else "" + extra = request.sampling_params.extra_args or {} + seed = extra.get("seed", None) + if isinstance(prompt, dict): + # Top-level keys (used by serving_speech.py /v1/audio/speech path) + text = prompt.get("input") or prompt.get("text") or prompt.get("prompt") + ref_audio = prompt.get("ref_audio") + ref_text = prompt.get("ref_text") + voice_name = prompt.get("voice_name") + lang = prompt.get("lang") + instruct = prompt.get("instruct") + # OmniTextPrompt format (used by offline Omni.generate path): + # ref_audio comes via multi_modal_data["audio"] and the rest via + # mm_processor_kwargs. Fall back to those when top-level keys are + # absent so both invocation styles work. + mm_data = prompt.get("multi_modal_data") or {} + mm_kwargs = prompt.get("mm_processor_kwargs") or {} + if ref_audio is None: + audio_field = mm_data.get("audio") + # Standard multimodal shape allows a list of audios; OmniVoice + # voice cloning conditions on a single reference clip, so + # unwrap a length-1 list and reject multi-reference prompts up + # front (otherwise a list would later crash inside + # ``_encode_ref_audio`` when it calls ``audio.dim()``). + if isinstance(audio_field, list): + if len(audio_field) == 1: + audio_field = audio_field[0] + elif len(audio_field) > 1: + return DiffusionOutput( + error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" # noqa: E501 + ) + else: + audio_field = None + if audio_field is not None: + if isinstance(audio_field, tuple) and len(audio_field) == 2: + ref_audio = audio_field + else: + sr = mm_kwargs.get("sample_rate") or self.sample_rate + ref_audio = (audio_field, int(sr)) + if ref_text is None: + ref_text = mm_kwargs.get("ref_text") + if lang is None: + lang = mm_kwargs.get("lang") + if instruct is None: + instruct = mm_kwargs.get("instruct") + + if not text: + return DiffusionOutput(error="Empty text prompt") + lang = lang or "None" + instruct = instruct or "None" + else: + text = str(prompt) + if not text: + return DiffusionOutput(error="Empty text prompt") + + # Estimate target duration + target_len = self.duration_estimator.estimate_duration(text, "Nice to meet you.", 25) + target_len = max(1, int(target_len)) + batch_target_len.append(target_len) + + # Build text prompt with control tokens + style_text = f"<|denoise|><|lang_start|>{lang}<|lang_end|><|instruct_start|>{instruct}<|instruct_end|>" + full_text = _combine_text(ref_text=ref_text, text=text) + wrapped_text = f"<|text_start|>{full_text}<|text_end|>" + style_tokens = self.tokenizer.encode(style_text).ids + text_tokens = _tokenize_with_nonverbal_tags(wrapped_text, self.tokenizer) + encoding_ids = style_tokens + text_tokens + text_tokens = torch.tensor(encoding_ids, dtype=torch.long, device=device) + text_len = text_tokens.shape[0] + + # Encode reference audio tokens if provided (with voice caching) + ref_audio_tokens = None + if ref_audio is not None: + if self.audio_tokenizer is None: + raise RuntimeError( + "Voice cloning requires transformers>=5.3.0. Try: uv pip install 'transformers>=5.3.0'" + ) + # Check speaker cache first + _cache_key = None + if voice_name: + _cache_key = self._speaker_cache.make_cache_key( + voice_name, + model_type="omnivoice", + created_at=int(prompt.get("voice_created_at") or 0), + ) + cached = self._speaker_cache.get(_cache_key) + if cached is not None: + ref_audio_tokens = cached["ref_audio_tokens"].to(device) + _cache_key = None # hit → don't store again + logger.debug("Speaker cache HIT for OmniVoice speaker '%s'", voice_name) + + if ref_audio_tokens is None: + audio_signal, sr = ref_audio + if isinstance(audio_signal, np.ndarray): + audio_signal = torch.from_numpy(audio_signal).float() + ref_audio_tokens = self._encode_ref_audio(audio_signal, int(sr)).to(device) + + # Store in cache for next request + if _cache_key is not None: + self._speaker_cache.put(_cache_key, {"ref_audio_tokens": ref_audio_tokens.cpu()}) + logger.debug("Speaker cache STORE for OmniVoice speaker '%s'", voice_name) + + # Build conditional + unconditional batches [2, 8, max_len] + text_ids = text_tokens.unsqueeze(0).repeat(num_cb, 1) + target_ids = torch.full((num_cb, target_len), mask_id, dtype=torch.long, device=device) + + if ref_audio_tokens is not None: + cond_ids = torch.cat([text_ids, ref_audio_tokens, target_ids], dim=1) + else: + cond_ids = torch.cat([text_ids, target_ids], dim=1) + cond_len = cond_ids.shape[1] + uncond_ids = target_ids.clone() + uncond_len = target_len + max_len = max(cond_len, uncond_len) + if uncond_len < max_len: + pad = torch.full( + (num_cb, max_len - uncond_len), + mask_id, + dtype=torch.long, + device=device, ) - # Check speaker cache first - _cache_key = None - if voice_name: - _cache_key = self._speaker_cache.make_cache_key( - voice_name, - model_type="omnivoice", - created_at=int(prompt.get("voice_created_at") or 0), + uncond_ids = torch.cat([uncond_ids, pad], dim=1) + input_ids = torch.stack([cond_ids, uncond_ids]) + batch_input_ids.append(input_ids) + + audio_mask = torch.zeros(2, max_len, dtype=torch.bool, device=device) + audio_mask[0, text_len:cond_len] = True + audio_mask[1, :uncond_len] = True + batch_audio_mask.append(audio_mask) + + attn_mask = torch.zeros(2, 1, max_len, max_len, dtype=torch.bool, device=device) + attn_mask[0, :, :cond_len, :cond_len] = True + attn_mask[1, :, :uncond_len, :uncond_len] = True + batch_attn_mask.append(attn_mask) + if len(batch_input_ids) > 1: + max_input_ids_len = max([ids.shape[-1] for ids in batch_input_ids]) + num_padded_len = [max_input_ids_len - ids.shape[-1] for ids in batch_input_ids] + for i, ids in enumerate(batch_input_ids): + batch_input_ids[i] = torch.cat( + [ + ids, + torch.full( + (ids.shape[0], ids.shape[1], num_padded_len[i]), mask_id, dtype=torch.long, device=device + ), + ], + dim=-1, ) - cached = self._speaker_cache.get(_cache_key) - if cached is not None: - ref_audio_tokens = cached["ref_audio_tokens"].to(device) - _cache_key = None # hit → don't store again - logger.debug("Speaker cache HIT for OmniVoice speaker '%s'", voice_name) - - if ref_audio_tokens is None: - audio_signal, sr = ref_audio - if isinstance(audio_signal, np.ndarray): - audio_signal = torch.from_numpy(audio_signal).float() - ref_audio_tokens = self._encode_ref_audio(audio_signal, int(sr)).to(device) - - # Store in cache for next request - if _cache_key is not None: - self._speaker_cache.put(_cache_key, {"ref_audio_tokens": ref_audio_tokens.cpu()}) - logger.debug("Speaker cache STORE for OmniVoice speaker '%s'", voice_name) - - # Build conditional + unconditional batches [2, 8, max_len] - text_ids = text_tokens.unsqueeze(0).repeat(num_cb, 1) - target_ids = torch.full((num_cb, target_len), mask_id, dtype=torch.long, device=device) - - if ref_audio_tokens is not None: - cond_ids = torch.cat([text_ids, ref_audio_tokens, target_ids], dim=1) - else: - cond_ids = torch.cat([text_ids, target_ids], dim=1) - cond_len = cond_ids.shape[1] - uncond_ids = target_ids.clone() - uncond_len = target_len - max_len = max(cond_len, uncond_len) - if uncond_len < max_len: - pad = torch.full( - (num_cb, max_len - uncond_len), - mask_id, - dtype=torch.long, - device=device, - ) - uncond_ids = torch.cat([uncond_ids, pad], dim=1) - - batch_input_ids = torch.stack([cond_ids, uncond_ids]) - - batch_audio_mask = torch.zeros(2, max_len, dtype=torch.bool, device=device) - batch_audio_mask[0, text_len:cond_len] = True - batch_audio_mask[1, :uncond_len] = True - - batch_attn_mask = torch.zeros(2, 1, max_len, max_len, dtype=torch.bool, device=device) - batch_attn_mask[0, :, :cond_len, :cond_len] = True - batch_attn_mask[1, :, :uncond_len, :uncond_len] = True - + for i, mask in enumerate(batch_audio_mask): + batch_audio_mask[i] = torch.cat( + [mask, torch.full((mask.shape[0], num_padded_len[i]), False, dtype=torch.bool, device=device)], + dim=-1, + ) + for i, mask in enumerate(batch_attn_mask): + n = num_padded_len[i] + batch_attn_mask[i] = F.pad(mask, (0, n, 0, n)) + # Each per-request tensor is ordered as [cond_i, uncond_i]. The + # generator expects the full batch to be ordered as + # [cond_0, ..., cond_B-1, uncond_0, ..., uncond_B-1] so that request i + # pairs logits i and B+i for classifier-free guidance. + input_id_pairs = torch.stack(batch_input_ids, dim=0) + audio_mask_pairs = torch.stack(batch_audio_mask, dim=0) + attn_mask_pairs = torch.stack(batch_attn_mask, dim=0) + batch_input_ids = torch.cat([input_id_pairs[:, 0], input_id_pairs[:, 1]], dim=0) + batch_audio_mask = torch.cat([audio_mask_pairs[:, 0], audio_mask_pairs[:, 1]], dim=0) + batch_attn_mask = torch.cat([attn_mask_pairs[:, 0], attn_mask_pairs[:, 1]], dim=0) # Run 32-step iterative unmasking tokens = self.generator( input_ids=batch_input_ids, audio_mask=batch_audio_mask, attention_mask=batch_attn_mask, - target_lens=[target_len], + target_lens=batch_target_len, num_step=self.num_step, guidance_scale=self.guidance_scale, t_shift=self.t_shift, @@ -369,10 +796,12 @@ def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: class_temperature=self.class_temperature, seed=seed, ) - # Decode tokens to audio - audio = self.decoder(tokens) # [1, 1, samples] - return DiffusionOutput(output=audio) + audio = self.decoder(tokens) # [1, 1, target_len * 960] + for i in range(len(batch_target_len)): + audio_output = audio[i : i + 1, :, : batch_target_len[i] * 960] + outputs.append(DiffusionOutput(output=audio_output)) + return outputs def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: """Load weights from model directory (not from the iterator). diff --git a/vllm_omni/entrypoints/openai/serving_speech.py b/vllm_omni/entrypoints/openai/serving_speech.py index 813d916feb7..767077381b6 100644 --- a/vllm_omni/entrypoints/openai/serving_speech.py +++ b/vllm_omni/entrypoints/openai/serving_speech.py @@ -3505,6 +3505,15 @@ async def _create_diffusion_speech( if sampling_params_list[0].extra_args is None: sampling_params_list[0].extra_args = {} sampling_params_list[0].extra_args.update(extra) + + sampling = sampling_params_list[0] + + if "num_inference_steps" in extra: + sampling.num_inference_steps = int(extra["num_inference_steps"]) + + if "guidance_scale" in extra: + sampling.guidance_scale = float(extra["guidance_scale"]) + sampling.guidance_scale_provided = True logger.info("Applied extra_params to diffusion: %s", extra) generator = self._diffusion_engine.generate( diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py index efb6d9dc7af..c93b5512643 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py @@ -21,7 +21,9 @@ import torch import torch.nn as nn +from torch.cuda.graphs import CUDAGraph from vllm.logger import init_logger +from vllm.platforms import current_platform from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig @@ -68,6 +70,68 @@ def decode(self, codes: torch.Tensor) -> torch.Tensor: return result +class DecoderGraph: + def __init__( + self, + decode_fn, + batch_sizes: tuple[int] = (1, 2, 4, 8, 16), + bucket_sizes: tuple[int] = (64, 128, 256), + device: torch.device = torch.device("cuda"), + ): + self.graphs: dict[tuple[int, int], CUDAGraph] = {} + self.capture_batch_size = batch_sizes + self.device = device + self.capture_bucket_size = bucket_sizes + self.decode_fn = decode_fn + self.static_inputs: dict[tuple[int, int], torch.Tensor] = {} + self.static_outputs: dict[tuple[int, int], torch.Tensor] = {} + self.captured = False + + def warmup(self): + for batch_size in self.capture_batch_size: + for bucket_size in self.capture_bucket_size: + self._capture(batch_size, bucket_size) + self.captured = bool(self.graphs) + + def _capture(self, batch_size: int, bucket_size: int): + if torch.cuda.is_current_stream_capturing() or self.device.type != "cuda": + return + static_inputs = torch.zeros((batch_size, 8, bucket_size), device=self.device, dtype=torch.long) + self.static_inputs[(batch_size, bucket_size)] = static_inputs + for _ in range(3): + self.decode_fn(static_inputs) + graph = CUDAGraph() + with torch.cuda.graph(graph, pool=current_platform.get_global_graph_pool()): + static_outputs = self.decode_fn(static_inputs) + self.graphs[(batch_size, bucket_size)] = graph + self.static_outputs[(batch_size, bucket_size)] = static_outputs + logger.info(f"Captured graph for batch_size={batch_size}, bucket_size={bucket_size}") + + def find_nearest_padding(self, batch, bucket): + candidates = [(b, bk) for b, bk in self.graphs.keys() if b >= batch and bk >= bucket] + return ( + min(candidates, key=lambda x: x[0] * x[1]) + if candidates + else max(self.graphs.keys(), key=lambda x: x[0] * x[1]) + ) + + def forward(self, codes): + batch_size, _, bucket_size = codes.shape + if batch_size > max(self.capture_batch_size) or bucket_size > max(self.capture_bucket_size): + return self.decode_fn(codes) + if (batch_size, bucket_size) in self.graphs: + self.static_inputs[(batch_size, bucket_size)].copy_(codes) + self.graphs[(batch_size, bucket_size)].replay() + return self.static_outputs[(batch_size, bucket_size)].clone() + else: + nearest_batch, nearest_bucket = self.find_nearest_padding(batch_size, bucket_size) + static_input = self.static_inputs[(nearest_batch, nearest_bucket)] + static_input.zero_() + static_input[:batch_size, :, :bucket_size].copy_(codes) + self.graphs[(nearest_batch, nearest_bucket)].replay() + return self.static_outputs[(nearest_batch, nearest_bucket)][:batch_size].clone() + + class OmniVoiceDecoder(nn.Module): """OmniVoice Stage 1: Token-to-audio decoder. @@ -86,8 +150,7 @@ def __init__(self, config: OmniVoiceConfig): self.fc2 = None self.acoustic_decoder = None - @torch.inference_mode() - def forward(self, audio_codes: torch.Tensor) -> torch.Tensor: + def _decode_impl(self, audio_codes: torch.Tensor): """Decode audio tokens to waveform. Args: @@ -96,11 +159,6 @@ def forward(self, audio_codes: torch.Tensor) -> torch.Tensor: Returns: waveform: [B, 1, audio_samples] at 24kHz """ - if not self._loaded: - raise RuntimeError("Decoder not loaded. Call load_weights() first.") - - device = audio_codes.device - # Transpose: [B, 8, T] → [8, B, T] codes = audio_codes.transpose(0, 1).long() @@ -120,7 +178,27 @@ def forward(self, audio_codes: torch.Tensor) -> torch.Tensor: if audio.dim() == 2: audio = audio.unsqueeze(1) - return audio.to(device) + return audio + + @torch.inference_mode() + def forward(self, audio_codes: torch.Tensor) -> torch.Tensor: + """Decode audio tokens to waveform. + + Args: + audio_codes: [B, 8, T] - 8-codebook audio token IDs + + Returns: + waveform: [B, 1, audio_samples] at 24kHz + """ + if not self._loaded: + raise RuntimeError("Decoder not loaded. Call load_weights() first.") + + device = audio_codes.device + + if self.graphs.captured: + return self.graphs.forward(audio_codes).to(device) + else: + return self._decode_impl(audio_codes).to(device) def _adjust_output_padding(self, decoder: nn.Module): """Adjust ConvTranspose1d output_padding (HiggsAudioV2 modification).""" @@ -210,6 +288,14 @@ def load_weights(self, model_dir: str, device: torch.device) -> None: self.acoustic_decoder.eval() self._loaded = True + self.graphs = DecoderGraph( + self._decode_impl, + batch_sizes=(1, 2, 4, 8, 16), + bucket_sizes=(64, 128, 256), + device=device, + ) + self.graphs.warmup() + logger.info( "Loaded OmniVoice decoder: %d quantizers, fc2(%d→%d), acoustic decoder (%d weights)", num_quantizers, diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py index a75c4e6f7fc..01bf3de734e 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py @@ -910,12 +910,10 @@ def forward( generator = torch.Generator(device=device).manual_seed(seed) # Initialize all target tokens as [MASK] - tokens = torch.full( - (B, num_codebooks, max_target_len), - mask_id, - dtype=torch.long, - device=device, - ) + positions = torch.arange(max_target_len, device=device).unsqueeze(0) + valid_target_mask = positions < torch.tensor(target_lens, device=device).unsqueeze(1) + tokens = torch.zeros((B, num_codebooks, max_target_len), dtype=torch.long, device=device) + tokens.masked_fill_(valid_target_mask.unsqueeze(1), mask_id) # Compute unmasking schedule timesteps = _get_time_steps(0.0, 1.0, num_step + 1, t_shift).tolist() From 6d8fa892c9f370fe9eb899653ee53d160f444625 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Wed, 19 Aug 2026 10:02:46 +0800 Subject: [PATCH 02/17] support bit-exact for padding scenarios Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../models/omnivoice/test_decoder.py | 57 ++++++++ .../models/omnivoice/pipeline_omnivoice.py | 2 +- .../models/omnivoice/omnivoice_decoder.py | 134 ++++++++++++++---- 3 files changed, 161 insertions(+), 32 deletions(-) diff --git a/tests/model_executor/models/omnivoice/test_decoder.py b/tests/model_executor/models/omnivoice/test_decoder.py index db8e7328ea1..040fde89612 100644 --- a/tests/model_executor/models/omnivoice/test_decoder.py +++ b/tests/model_executor/models/omnivoice/test_decoder.py @@ -212,3 +212,60 @@ def test_acoustic_decoder_weights_are_float32(): assert param.dtype == torch.float32, ( f"acoustic_decoder.{name} is {param.dtype} — must be float32 (see load_weights fix)" ) + + +# --------------------------------------------------------------------------- +# 4. Bit-Exact when Input is padded +# --------------------------------------------------------------------------- + + +def _build_padded_decoder() -> OmniVoiceDecoder: + from transformers import DacConfig, DacModel + + torch.manual_seed(42) + + decoder = OmniVoiceDecoder(OmniVoiceConfig()) + + decoder.quantizer = HiggsAudioRVQ( + num_quantizers=8, + codebook_size=1024, + codebook_dim=64, + hidden_size=1024, + ).to(DEVICE) + + decoder.fc2 = nn.Linear(1024, 256).to(DEVICE).float() + + dac_config = DacConfig( + hidden_size=256, + decoder_hidden_size=64, + upsampling_ratios=[2, 2], + ) + decoder.acoustic_decoder = DacModel(dac_config).decoder.to(DEVICE).float().eval() + + decoder._adjust_output_padding(decoder.acoustic_decoder) + decoder.acoustic_decoder.tanh = nn.Identity() + + decoder.graphs = None + decoder._loaded = True + return decoder + + +def test_output_bit_exact_when_frames_are_padded(): + decoder = _build_padded_decoder() + exact_inputs = torch.randint( + 0, + 1024, + (1, 8, T_FRAMES), + device=DEVICE, + ) + + padded_inputs = torch.nn.functional.pad( + exact_inputs, + (0, 1), + value=0, + ) + padded_length = [T_FRAMES] + exact_output = decoder(exact_inputs) + padded_output = decoder(padded_inputs, padded_length) + padded_output = padded_output[:, :, : T_FRAMES * UPSAMPLE] + torch.testing.assert_close(exact_output, padded_output) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index e6b55af65c6..719c2332d47 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -797,7 +797,7 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: seed=seed, ) - audio = self.decoder(tokens) # [1, 1, target_len * 960] + audio = self.decoder(tokens, batch_target_len) # [B, 1, max_target_len * 960] for i in range(len(batch_target_len)): audio_output = audio[i : i + 1, :, : batch_target_len[i] * 960] outputs.append(DiffusionOutput(output=audio_output)) diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py index c93b5512643..b262cb46082 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py @@ -74,7 +74,7 @@ class DecoderGraph: def __init__( self, decode_fn, - batch_sizes: tuple[int] = (1, 2, 4, 8, 16), + batch_sizes: tuple[int] = (1, 2), bucket_sizes: tuple[int] = (64, 128, 256), device: torch.device = torch.device("cuda"), ): @@ -84,6 +84,7 @@ def __init__( self.capture_bucket_size = bucket_sizes self.decode_fn = decode_fn self.static_inputs: dict[tuple[int, int], torch.Tensor] = {} + self.static_lengths: dict[tuple[int, int], torch.Tensor] = {} self.static_outputs: dict[tuple[int, int], torch.Tensor] = {} self.captured = False @@ -97,39 +98,44 @@ def _capture(self, batch_size: int, bucket_size: int): if torch.cuda.is_current_stream_capturing() or self.device.type != "cuda": return static_inputs = torch.zeros((batch_size, 8, bucket_size), device=self.device, dtype=torch.long) + static_lengths = torch.full((batch_size,), bucket_size, device=self.device, dtype=torch.long) self.static_inputs[(batch_size, bucket_size)] = static_inputs + self.static_lengths[(batch_size, bucket_size)] = static_lengths for _ in range(3): - self.decode_fn(static_inputs) + self.decode_fn(static_inputs, static_lengths) graph = CUDAGraph() with torch.cuda.graph(graph, pool=current_platform.get_global_graph_pool()): - static_outputs = self.decode_fn(static_inputs) + static_outputs = self.decode_fn(static_inputs, static_lengths) self.graphs[(batch_size, bucket_size)] = graph self.static_outputs[(batch_size, bucket_size)] = static_outputs logger.info(f"Captured graph for batch_size={batch_size}, bucket_size={bucket_size}") def find_nearest_padding(self, batch, bucket): - candidates = [(b, bk) for b, bk in self.graphs.keys() if b >= batch and bk >= bucket] - return ( - min(candidates, key=lambda x: x[0] * x[1]) - if candidates - else max(self.graphs.keys(), key=lambda x: x[0] * x[1]) - ) + candidates = [ + (b, bk) for b in self.capture_batch_size for bk in self.capture_bucket_size if b >= batch and bk >= bucket + ] + return min(candidates, key=lambda x: x[0] * x[1]) if candidates else None - def forward(self, codes): + def forward(self, codes, lengths): batch_size, _, bucket_size = codes.shape if batch_size > max(self.capture_batch_size) or bucket_size > max(self.capture_bucket_size): - return self.decode_fn(codes) - if (batch_size, bucket_size) in self.graphs: - self.static_inputs[(batch_size, bucket_size)].copy_(codes) - self.graphs[(batch_size, bucket_size)].replay() - return self.static_outputs[(batch_size, bucket_size)].clone() - else: - nearest_batch, nearest_bucket = self.find_nearest_padding(batch_size, bucket_size) - static_input = self.static_inputs[(nearest_batch, nearest_bucket)] - static_input.zero_() - static_input[:batch_size, :, :bucket_size].copy_(codes) - self.graphs[(nearest_batch, nearest_bucket)].replay() - return self.static_outputs[(nearest_batch, nearest_bucket)][:batch_size].clone() + return self.decode_fn(codes, lengths) + + graph_key = self.find_nearest_padding(batch_size, bucket_size) + if graph_key is None: + return self.decode_fn(codes, lengths) + + if graph_key not in self.graphs: + return self.decode_fn(codes, lengths) + + static_input = self.static_inputs[graph_key] + static_lengths = self.static_lengths[graph_key] + static_input.zero_() + static_lengths.zero_() + static_input[:batch_size, :, :bucket_size].copy_(codes) + static_lengths[:batch_size].copy_(lengths) + self.graphs[graph_key].replay() + return self.static_outputs[graph_key][:batch_size].clone() class OmniVoiceDecoder(nn.Module): @@ -149,8 +155,48 @@ def __init__(self, config: OmniVoiceConfig): self.quantizer = None self.fc2 = None self.acoustic_decoder = None + self.graphs: DecoderGraph | None = None + + @staticmethod + def _mask_by_lengths(hidden_states: torch.Tensor, lengths: torch.Tensor) -> torch.Tensor: + positions = torch.arange(hidden_states.shape[-1], device=hidden_states.device) + valid = positions.unsqueeze(0) < lengths.unsqueeze(1) + return hidden_states.masked_fill(~valid.unsqueeze(1), 0.0) - def _decode_impl(self, audio_codes: torch.Tensor): + def _decode_acoustic_length_aware( + self, + hidden_states: torch.Tensor, + lengths: torch.Tensor, + ) -> torch.Tensor: + """Run the non-causal DAC decoder while zeroing padded activations.""" + decoder = self.acoustic_decoder + hidden_states = self._mask_by_lengths(hidden_states, lengths) + hidden_states = decoder.conv1(hidden_states) + hidden_states = self._mask_by_lengths(hidden_states, lengths) + + for block in decoder.block: + hidden_states = block.snake1(hidden_states) + hidden_states = block.conv_t1(hidden_states) + stride = block.conv_t1.stride[0] + lengths = lengths * stride + hidden_states = self._mask_by_lengths(hidden_states, lengths) + + for residual_unit in (block.res_unit1, block.res_unit2, block.res_unit3): + hidden_states = residual_unit(hidden_states) + hidden_states = self._mask_by_lengths(hidden_states, lengths) + + hidden_states = decoder.snake1(hidden_states) + hidden_states = self._mask_by_lengths(hidden_states, lengths) + hidden_states = decoder.conv2(hidden_states) + hidden_states = self._mask_by_lengths(hidden_states, lengths) + hidden_states = decoder.tanh(hidden_states) + return self._mask_by_lengths(hidden_states, lengths) + + def _decode_impl( + self, + audio_codes: torch.Tensor, + target_lens: torch.Tensor | None = None, + ) -> torch.Tensor: """Decode audio tokens to waveform. Args: @@ -172,7 +218,12 @@ def _decode_impl(self, audio_codes: torch.Tensor): quantized = self.fc2(quantized.transpose(1, 2).to(self.fc2.weight.dtype)).transpose(1, 2).float() # Acoustic decoder: [B, 256, T] → [B, 1, T*960] - audio = self.acoustic_decoder(quantized) + if target_lens is None or not all( + hasattr(self.acoustic_decoder, name) for name in ("conv1", "block", "snake1", "conv2", "tanh") + ): + audio = self.acoustic_decoder(quantized) + else: + audio = self._decode_acoustic_length_aware(quantized, target_lens) # Ensure [B, 1, samples] if audio.dim() == 2: @@ -181,7 +232,11 @@ def _decode_impl(self, audio_codes: torch.Tensor): return audio @torch.inference_mode() - def forward(self, audio_codes: torch.Tensor) -> torch.Tensor: + def forward( + self, + audio_codes: torch.Tensor, + target_lens: list[int] | torch.Tensor | None = None, + ) -> torch.Tensor: """Decode audio tokens to waveform. Args: @@ -194,11 +249,30 @@ def forward(self, audio_codes: torch.Tensor) -> torch.Tensor: raise RuntimeError("Decoder not loaded. Call load_weights() first.") device = audio_codes.device - - if self.graphs.captured: - return self.graphs.forward(audio_codes).to(device) + if target_lens is None: + lengths = torch.full( + (audio_codes.shape[0],), + audio_codes.shape[-1], + dtype=torch.long, + device=device, + ) + elif isinstance(target_lens, torch.Tensor): + lengths = target_lens.to(device=device, dtype=torch.long) + else: + lengths = torch.tensor(target_lens, device=device, dtype=torch.long) + + if lengths.shape != (audio_codes.shape[0],): + raise ValueError( + f"Expected one target length per request, got shape {tuple(lengths.shape)} " + f"for batch size {audio_codes.shape[0]}." + ) + if torch.any(lengths <= 0) or torch.any(lengths > audio_codes.shape[-1]): + raise ValueError(f"Target lengths must be in [1, {audio_codes.shape[-1]}], got {lengths.tolist()}.") + + if self.graphs is not None and self.graphs.captured: + return self.graphs.forward(audio_codes, lengths).to(device) else: - return self._decode_impl(audio_codes).to(device) + return self._decode_impl(audio_codes, lengths).to(device) def _adjust_output_padding(self, decoder: nn.Module): """Adjust ConvTranspose1d output_padding (HiggsAudioV2 modification).""" @@ -290,8 +364,6 @@ def load_weights(self, model_dir: str, device: torch.device) -> None: self.graphs = DecoderGraph( self._decode_impl, - batch_sizes=(1, 2, 4, 8, 16), - bucket_sizes=(64, 128, 256), device=device, ) self.graphs.warmup() From 042ea81b81f1fd6bb579d07ce101b1d7efb04909 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 20 Aug 2026 18:04:02 +0800 Subject: [PATCH 03/17] support batched step-execution Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../models/omnivoice/test_decoder.py | 1 - .../models/omnivoice/pipeline_omnivoice.py | 698 +++++++----------- .../models/omnivoice/omnivoice_decoder.py | 137 +--- 3 files changed, 301 insertions(+), 535 deletions(-) diff --git a/tests/model_executor/models/omnivoice/test_decoder.py b/tests/model_executor/models/omnivoice/test_decoder.py index 040fde89612..acfb8a4a062 100644 --- a/tests/model_executor/models/omnivoice/test_decoder.py +++ b/tests/model_executor/models/omnivoice/test_decoder.py @@ -245,7 +245,6 @@ def _build_padded_decoder() -> OmniVoiceDecoder: decoder._adjust_output_padding(decoder.acoustic_decoder) decoder.acoustic_decoder.tanh = nn.Identity() - decoder.graphs = None decoder._loaded = True return decoder diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 719c2332d47..7dbe61e9ba7 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -17,6 +17,7 @@ import random import re from collections.abc import Iterable, Sequence +from dataclasses import dataclass from typing import Any, ClassVar import numpy as np @@ -52,6 +53,15 @@ logger = init_logger(__name__) +@dataclass +class _PreparedOmniVoiceRequest: + input_ids: torch.Tensor + audio_mask: torch.Tensor + attention_mask: torch.Tensor + target_len: int + seed: int | None + + def get_omnivoice_post_process_func(od_config: OmniDiffusionConfig): """Post-processing: convert audio tensor to numpy for WAV encoding.""" @@ -216,184 +226,226 @@ def _encode_ref_audio(self, audio_signal: torch.Tensor, sr: int) -> torch.Tensor tokens = tokens.squeeze(0) # [8, T_ref] return tokens - def prepare_encode(self, states: list[StepRequestState]) -> DiffusionRequestBatch: + def _prepare_request_input( + self, + prompt: Any, + extra: dict[str, Any], + ) -> _PreparedOmniVoiceRequest | DiffusionOutput: + """Build one request's conditional/unconditional model inputs.""" ref_audio = None ref_text = None lang = "None" instruct = "None" voice_name = None - device = self.device - num_cb = self.config.num_audio_codebook - mask_id = self.config.audio_mask_id - batch_target_len: list[int] = [] - batch_input_ids: list[torch.Tensor] = [] - batch_audio_mask: list[torch.Tensor] = [] - batch_attn_mask: list[torch.Tensor] = [] - if isinstance(states, StepRequestState): - states = [states] - for state in states: - prompt = state.prompt if state.prompt else "" - extra = state.sampling.extra_args or {} - seed = extra.get("seed", None) - if isinstance(prompt, dict): - # Top-level keys (used by serving_speech.py /v1/audio/speech path) - text = prompt.get("input") or prompt.get("text") or prompt.get("prompt") - ref_audio = prompt.get("ref_audio") - ref_text = prompt.get("ref_text") - voice_name = prompt.get("voice_name") - lang = prompt.get("lang") - instruct = prompt.get("instruct") - # OmniTextPrompt format (used by offline Omni.generate path): - # ref_audio comes via multi_modal_data["audio"] and the rest via - # mm_processor_kwargs. Fall back to those when top-level keys are - # absent so both invocation styles work. - mm_data = prompt.get("multi_modal_data") or {} - mm_kwargs = prompt.get("mm_processor_kwargs") or {} - if ref_audio is None: - audio_field = mm_data.get("audio") - # Standard multimodal shape allows a list of audios; OmniVoice - # voice cloning conditions on a single reference clip, so - # unwrap a length-1 list and reject multi-reference prompts up - # front (otherwise a list would later crash inside - # ``_encode_ref_audio`` when it calls ``audio.dim()``). - if isinstance(audio_field, list): - if len(audio_field) == 1: - audio_field = audio_field[0] - elif len(audio_field) > 1: - return DiffusionOutput( - error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" # noqa: E501 - ) - else: - audio_field = None - if audio_field is not None: - if isinstance(audio_field, tuple) and len(audio_field) == 2: - ref_audio = audio_field - else: - sr = mm_kwargs.get("sample_rate") or self.sample_rate - ref_audio = (audio_field, int(sr)) - if ref_text is None: - ref_text = mm_kwargs.get("ref_text") - if lang is None: - lang = mm_kwargs.get("lang") - if instruct is None: - instruct = mm_kwargs.get("instruct") - - if not text: - return DiffusionOutput(error="Empty text prompt") - lang = lang or "None" - instruct = instruct or "None" - else: - text = str(prompt) - if not text: - return DiffusionOutput(error="Empty text prompt") - - # Estimate target duration - target_len = self.duration_estimator.estimate_duration(text, "Nice to meet you.", 25) - target_len = max(1, int(target_len)) - batch_target_len.append(target_len) - - # Build text prompt with control tokens - style_text = f"<|denoise|><|lang_start|>{lang}<|lang_end|><|instruct_start|>{instruct}<|instruct_end|>" - full_text = _combine_text(ref_text=ref_text, text=text) - wrapped_text = f"<|text_start|>{full_text}<|text_end|>" - style_tokens = self.tokenizer.encode(style_text).ids - text_tokens = _tokenize_with_nonverbal_tags(wrapped_text, self.tokenizer) - encoding_ids = style_tokens + text_tokens - text_tokens = torch.tensor(encoding_ids, dtype=torch.long, device=device) - text_len = text_tokens.shape[0] - - # Encode reference audio tokens if provided (with voice caching) - ref_audio_tokens = None - if ref_audio is not None: - if self.audio_tokenizer is None: - raise RuntimeError( - "Voice cloning requires transformers>=5.3.0. Try: uv pip install 'transformers>=5.3.0'" - ) - # Check speaker cache first - _cache_key = None - if voice_name: - _cache_key = self._speaker_cache.make_cache_key( - voice_name, - model_type="omnivoice", - created_at=int(prompt.get("voice_created_at") or 0), - ) - cached = self._speaker_cache.get(_cache_key) - if cached is not None: - ref_audio_tokens = cached["ref_audio_tokens"].to(device) - _cache_key = None # hit → don't store again - logger.debug("Speaker cache HIT for OmniVoice speaker '%s'", voice_name) - - if ref_audio_tokens is None: - audio_signal, sr = ref_audio - if isinstance(audio_signal, np.ndarray): - audio_signal = torch.from_numpy(audio_signal).float() - ref_audio_tokens = self._encode_ref_audio(audio_signal, int(sr)).to(device) - - # Store in cache for next request - if _cache_key is not None: - self._speaker_cache.put(_cache_key, {"ref_audio_tokens": ref_audio_tokens.cpu()}) - logger.debug("Speaker cache STORE for OmniVoice speaker '%s'", voice_name) - - # Build conditional + unconditional batches [2, 8, max_len] - text_ids = text_tokens.unsqueeze(0).repeat(num_cb, 1) - target_ids = torch.full((num_cb, target_len), mask_id, dtype=torch.long, device=device) - - if ref_audio_tokens is not None: - cond_ids = torch.cat([text_ids, ref_audio_tokens, target_ids], dim=1) - else: - cond_ids = torch.cat([text_ids, target_ids], dim=1) - - cond_len = cond_ids.shape[1] - uncond_ids = target_ids.clone() - uncond_len = target_len - max_len = max(cond_len, uncond_len) - if uncond_len < max_len: - pad = torch.full( - (num_cb, max_len - uncond_len), - mask_id, - dtype=torch.long, - device=device, - ) - uncond_ids = torch.cat([uncond_ids, pad], dim=1) - input_ids = torch.stack([cond_ids, uncond_ids]) - batch_input_ids.append(input_ids) - - audio_mask = torch.zeros(2, max_len, dtype=torch.bool, device=device) - audio_mask[0, text_len:cond_len] = True - audio_mask[1, :uncond_len] = True - batch_audio_mask.append(audio_mask) - - attn_mask = torch.zeros(2, 1, max_len, max_len, dtype=torch.bool, device=device) - attn_mask[0, :, :cond_len, :cond_len] = True - attn_mask[1, :, :uncond_len, :uncond_len] = True - batch_attn_mask.append(attn_mask) - - if len(batch_input_ids) > 1: - max_input_ids_len = max([ids.shape[-1] for ids in batch_input_ids]) - num_padded_len = [max_input_ids_len - ids.shape[-1] for ids in batch_input_ids] - for i, ids in enumerate(batch_input_ids): - batch_input_ids[i] = torch.cat( - [ - ids, - torch.full( - (ids.shape[0], ids.shape[1], num_padded_len[i]), mask_id, dtype=torch.long, device=device - ), - ], - dim=-1, + seed = extra.get("seed", None) + + if isinstance(prompt, dict): + text = prompt.get("input") or prompt.get("text") or prompt.get("prompt") + ref_audio = prompt.get("ref_audio") + ref_text = prompt.get("ref_text") + voice_name = prompt.get("voice_name") + lang = prompt.get("lang") + instruct = prompt.get("instruct") + mm_data = prompt.get("multi_modal_data") or {} + mm_kwargs = prompt.get("mm_processor_kwargs") or {} + if ref_audio is None: + audio_field = mm_data.get("audio") + if isinstance(audio_field, list): + if len(audio_field) == 1: + audio_field = audio_field[0] + elif len(audio_field) > 1: + return DiffusionOutput( + error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" + ) + else: + audio_field = None + if audio_field is not None: + if isinstance(audio_field, tuple) and len(audio_field) == 2: + ref_audio = audio_field + else: + sr = mm_kwargs.get("sample_rate") or self.sample_rate + ref_audio = (audio_field, int(sr)) + if ref_text is None: + ref_text = mm_kwargs.get("ref_text") + if lang is None: + lang = mm_kwargs.get("lang") + if instruct is None: + instruct = mm_kwargs.get("instruct") + if not text: + return DiffusionOutput(error="Empty text prompt") + lang = lang or "None" + instruct = instruct or "None" + else: + text = str(prompt) + if not text: + return DiffusionOutput(error="Empty text prompt") + + target_len = self.duration_estimator.estimate_duration(text, "Nice to meet you.", 25) + target_len = max(1, int(target_len)) + + style_text = f"<|denoise|><|lang_start|>{lang}<|lang_end|><|instruct_start|>{instruct}<|instruct_end|>" + full_text = _combine_text(ref_text=ref_text, text=text) + wrapped_text = f"<|text_start|>{full_text}<|text_end|>" + style_tokens = self.tokenizer.encode(style_text).ids + text_tokens = _tokenize_with_nonverbal_tags(wrapped_text, self.tokenizer) + encoding_ids = style_tokens + text_tokens + text_tokens_tensor = torch.tensor(encoding_ids, dtype=torch.long, device=self.device) + text_len = text_tokens_tensor.shape[0] + + ref_audio_tokens = None + if ref_audio is not None: + if self.audio_tokenizer is None: + raise RuntimeError( + "Voice cloning requires transformers>=5.3.0. Try: uv pip install 'transformers>=5.3.0'" ) - for i, mask in enumerate(batch_audio_mask): - batch_audio_mask[i] = torch.cat( - [mask, torch.full((mask.shape[0], num_padded_len[i]), False, dtype=torch.bool, device=device)], - dim=-1, + cache_key = None + if voice_name: + cache_key = self._speaker_cache.make_cache_key( + voice_name, + model_type="omnivoice", + created_at=int(prompt.get("voice_created_at") or 0), ) - for i, mask in enumerate(batch_attn_mask): - n = num_padded_len[i] - batch_attn_mask[i] = F.pad(mask, (0, n, 0, n)) + cached = self._speaker_cache.get(cache_key) + if cached is not None: + ref_audio_tokens = cached["ref_audio_tokens"].to(self.device) + cache_key = None + logger.debug("Speaker cache HIT for OmniVoice speaker '%s'", voice_name) + + if ref_audio_tokens is None: + audio_signal, sr = ref_audio + if isinstance(audio_signal, np.ndarray): + audio_signal = torch.from_numpy(audio_signal).float() + ref_audio_tokens = self._encode_ref_audio(audio_signal, int(sr)).to(self.device) + if cache_key is not None: + self._speaker_cache.put(cache_key, {"ref_audio_tokens": ref_audio_tokens.cpu()}) + logger.debug("Speaker cache STORE for OmniVoice speaker '%s'", voice_name) - target_lens = batch_target_len - B = len(target_lens) - device = input_ids.device - max_target_len = max(target_lens) + num_cb = self.config.num_audio_codebook + mask_id = self.config.audio_mask_id + text_ids = text_tokens_tensor.unsqueeze(0).repeat(num_cb, 1) + target_ids = torch.full((num_cb, target_len), mask_id, dtype=torch.long, device=self.device) + cond_ids = ( + torch.cat([text_ids, ref_audio_tokens, target_ids], dim=1) + if ref_audio_tokens is not None + else torch.cat([text_ids, target_ids], dim=1) + ) + cond_len = cond_ids.shape[1] + uncond_ids = target_ids.clone() + uncond_len = target_len + max_len = max(cond_len, uncond_len) + if uncond_len < max_len: + pad = torch.full( + (num_cb, max_len - uncond_len), + mask_id, + dtype=torch.long, + device=self.device, + ) + uncond_ids = torch.cat([uncond_ids, pad], dim=1) + + input_ids = torch.stack([cond_ids, uncond_ids]) + audio_mask = torch.zeros(2, max_len, dtype=torch.bool, device=self.device) + audio_mask[0, text_len:cond_len] = True + audio_mask[1, :uncond_len] = True + attention_mask = torch.zeros(2, 1, max_len, max_len, dtype=torch.bool, device=self.device) + attention_mask[0, :, :cond_len, :cond_len] = True + attention_mask[1, :, :uncond_len, :uncond_len] = True + + return _PreparedOmniVoiceRequest( + input_ids=input_ids, + audio_mask=audio_mask, + attention_mask=attention_mask, + target_len=target_len, + seed=seed, + ) + + def _collate_request_inputs( + self, + prepared_requests: Sequence[_PreparedOmniVoiceRequest], + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Right-pad requests and arrange rows as all cond followed by all uncond.""" + max_len = max(request.input_ids.shape[-1] for request in prepared_requests) + input_pairs: list[torch.Tensor] = [] + audio_mask_pairs: list[torch.Tensor] = [] + attention_mask_pairs: list[torch.Tensor] = [] + mask_id = self.config.audio_mask_id + + for request in prepared_requests: + pad_len = max_len - request.input_ids.shape[-1] + input_ids = request.input_ids + audio_mask = request.audio_mask + attention_mask = request.attention_mask + if pad_len: + input_ids = F.pad(input_ids, (0, pad_len), value=mask_id) + audio_mask = F.pad(audio_mask, (0, pad_len), value=False) + attention_mask = F.pad(attention_mask, (0, pad_len, 0, pad_len), value=False) + input_pairs.append(input_ids) + audio_mask_pairs.append(audio_mask) + attention_mask_pairs.append(attention_mask) + + input_pairs_tensor = torch.stack(input_pairs, dim=0) + audio_mask_pairs_tensor = torch.stack(audio_mask_pairs, dim=0) + attention_mask_pairs_tensor = torch.stack(attention_mask_pairs, dim=0) + return ( + torch.cat([input_pairs_tensor[:, 0], input_pairs_tensor[:, 1]], dim=0), + torch.cat([audio_mask_pairs_tensor[:, 0], audio_mask_pairs_tensor[:, 1]], dim=0), + torch.cat([attention_mask_pairs_tensor[:, 0], attention_mask_pairs_tensor[:, 1]], dim=0), + ) + + @staticmethod + def _request_major_to_cfg_major(x: torch.Tensor, batch_size: int) -> torch.Tensor: + """Convert [cond0, uncond0, ...] to [cond0, ..., uncond0, ...].""" + if x.shape[0] != 2 * batch_size: + raise ValueError(f"Expected {2 * batch_size} CFG rows for batch size {batch_size}, got {x.shape[0]}.") + pairs = x.reshape(batch_size, 2, *x.shape[1:]) + return torch.cat([pairs[:, 0], pairs[:, 1]], dim=0) + + @staticmethod + def _cfg_major_to_request_major(x: torch.Tensor, batch_size: int) -> torch.Tensor: + """Convert [cond0, ..., uncond0, ...] to [cond0, uncond0, ...].""" + if x.shape[0] != 2 * batch_size: + raise ValueError(f"Expected {2 * batch_size} CFG rows for batch size {batch_size}, got {x.shape[0]}.") + pairs = torch.stack([x[:batch_size], x[batch_size:]], dim=1) + return pairs.reshape(2 * batch_size, *x.shape[1:]) + + def prepare_state_batch(self, states: list[StepRequestState], new_request_ids: list[str]) -> list[StepRequestState]: + if not new_request_ids: + return states + else: + max_target_len = max([state.latents.shape[-1] for state in states]) + if states[0].extra.get("max_target_len", None) == max_target_len: + states_to_repad = [state for state in states if state.request_id in new_request_ids] + self.repadding_state(states_to_repad, max_target_len) + return states + + self.repadding_state(states, max_target_len) + return states + + def repadding_state(self, states: list[StepRequestState], max_new_target_len: int) -> list[StepRequestState]: + for state in states: + num_pads = max_new_target_len - state.latents.shape[-1] + state.latents = F.pad( + state.latents, + (0, num_pads), + value=self.config.audio_mask_id, + ) + state.extra["audio_mask"] = F.pad(state.extra["audio_mask"], (0, num_pads), value=False) + state.extra["attn_mask"] = F.pad(state.extra["attn_mask"], (0, num_pads, 0, num_pads), value=False) + return states + + def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: + prompt = state.prompt if state.prompt else "" + extra = state.sampling.extra_args or {} + prepared = self._prepare_request_input(prompt, extra) + if isinstance(prepared, DiffusionOutput): + return prepared + prepared_request = prepared + + target_lens = prepared_request.target_len + input_ids = prepared_request.input_ids + audio_mask = prepared_request.audio_mask + attn_mask = prepared_request.attention_mask + seed = prepared_request.seed + device = self.device mask_id = self.config.audio_mask_id num_codebooks = self.config.num_audio_codebook if seed is None: @@ -402,71 +454,50 @@ def prepare_encode(self, states: list[StepRequestState]) -> DiffusionRequestBatc t_shift = self.t_shift # Initialize all target tokens as [MASK] - positions = torch.arange(max_target_len, device=device).unsqueeze(0) - valid_target_mask = positions < torch.tensor(target_lens, device=device).unsqueeze(1) - tokens = torch.zeros((B, num_codebooks, max_target_len), dtype=torch.long, device=device) + positions = torch.arange(target_lens, device=device).unsqueeze(0) + valid_target_mask = positions < torch.tensor(target_lens, device=device) + tokens = torch.zeros((1, num_codebooks, target_lens), dtype=torch.long, device=device) tokens.masked_fill_(valid_target_mask.unsqueeze(1), mask_id) timesteps = _get_time_steps(0.0, 1.0, num_step + 1, t_shift) # Compute unmasking schedule schedules = [] - for t_len in target_lens: - total_mask = t_len * num_codebooks - rem = total_mask - sched = [] - for step in range(num_step): - num = ( - rem - if step == num_step - 1 - else min( - math.ceil(total_mask * (timesteps[step + 1] - timesteps[step])), - rem, - ) + total_mask = target_lens * num_codebooks + rem = total_mask + sched = [] + for step in range(num_step): + num = ( + rem + if step == num_step - 1 + else min( + math.ceil(total_mask * (timesteps[step + 1] - timesteps[step])), + rem, ) - sched.append(int(num)) - rem -= int(num) - schedules.append(sched) - schedules = torch.tensor(schedules, dtype=torch.long, device=device) + ) + sched.append(int(num)) + rem -= int(num) + schedules = torch.tensor(sched, dtype=torch.long, device=device) layer_ids = torch.arange(num_codebooks, device=device).view(1, -1, 1) generator = torch.Generator(device=device).manual_seed(seed) - for i in range(B): - states[i].latents = batch_input_ids[i] - states[i].timesteps = schedules[i] - states[i].guidance = self.guidance_scale - states[i].extra["schedules"] = schedules - states[i].extra["layer_ids"] = layer_ids - states[i].extra["generator"] = generator - states[i].extra["t_shift"] = t_shift - states[i].extra["target_len"] = target_lens[i] - states[i].extra["audio_mask"] = batch_audio_mask[i] - states[i].extra["attn_mask"] = batch_attn_mask[i] - states[i].extra["tokens"] = tokens[i] - - use_cuda_graph = self.generator._cuda_graph_fwd is not None and input_ids.is_cuda - if not use_cuda_graph: - # Eager-path-only constants (the cuda-graph captures its own). - text_embeds_cached = self.text_embedding(input_ids[:, 0, :]) - audio_mask_3d = audio_mask.unsqueeze(-1) - self._ensure_rope(input_ids.shape[-1], device) - target_dtype = text_embeds_cached.dtype - cos = self._rope_cos.to(device=device, dtype=target_dtype) - sin = self._rope_sin.to(device=device, dtype=target_dtype) - - for i in range(B): - states[i].extra["text_embeds"] = text_embeds_cached - states[i].extra["audio_mask_3d"] = audio_mask_3d - states[i].extra["cos"] = cos - states[i].extra["sin"] = sin + state.latents = input_ids + state.timesteps = schedules + state.guidance = self.guidance_scale + state.extra["schedules"] = schedules + state.extra["layer_ids"] = layer_ids + state.extra["generator"] = generator + state.extra["t_shift"] = t_shift + state.extra["target_len"] = target_lens + state.extra["audio_mask"] = audio_mask + state.extra["attn_mask"] = attn_mask + state.extra["tokens"] = tokens def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestState] | None = None, **kwargs: Any): use_cuda_graph = self.generator._cuda_graph_fwd is not None input_ids = input_batch.latents layer_ids = states[0].extra["layer_ids"] - generator = states[0].extra["generator"] - schedules = states[0].extra["schedules"] batch_audio_mask: list[torch.Tesor] = [] batch_attn_mask: list[torch.Tensor] = [] @@ -474,6 +505,8 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS batch_tokens: list[torch.Tensor] = [] steps: list[int] = [] + schedules: list[torch.Tensor] = [] + generators: list[torch.Generator] = [] for state in states: batch_audio_mask.append(state.extra.get("audio_mask", None)) @@ -481,24 +514,25 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS guidance_scale = state.extra.get("guidance", self.guidance_scale) batch_target_len.append(state.extra["target_len"]) batch_tokens.append(state.extra["tokens"]) + generators.append(state.extra.get("generator", None)) + schedules.append(state.extra["schedules"]) steps.append(state.step_index) - batch_tokens = torch.stack(batch_tokens, dim=0) - if batch_tokens.dim() == 4: - batch_tokens = batch_tokens.squeeze(0) if len(batch_audio_mask) == 1: - audio_mask = batch_audio_mask[0] + batch_audio_mask = batch_audio_mask[0] batch_attn_mask = batch_attn_mask[0] else: - audio_mask = torch.stack(batch_audio_mask, dim=0) + batch_audio_mask = torch.stack(batch_audio_mask, dim=0) batch_attn_mask = torch.stack(batch_attn_mask, dim=0) + if len(batch_target_len) > 1: + batch_audio_mask = batch_audio_mask.reshape(-1, *batch_audio_mask.shape[2:]) + batch_attn_mask = batch_attn_mask.reshape(-1, *batch_attn_mask.shape[2:]) + B = len(batch_target_len) target_lens = batch_target_len - - text_embeds_cached = states[0].extra.get("text_embeds", None) - audio_mask_3d = states[0].extra.get("audio_mask_3d", None) - cos = states[0].extra.get("cos", None) - sin = states[0].extra.get("sin", None) + input_ids = self._request_major_to_cfg_major(input_ids, B) + batch_audio_mask = self._request_major_to_cfg_major(batch_audio_mask, B) + batch_attn_mask = self._request_major_to_cfg_major(batch_attn_mask, B) mask_id = self.config.audio_mask_id position_temperature = self.position_temperature @@ -513,13 +547,12 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS ) if use_cuda_graph: # Float mask skips per-layer conversion; fp32 cast deferred to the per-item slices below. - batch_logits = self.generator._cuda_graph_fwd(input_ids, audio_mask, sdpa_attn_mask) + batch_logits = self.generator._cuda_graph_fwd(input_ids, batch_audio_mask, sdpa_attn_mask) else: - # Eager fallback reuses hoisted constants (text embeds, sdpa mask, cos/sin). - inputs_embeds = self.generator._prepare_embeddings( - input_ids, audio_mask, text_embeds=text_embeds_cached, audio_mask_3d=audio_mask_3d - ) - hidden_states = self.generator._transformer_forward(inputs_embeds, sdpa_attn_mask, cos=cos, sin=sin) + # Recompute embeddings and RoPE from the current dynamically + # padded/reordered batch. + inputs_embeds = self.generator._prepare_embeddings(input_ids, batch_audio_mask) + hidden_states = self.generator._transformer_forward(inputs_embeds, sdpa_attn_mask) # fp32 cast deferred to the per-item slices below. batch_logits = self.generator._get_logits(hidden_states) # batch_logits: [2*B, 8, S, 1025] @@ -553,7 +586,7 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS # Token prediction if class_temperature > 0.0: - pred_tokens = _gumbel_sample(log_probs, class_temperature, generator).argmax(dim=-1) + pred_tokens = _gumbel_sample(log_probs, class_temperature, generators[i]).argmax(dim=-1) else: pred_tokens = log_probs.argmax(dim=-1) # [1, 8, T] @@ -565,10 +598,11 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS # Gumbel noise for position selection if position_temperature > 0.0: - scores = _gumbel_sample(scores, position_temperature, generator) + scores = _gumbel_sample(scores, position_temperature, generators[i]) # Mask out already unmasked positions - sample_tokens = batch_tokens[i : i + 1, :, :t_len] + sample = batch_tokens[i] + sample_tokens = sample[..., :t_len] scores.masked_fill_(sample_tokens != mask_id, -float("inf")) # Select top-k positions to unmask. .flatten() on this non-contiguous view already copies. @@ -582,14 +616,14 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS input_ids[i, :, c_len - t_len : c_len] = sample_tokens.squeeze(0) input_ids[B + i, :, :t_len] = sample_tokens.squeeze(0) - return input_ids + return self._cfg_major_to_request_major(input_ids, B) def step_scheduler(self, state: StepRequestState, noise_pred: torch.Tensor, **kwargs: Any): state.latents = noise_pred state.step_index += 1 def post_decode(self, state: StepRequestState, **kwargs: Any): - tokens = state.extra.get("tokens", None) + tokens = state.extra["tokens"] if tokens.dim() == 2: tokens = tokens.unsqueeze(0) audio = self.decoder(tokens) @@ -603,185 +637,19 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: {"text": "...", "ref_audio": (samples, sr), "ref_text": "...", "lang": "...", "instruct": "..."} """ - ref_audio = None - ref_text = None - lang = "None" - instruct = "None" - voice_name = None outputs: list[DiffusionOutput] = [] - device = self.device - num_cb = self.config.num_audio_codebook - mask_id = self.config.audio_mask_id - batch_target_len: list[int] = [] - batch_input_ids: list[torch.Tensor] = [] - batch_audio_mask: list[torch.Tensor] = [] - batch_attn_mask: list[torch.Tensor] = [] + prepared_requests: list[_PreparedOmniVoiceRequest] = [] for request in req.requests: prompt = request.prompt if request.prompt else "" extra = request.sampling_params.extra_args or {} - seed = extra.get("seed", None) - if isinstance(prompt, dict): - # Top-level keys (used by serving_speech.py /v1/audio/speech path) - text = prompt.get("input") or prompt.get("text") or prompt.get("prompt") - ref_audio = prompt.get("ref_audio") - ref_text = prompt.get("ref_text") - voice_name = prompt.get("voice_name") - lang = prompt.get("lang") - instruct = prompt.get("instruct") - # OmniTextPrompt format (used by offline Omni.generate path): - # ref_audio comes via multi_modal_data["audio"] and the rest via - # mm_processor_kwargs. Fall back to those when top-level keys are - # absent so both invocation styles work. - mm_data = prompt.get("multi_modal_data") or {} - mm_kwargs = prompt.get("mm_processor_kwargs") or {} - if ref_audio is None: - audio_field = mm_data.get("audio") - # Standard multimodal shape allows a list of audios; OmniVoice - # voice cloning conditions on a single reference clip, so - # unwrap a length-1 list and reject multi-reference prompts up - # front (otherwise a list would later crash inside - # ``_encode_ref_audio`` when it calls ``audio.dim()``). - if isinstance(audio_field, list): - if len(audio_field) == 1: - audio_field = audio_field[0] - elif len(audio_field) > 1: - return DiffusionOutput( - error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" # noqa: E501 - ) - else: - audio_field = None - if audio_field is not None: - if isinstance(audio_field, tuple) and len(audio_field) == 2: - ref_audio = audio_field - else: - sr = mm_kwargs.get("sample_rate") or self.sample_rate - ref_audio = (audio_field, int(sr)) - if ref_text is None: - ref_text = mm_kwargs.get("ref_text") - if lang is None: - lang = mm_kwargs.get("lang") - if instruct is None: - instruct = mm_kwargs.get("instruct") - - if not text: - return DiffusionOutput(error="Empty text prompt") - lang = lang or "None" - instruct = instruct or "None" - else: - text = str(prompt) - if not text: - return DiffusionOutput(error="Empty text prompt") - - # Estimate target duration - target_len = self.duration_estimator.estimate_duration(text, "Nice to meet you.", 25) - target_len = max(1, int(target_len)) - batch_target_len.append(target_len) - - # Build text prompt with control tokens - style_text = f"<|denoise|><|lang_start|>{lang}<|lang_end|><|instruct_start|>{instruct}<|instruct_end|>" - full_text = _combine_text(ref_text=ref_text, text=text) - wrapped_text = f"<|text_start|>{full_text}<|text_end|>" - style_tokens = self.tokenizer.encode(style_text).ids - text_tokens = _tokenize_with_nonverbal_tags(wrapped_text, self.tokenizer) - encoding_ids = style_tokens + text_tokens - text_tokens = torch.tensor(encoding_ids, dtype=torch.long, device=device) - text_len = text_tokens.shape[0] - - # Encode reference audio tokens if provided (with voice caching) - ref_audio_tokens = None - if ref_audio is not None: - if self.audio_tokenizer is None: - raise RuntimeError( - "Voice cloning requires transformers>=5.3.0. Try: uv pip install 'transformers>=5.3.0'" - ) - # Check speaker cache first - _cache_key = None - if voice_name: - _cache_key = self._speaker_cache.make_cache_key( - voice_name, - model_type="omnivoice", - created_at=int(prompt.get("voice_created_at") or 0), - ) - cached = self._speaker_cache.get(_cache_key) - if cached is not None: - ref_audio_tokens = cached["ref_audio_tokens"].to(device) - _cache_key = None # hit → don't store again - logger.debug("Speaker cache HIT for OmniVoice speaker '%s'", voice_name) - - if ref_audio_tokens is None: - audio_signal, sr = ref_audio - if isinstance(audio_signal, np.ndarray): - audio_signal = torch.from_numpy(audio_signal).float() - ref_audio_tokens = self._encode_ref_audio(audio_signal, int(sr)).to(device) - - # Store in cache for next request - if _cache_key is not None: - self._speaker_cache.put(_cache_key, {"ref_audio_tokens": ref_audio_tokens.cpu()}) - logger.debug("Speaker cache STORE for OmniVoice speaker '%s'", voice_name) - - # Build conditional + unconditional batches [2, 8, max_len] - text_ids = text_tokens.unsqueeze(0).repeat(num_cb, 1) - target_ids = torch.full((num_cb, target_len), mask_id, dtype=torch.long, device=device) - - if ref_audio_tokens is not None: - cond_ids = torch.cat([text_ids, ref_audio_tokens, target_ids], dim=1) - else: - cond_ids = torch.cat([text_ids, target_ids], dim=1) - cond_len = cond_ids.shape[1] - uncond_ids = target_ids.clone() - uncond_len = target_len - max_len = max(cond_len, uncond_len) - if uncond_len < max_len: - pad = torch.full( - (num_cb, max_len - uncond_len), - mask_id, - dtype=torch.long, - device=device, - ) - uncond_ids = torch.cat([uncond_ids, pad], dim=1) - input_ids = torch.stack([cond_ids, uncond_ids]) - batch_input_ids.append(input_ids) - - audio_mask = torch.zeros(2, max_len, dtype=torch.bool, device=device) - audio_mask[0, text_len:cond_len] = True - audio_mask[1, :uncond_len] = True - batch_audio_mask.append(audio_mask) - - attn_mask = torch.zeros(2, 1, max_len, max_len, dtype=torch.bool, device=device) - attn_mask[0, :, :cond_len, :cond_len] = True - attn_mask[1, :, :uncond_len, :uncond_len] = True - batch_attn_mask.append(attn_mask) - if len(batch_input_ids) > 1: - max_input_ids_len = max([ids.shape[-1] for ids in batch_input_ids]) - num_padded_len = [max_input_ids_len - ids.shape[-1] for ids in batch_input_ids] - for i, ids in enumerate(batch_input_ids): - batch_input_ids[i] = torch.cat( - [ - ids, - torch.full( - (ids.shape[0], ids.shape[1], num_padded_len[i]), mask_id, dtype=torch.long, device=device - ), - ], - dim=-1, - ) - for i, mask in enumerate(batch_audio_mask): - batch_audio_mask[i] = torch.cat( - [mask, torch.full((mask.shape[0], num_padded_len[i]), False, dtype=torch.bool, device=device)], - dim=-1, - ) - for i, mask in enumerate(batch_attn_mask): - n = num_padded_len[i] - batch_attn_mask[i] = F.pad(mask, (0, n, 0, n)) - # Each per-request tensor is ordered as [cond_i, uncond_i]. The - # generator expects the full batch to be ordered as - # [cond_0, ..., cond_B-1, uncond_0, ..., uncond_B-1] so that request i - # pairs logits i and B+i for classifier-free guidance. - input_id_pairs = torch.stack(batch_input_ids, dim=0) - audio_mask_pairs = torch.stack(batch_audio_mask, dim=0) - attn_mask_pairs = torch.stack(batch_attn_mask, dim=0) - batch_input_ids = torch.cat([input_id_pairs[:, 0], input_id_pairs[:, 1]], dim=0) - batch_audio_mask = torch.cat([audio_mask_pairs[:, 0], audio_mask_pairs[:, 1]], dim=0) - batch_attn_mask = torch.cat([attn_mask_pairs[:, 0], attn_mask_pairs[:, 1]], dim=0) + prepared = self._prepare_request_input(prompt, extra) + if isinstance(prepared, DiffusionOutput): + return prepared + prepared_requests.append(prepared) + + batch_target_len = [request.target_len for request in prepared_requests] + seed = prepared_requests[-1].seed + batch_input_ids, batch_audio_mask, batch_attn_mask = self._collate_request_inputs(prepared_requests) # Run 32-step iterative unmasking tokens = self.generator( input_ids=batch_input_ids, diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py index b262cb46082..d40a67be3f4 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py @@ -21,9 +21,7 @@ import torch import torch.nn as nn -from torch.cuda.graphs import CUDAGraph from vllm.logger import init_logger -from vllm.platforms import current_platform from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig @@ -70,74 +68,6 @@ def decode(self, codes: torch.Tensor) -> torch.Tensor: return result -class DecoderGraph: - def __init__( - self, - decode_fn, - batch_sizes: tuple[int] = (1, 2), - bucket_sizes: tuple[int] = (64, 128, 256), - device: torch.device = torch.device("cuda"), - ): - self.graphs: dict[tuple[int, int], CUDAGraph] = {} - self.capture_batch_size = batch_sizes - self.device = device - self.capture_bucket_size = bucket_sizes - self.decode_fn = decode_fn - self.static_inputs: dict[tuple[int, int], torch.Tensor] = {} - self.static_lengths: dict[tuple[int, int], torch.Tensor] = {} - self.static_outputs: dict[tuple[int, int], torch.Tensor] = {} - self.captured = False - - def warmup(self): - for batch_size in self.capture_batch_size: - for bucket_size in self.capture_bucket_size: - self._capture(batch_size, bucket_size) - self.captured = bool(self.graphs) - - def _capture(self, batch_size: int, bucket_size: int): - if torch.cuda.is_current_stream_capturing() or self.device.type != "cuda": - return - static_inputs = torch.zeros((batch_size, 8, bucket_size), device=self.device, dtype=torch.long) - static_lengths = torch.full((batch_size,), bucket_size, device=self.device, dtype=torch.long) - self.static_inputs[(batch_size, bucket_size)] = static_inputs - self.static_lengths[(batch_size, bucket_size)] = static_lengths - for _ in range(3): - self.decode_fn(static_inputs, static_lengths) - graph = CUDAGraph() - with torch.cuda.graph(graph, pool=current_platform.get_global_graph_pool()): - static_outputs = self.decode_fn(static_inputs, static_lengths) - self.graphs[(batch_size, bucket_size)] = graph - self.static_outputs[(batch_size, bucket_size)] = static_outputs - logger.info(f"Captured graph for batch_size={batch_size}, bucket_size={bucket_size}") - - def find_nearest_padding(self, batch, bucket): - candidates = [ - (b, bk) for b in self.capture_batch_size for bk in self.capture_bucket_size if b >= batch and bk >= bucket - ] - return min(candidates, key=lambda x: x[0] * x[1]) if candidates else None - - def forward(self, codes, lengths): - batch_size, _, bucket_size = codes.shape - if batch_size > max(self.capture_batch_size) or bucket_size > max(self.capture_bucket_size): - return self.decode_fn(codes, lengths) - - graph_key = self.find_nearest_padding(batch_size, bucket_size) - if graph_key is None: - return self.decode_fn(codes, lengths) - - if graph_key not in self.graphs: - return self.decode_fn(codes, lengths) - - static_input = self.static_inputs[graph_key] - static_lengths = self.static_lengths[graph_key] - static_input.zero_() - static_lengths.zero_() - static_input[:batch_size, :, :bucket_size].copy_(codes) - static_lengths[:batch_size].copy_(lengths) - self.graphs[graph_key].replay() - return self.static_outputs[graph_key][:batch_size].clone() - - class OmniVoiceDecoder(nn.Module): """OmniVoice Stage 1: Token-to-audio decoder. @@ -155,7 +85,6 @@ def __init__(self, config: OmniVoiceConfig): self.quantizer = None self.fc2 = None self.acoustic_decoder = None - self.graphs: DecoderGraph | None = None @staticmethod def _mask_by_lengths(hidden_states: torch.Tensor, lengths: torch.Tensor) -> torch.Tensor: @@ -192,45 +121,6 @@ def _decode_acoustic_length_aware( hidden_states = decoder.tanh(hidden_states) return self._mask_by_lengths(hidden_states, lengths) - def _decode_impl( - self, - audio_codes: torch.Tensor, - target_lens: torch.Tensor | None = None, - ) -> torch.Tensor: - """Decode audio tokens to waveform. - - Args: - audio_codes: [B, 8, T] - 8-codebook audio token IDs - - Returns: - waveform: [B, 1, audio_samples] at 24kHz - """ - # Transpose: [B, 8, T] → [8, B, T] - codes = audio_codes.transpose(0, 1).long() - - # RVQ decode: sum codebook embeddings → [B, 1024, T] - quantized = self.quantizer.decode(codes) - - # Project: [B, 1024, T] → fc2 → [B, 256, T] - # Cast to fc2 weight dtype (may be fp16 when checkpoint stores weights as fp16), - # then upcast back to float32 — acoustic decoder ConvTranspose1d upsampling - # produces intermediate values that exceed the fp16 range (~65504), causing NaN. - quantized = self.fc2(quantized.transpose(1, 2).to(self.fc2.weight.dtype)).transpose(1, 2).float() - - # Acoustic decoder: [B, 256, T] → [B, 1, T*960] - if target_lens is None or not all( - hasattr(self.acoustic_decoder, name) for name in ("conv1", "block", "snake1", "conv2", "tanh") - ): - audio = self.acoustic_decoder(quantized) - else: - audio = self._decode_acoustic_length_aware(quantized, target_lens) - - # Ensure [B, 1, samples] - if audio.dim() == 2: - audio = audio.unsqueeze(1) - - return audio - @torch.inference_mode() def forward( self, @@ -269,10 +159,25 @@ def forward( if torch.any(lengths <= 0) or torch.any(lengths > audio_codes.shape[-1]): raise ValueError(f"Target lengths must be in [1, {audio_codes.shape[-1]}], got {lengths.tolist()}.") - if self.graphs is not None and self.graphs.captured: - return self.graphs.forward(audio_codes, lengths).to(device) + # Transpose: [B, 8, T] → [8, B, T] + codes = audio_codes.transpose(0, 1).long() + + # RVQ decode: sum codebook embeddings → [B, 1024, T] + quantized = self.quantizer.decode(codes) + + # Project: [B, 1024, T] → fc2 → [B, 256, T]. Keep the acoustic + # decoder in float32 to avoid ConvTranspose1d intermediate overflow. + quantized = self.fc2(quantized.transpose(1, 2).to(self.fc2.weight.dtype)).transpose(1, 2).float() + + # Acoustic decoder: [B, 256, T] → [B, 1, T*960] + if all(hasattr(self.acoustic_decoder, name) for name in ("conv1", "block", "snake1", "conv2", "tanh")): + audio = self._decode_acoustic_length_aware(quantized, lengths) else: - return self._decode_impl(audio_codes, lengths).to(device) + audio = self.acoustic_decoder(quantized) + + if audio.dim() == 2: + audio = audio.unsqueeze(1) + return audio.to(device) def _adjust_output_padding(self, decoder: nn.Module): """Adjust ConvTranspose1d output_padding (HiggsAudioV2 modification).""" @@ -362,12 +267,6 @@ def load_weights(self, model_dir: str, device: torch.device) -> None: self.acoustic_decoder.eval() self._loaded = True - self.graphs = DecoderGraph( - self._decode_impl, - device=device, - ) - self.graphs.warmup() - logger.info( "Loaded OmniVoice decoder: %d quantizers, fc2(%d→%d), acoustic decoder (%d weights)", num_quantizers, From 08b7ef9b1d4cc5cf5f86bda749e261fbefca477e Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 20 Aug 2026 23:48:13 +0800 Subject: [PATCH 04/17] add e2e test and unit test Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../test_omnivoice_expansion.py | 55 ++++++++++- .../online_serving/test_omnivoice_parity.py | 63 ++++++++++++ .../omnivoice/test_pipeline_batching.py | 97 +++++++++++++++++++ .../models/omnivoice/pipeline_omnivoice.py | 13 ++- 4 files changed, 218 insertions(+), 10 deletions(-) create mode 100644 tests/e2e/online_serving/test_omnivoice_parity.py create mode 100644 tests/model_executor/models/omnivoice/test_pipeline_batching.py diff --git a/tests/e2e/online_serving/test_omnivoice_expansion.py b/tests/e2e/online_serving/test_omnivoice_expansion.py index 3a042082e63..d2e3dd76940 100644 --- a/tests/e2e/online_serving/test_omnivoice_expansion.py +++ b/tests/e2e/online_serving/test_omnivoice_expansion.py @@ -31,10 +31,7 @@ MODEL = "k2-fsa/OmniVoice" STAGE_CONFIG = get_deploy_config_path("omnivoice.yaml") -EXTRA_ARGS = [ - "--trust-remote-code", - "--disable-log-stats", -] +EXTRA_ARGS = ["--trust-remote-code", "--disable-log-stats", "--max-num-seqs", "8"] TEST_PARAMS = [ OmniServerParams( model=MODEL, @@ -42,7 +39,18 @@ server_args=EXTRA_ARGS, ) ] - +STEP_EXECUTION_ARGS = [ + "--trust-remote-code", + "--enforce-eager", + "--step-execution", +] +STEP_EXECUTION_PARAMS = [ + OmniServerParams( + model=MODEL, + stage_config_path=STAGE_CONFIG, + server_args=STEP_EXECUTION_ARGS, + ) +] # Lower this in ``request_config`` via ``min_audio_bytes`` if a run produces legitimately short WAVs. _DEFAULT_MIN_AUDIO_BYTES = 5000 @@ -74,6 +82,43 @@ def test_speech_auto_voice(self, omni_server, online_client) -> None: } online_client.send_audio_speech_request(request_config) + @hardware_test(res={"cuda": "L4"}, num_cards=1) + def test_speech_auto_voice_batch(self, omni_server, openai_client) -> None: + """The same seeded request must match in request-batch and per-request execution.""" + batch_request_config = { + "model": omni_server.model, + "input": get_prompt("text"), + "response_format": "wav", + "seed": 42, + "min_audio_bytes": _DEFAULT_MIN_AUDIO_BYTES, + } + single_request_config = dict(batch_request_config) + batch_r = openai_client.send_audio_speech_request(batch_request_config, request_num=2) + single_r = openai_client.send_audio_speech_request(single_request_config)[0] + assert len(batch_r) == 2 + assert batch_r[0].audio_bytes == single_r.audio_bytes + assert batch_r[1].audio_bytes == single_r.audio_bytes + + +@pytest.mark.parametrize("omni_server", STEP_EXECUTION_PARAMS, indirect=True) +class TestOmniVoiceStepExecution: + """E2E tests for OmniVoice TTS model.""" + + @hardware_test(res={"cuda": "L4"}, num_cards=1) + def test_speech_auto_voice_step_execution(self, omni_server, openai_client) -> None: + """Test auto voice TTS generation (text only, no reference audio).""" + request_config = { + "model": omni_server.model, + "input": get_prompt("text"), + "response_format": "wav", + "timeout": 180.0, + "min_audio_bytes": _DEFAULT_MIN_AUDIO_BYTES, + "extra_params": { + "num_inference_steps": 32, + }, + } + openai_client.send_audio_speech_request(request_config) + @pytest.mark.parametrize("omni_server", TEST_PARAMS, indirect=True) class TestOmniVoiceSeed: diff --git a/tests/e2e/online_serving/test_omnivoice_parity.py b/tests/e2e/online_serving/test_omnivoice_parity.py new file mode 100644 index 00000000000..f06ce6b7c15 --- /dev/null +++ b/tests/e2e/online_serving/test_omnivoice_parity.py @@ -0,0 +1,63 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Cross-mode full-model parity tests for OmniVoice.""" + +import os + +os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn" + +import pytest +import requests + +from tests.helpers.mark import hardware_test +from tests.helpers.runtime import OmniServer +from tests.helpers.stage_config import get_deploy_config_path + +pytestmark = [pytest.mark.slow, pytest.mark.tts] + +MODEL = "k2-fsa/OmniVoice" +STAGE_CONFIG = get_deploy_config_path("omnivoice.yaml") +PROMPT = "The weather is nice today, perfect for a walk in the park." + + +def _generate_once(server_args: list[str]) -> bytes: + payload = { + "model": MODEL, + "input": PROMPT, + "language": "English", + "seed": 42, + "response_format": "wav", + "extra_params": {"num_inference_steps": 32}, + } + with OmniServer( + MODEL, + server_args, + use_omni=True, + env_dict={"OMNIVOICE_CUDA_GRAPH": "0"}, + ) as server: + response = requests.post( + f"http://{server.host}:{server.port}/v1/audio/speech", + json=payload, + timeout=600, + ) + response.raise_for_status() + assert response.content.startswith(b"RIFF") + return response.content + + +@hardware_test(res={"cuda": "L4"}, num_cards=1) +def test_request_mode_and_step_execution_b1_parity() -> None: + """B=1 request mode and step execution must produce identical seeded WAV bytes.""" + common_args = [ + "--trust-remote-code", + "--disable-log-stats", + "--deploy-config", + STAGE_CONFIG, + "--max-num-seqs", + "1", + ] + + request_audio = _generate_once(common_args) + step_audio = _generate_once([*common_args, "--step-execution", "--enforce-eager"]) + + assert request_audio == step_audio diff --git a/tests/model_executor/models/omnivoice/test_pipeline_batching.py b/tests/model_executor/models/omnivoice/test_pipeline_batching.py new file mode 100644 index 00000000000..c7519935de0 --- /dev/null +++ b/tests/model_executor/models/omnivoice/test_pipeline_batching.py @@ -0,0 +1,97 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Lightweight batching invariants for the OmniVoice pipeline.""" + +from types import SimpleNamespace + +import pytest +import torch +import torch.nn as nn + +from vllm_omni.diffusion.models.omnivoice.pipeline_omnivoice import OmniVoicePipeline + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + + +def _minimal_pipeline() -> OmniVoicePipeline: + pipeline = OmniVoicePipeline.__new__(OmniVoicePipeline) + nn.Module.__init__(pipeline) + pipeline.config = SimpleNamespace(audio_mask_id=1024) + return pipeline + + +def _state(request_id: str, seq_len: int, fill: int): + return SimpleNamespace( + request_id=request_id, + latents=torch.full((2, 8, seq_len), fill, dtype=torch.long), + extra={ + "audio_mask": torch.ones(2, seq_len, dtype=torch.bool), + "attn_mask": torch.ones(2, 1, seq_len, seq_len, dtype=torch.bool), + }, + ) + + +@pytest.mark.parametrize("batch_size", [1, 2, 3, 4]) +def test_cfg_layout_round_trip(batch_size: int) -> None: + request_major = torch.arange(2 * batch_size).reshape(2 * batch_size, 1) + + cfg_major = OmniVoicePipeline._request_major_to_cfg_major(request_major, batch_size) + restored = OmniVoicePipeline._cfg_major_to_request_major(cfg_major, batch_size) + + expected_cfg = torch.tensor([*range(0, 2 * batch_size, 2), *range(1, 2 * batch_size, 2)]).reshape(2 * batch_size, 1) + assert torch.equal(cfg_major, expected_cfg) + assert torch.equal(restored, request_major) + + +def test_new_longer_request_repads_existing_state_and_masks() -> None: + pipeline = _minimal_pipeline() + old = _state("old", seq_len=5, fill=11) + new = _state("new", seq_len=8, fill=22) + + old_latents = old.latents.clone() + old_audio_mask = old.extra["audio_mask"].clone() + old_attn_mask = old.extra["attn_mask"].clone() + + states = pipeline.prepare_state_batch([old, new], new_request_ids=["new"]) + + assert states[0] is old + assert states[1] is new + assert old.latents.shape == new.latents.shape == (2, 8, 8) + assert old.extra["audio_mask"].shape == new.extra["audio_mask"].shape == (2, 8) + assert old.extra["attn_mask"].shape == new.extra["attn_mask"].shape == (2, 1, 8, 8) + + torch.testing.assert_close(old.latents[..., :5], old_latents) + assert torch.all(old.latents[..., 5:] == 1024) + assert torch.equal(old.extra["audio_mask"][..., :5], old_audio_mask) + assert not torch.any(old.extra["audio_mask"][..., 5:]) + assert torch.equal(old.extra["attn_mask"][..., :5, :5], old_attn_mask) + assert not torch.any(old.extra["attn_mask"][..., 5:, :]) + assert not torch.any(old.extra["attn_mask"][..., :, 5:]) + + +def test_later_short_request_is_padded_to_existing_active_length() -> None: + pipeline = _minimal_pipeline() + first = _state("first", seq_len=8, fill=11) + second = _state("second", seq_len=8, fill=22) + newcomer = _state("new", seq_len=6, fill=33) + + pipeline.prepare_state_batch([first, second, newcomer], new_request_ids=["new"]) + + assert first.latents.shape == second.latents.shape == newcomer.latents.shape == (2, 8, 8) + assert torch.all(newcomer.latents[..., :6] == 33) + assert torch.all(newcomer.latents[..., 6:] == 1024) + assert not torch.any(newcomer.extra["audio_mask"][..., 6:]) + assert not torch.any(newcomer.extra["attn_mask"][..., 6:, :]) + assert not torch.any(newcomer.extra["attn_mask"][..., :, 6:]) + + +def test_audio_outputs_are_trimmed_to_each_target_length() -> None: + audio = torch.arange(2 * 10 * 960, dtype=torch.float32).reshape(2, 1, 10 * 960) + + outputs = OmniVoicePipeline._split_audio_outputs(audio, target_lens=[3, 7]) + + assert len(outputs) == 2 + assert outputs[0].output.shape == (1, 1, 3 * 960) + assert outputs[1].output.shape == (1, 1, 7 * 960) + assert torch.equal(outputs[0].output, audio[0:1, :, : 3 * 960]) + assert torch.equal(outputs[1].output, audio[1:2, :, : 7 * 960]) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 7dbe61e9ba7..441e5d91590 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -407,6 +407,13 @@ def _cfg_major_to_request_major(x: torch.Tensor, batch_size: int) -> torch.Tenso pairs = torch.stack([x[:batch_size], x[batch_size:]], dim=1) return pairs.reshape(2 * batch_size, *x.shape[1:]) + @staticmethod + def _split_audio_outputs(audio: torch.Tensor, target_lens: Sequence[int]) -> list[DiffusionOutput]: + """Split a padded waveform batch into one length-trimmed output per request.""" + return [ + DiffusionOutput(output=audio[i : i + 1, :, : target_len * 960]) for i, target_len in enumerate(target_lens) + ] + def prepare_state_batch(self, states: list[StepRequestState], new_request_ids: list[str]) -> list[StepRequestState]: if not new_request_ids: return states @@ -637,7 +644,6 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: {"text": "...", "ref_audio": (samples, sr), "ref_text": "...", "lang": "...", "instruct": "..."} """ - outputs: list[DiffusionOutput] = [] prepared_requests: list[_PreparedOmniVoiceRequest] = [] for request in req.requests: prompt = request.prompt if request.prompt else "" @@ -666,10 +672,7 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: ) audio = self.decoder(tokens, batch_target_len) # [B, 1, max_target_len * 960] - for i in range(len(batch_target_len)): - audio_output = audio[i : i + 1, :, : batch_target_len[i] * 960] - outputs.append(DiffusionOutput(output=audio_output)) - return outputs + return self._split_audio_outputs(audio, batch_target_len) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: """Load weights from model directory (not from the iterator). From 0f1b8d0c18a29b0e264d6629168e45089306623f Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Sun, 23 Aug 2026 20:16:54 +0800 Subject: [PATCH 05/17] support varlen_attn for both eager/graph and request-batch/step-execution Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../test_omnivoice_expansion.py | 20 + .../omnivoice/test_cuda_graph_generator.py | 291 +++++------ .../omnivoice/test_pipeline_batching.py | 97 ---- .../models/omnivoice/pipeline_omnivoice.py | 228 +++------ .../models/omnivoice/omnivoice_generator.py | 482 ++++++++++-------- 5 files changed, 524 insertions(+), 594 deletions(-) delete mode 100644 tests/model_executor/models/omnivoice/test_pipeline_batching.py diff --git a/tests/e2e/online_serving/test_omnivoice_expansion.py b/tests/e2e/online_serving/test_omnivoice_expansion.py index d2e3dd76940..1a17d11d45d 100644 --- a/tests/e2e/online_serving/test_omnivoice_expansion.py +++ b/tests/e2e/online_serving/test_omnivoice_expansion.py @@ -119,6 +119,26 @@ def test_speech_auto_voice_step_execution(self, omni_server, openai_client) -> N } openai_client.send_audio_speech_request(request_config) + @hardware_test(res={"cuda": "L4"}, num_cards=1) + def test_speech_auto_voice_batch_step_execution(self, omni_server, openai_client) -> None: + """The same seeded request must match in step-execution-batch and per-request execution.""" + batch_request_config = { + "model": omni_server.model, + "input": get_prompt("text"), + "response_format": "wav", + "seed": 42, + "min_audio_bytes": _DEFAULT_MIN_AUDIO_BYTES, + "extra_params": { + "num_inference_steps": 32, + }, + } + single_request_config = dict(batch_request_config) + batch_r = openai_client.send_audio_speech_request(batch_request_config, request_num=2) + single_r = openai_client.send_audio_speech_request(single_request_config)[0] + assert len(batch_r) == 2 + assert batch_r[0].audio_bytes == single_r.audio_bytes + assert batch_r[1].audio_bytes == single_r.audio_bytes + @pytest.mark.parametrize("omni_server", TEST_PARAMS, indirect=True) class TestOmniVoiceSeed: diff --git a/tests/model_executor/models/omnivoice/test_cuda_graph_generator.py b/tests/model_executor/models/omnivoice/test_cuda_graph_generator.py index b55f8279c44..ef05bf7383c 100644 --- a/tests/model_executor/models/omnivoice/test_cuda_graph_generator.py +++ b/tests/model_executor/models/omnivoice/test_cuda_graph_generator.py @@ -6,19 +6,27 @@ Verifies that _OmniVoiceCUDAGraphForward produces results equivalent to eager mode across three scenarios: - Exact-size inputs (no padding) → bit-identical - - Padded inputs (padded to nearest bucket) → correct slicing, exact match - for position-independent models - - Oversized inputs (fallback to eager) → bit-identical + - Padded inputs (padded to nearest bucket) → correct slicing and exact match + - Oversized inputs (128-aligned lazy capture) → bit-identical -Uses a lightweight synthetic generator to keep warmup fast and avoid -loading actual model weights. +Uses a small randomly initialized OmniVoiceGenerator so every test exercises +the production embedding, transformer, varlen-attention, and logits paths +without loading checkpoint weights. """ from __future__ import annotations +from types import SimpleNamespace + import pytest import torch -import torch.nn as nn +from vllm.utils.math_utils import round_up + +from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( + OmniVoiceGenerator, + _OmniVoiceCUDAGraphForward, +) +from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig pytestmark = [ pytest.mark.core_model, @@ -29,72 +37,45 @@ DEVICE = torch.device("cuda:0") NUM_CB = 8 VOCAB = 1025 -HIDDEN = 64 CAPTURE_SIZES = [32, 64, 128] -HEAD_DIM = 64 # --------------------------------------------------------------------------- -# Synthetic generator: matches _OmniVoiceCUDAGraphForward's interface +# Helpers # --------------------------------------------------------------------------- -class _SyntheticGenerator: - """Position-independent Embedding+Linear replacing OmniVoiceGenerator's 28-layer transformer. - - Each token position is processed independently (no attention), so padded - runs produce bit-identical results to non-padded runs on the same positions. - This makes the padding/slicing logic easy to verify exactly. - """ - - class _Cfg: - num_audio_codebook = NUM_CB - - def __init__(self, device: torch.device): - self.config = self._Cfg() - self.text_embedding = nn.Embedding(1000, HIDDEN).to(device).eval() - self._linear = nn.Linear(HIDDEN, NUM_CB * VOCAB, bias=False).to(device).eval() - self._rope_table: torch.Tensor | None = None - - @property - def model_dtype(self) -> torch.dtype: - """Part of the interface _OmniVoiceCUDAGraphForward expects of a generator.""" - return self.text_embedding.weight.dtype +def _eager(gen: OmniVoiceGenerator, ids: torch.Tensor, mask: torch.Tensor, cu_seqs) -> torch.Tensor: + """Run the production varlen step outside CUDA Graph.""" + seq_len = ids.shape[0] + rope_table = gen._rope_table_for(seq_len, ids.device, gen.model_dtype) + return gen._step_forward(ids, mask, cu_seqs, rope_table) - def _rope_table_for(self, seq_len: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor: - """Part of the interface _OmniVoiceCUDAGraphForward expects of a generator.""" - return torch.zeros(seq_len, HEAD_DIM, device=device, dtype=dtype) - - def _step_forward( - self, - input_ids: torch.Tensor, - audio_mask: torch.Tensor, - attention_mask: torch.Tensor | None, - rope_table: torch.Tensor, - ) -> torch.Tensor: - two_b, _, S = input_ids.shape - x = self.text_embedding(input_ids[:, 0, :].clamp(0, 999)) # [two_b, S, H] - logits = self._linear(x) # [two_b, S, 8*1025] - return logits.view(two_b, S, NUM_CB, VOCAB).permute(0, 2, 1, 3) # [two_b, 8, S, 1025] - - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - -def _eager(gen: _SyntheticGenerator, ids: torch.Tensor, mask: torch.Tensor, attn: torch.Tensor) -> torch.Tensor: - """Run _step_forward with a zero RoPE table (synthetic gen ignores it).""" - two_b, _, S = ids.shape - return gen._step_forward(ids, mask, attn, gen._rope_table_for(S, ids.device, torch.float32)) +def _bucketed_eager(gen, wrapper, ids, mask, cu_seqs, batch_size) -> torch.Tensor: + """Run eager with the exact padded shape and metadata used by Graph.""" + seq_len = ids.shape[0] + bucket = wrapper._find_bucket(batch_size, seq_len) + if bucket is None: + bucket = round_up(seq_len, wrapper._LAZY_CAPTURE_ALIGNMENT) + padded_ids, padded_mask = wrapper._pad_inputs(ids, mask, bucket) + padded_cu_seqs = cu_seqs.clone() + padded_cu_seqs[-1] = bucket + return _eager(gen, padded_ids, padded_mask, padded_cu_seqs)[:, :seq_len, :] def _make_inputs(seq_len: int, device: torch.device = DEVICE): - two_b = 2 - ids = torch.randint(0, 100, (two_b, NUM_CB, seq_len), dtype=torch.long, device=device) - mask = torch.ones(two_b, seq_len, dtype=torch.bool, device=device) - attn = torch.ones(two_b, 1, seq_len, seq_len, dtype=torch.bool, device=device) - return ids, mask, attn + ids = torch.randint(0, 100, (seq_len, NUM_CB), dtype=torch.long, device=device) + mask = torch.ones(seq_len, dtype=torch.bool, device=device) + cond_len = (seq_len + 1) // 2 + uncond_len = seq_len - cond_len + + cu_seqs = torch.tensor( + [0, cond_len, cond_len + uncond_len, cond_len + uncond_len], + dtype=torch.int32, + device=device, + ) + return ids, mask, cu_seqs # --------------------------------------------------------------------------- @@ -105,15 +86,26 @@ def _make_inputs(seq_len: int, device: torch.device = DEVICE): @pytest.fixture(scope="module") def gen(): torch.manual_seed(42) - return _SyntheticGenerator(DEVICE) + config = OmniVoiceConfig( + audio_vocab_size=VOCAB, + num_audio_codebook=NUM_CB, + llm_config={ + "hidden_size": 64, + "num_hidden_layers": 2, + "num_attention_heads": 4, + "num_key_value_heads": 2, + "intermediate_size": 128, + "vocab_size": 128, + "max_position_embeddings": 4096, + "head_dim": 16, + }, + enable_cuda_graph=False, + ) + return OmniVoiceGenerator(config, SimpleNamespace(max_num_seqs=2)).to(DEVICE).eval() @pytest.fixture(scope="module") def wrapper(gen): - from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( - _OmniVoiceCUDAGraphForward, - ) - w = _OmniVoiceCUDAGraphForward(gen, capture_sizes=CAPTURE_SIZES) w.warmup(DEVICE) return w @@ -127,10 +119,10 @@ def wrapper(gen): @pytest.mark.parametrize("seq_len", CAPTURE_SIZES) def test_exact_size_bit_identical(gen, wrapper, seq_len): """When input exactly matches a captured bucket, output must be bit-identical to eager.""" - ids, mask, attn = _make_inputs(seq_len) + ids, mask, cu_seqs = _make_inputs(seq_len) with torch.no_grad(): - eager_out = _eager(gen, ids, mask, attn) - graph_out = wrapper(ids, mask, attn) + eager_out = _eager(gen, ids, mask, cu_seqs) + graph_out = wrapper(ids, mask, cu_seqs, 1) torch.testing.assert_close(graph_out, eager_out, atol=0, rtol=0) @@ -140,40 +132,36 @@ def test_exact_size_bit_identical(gen, wrapper, seq_len): @pytest.mark.parametrize("seq_len", [1, 15, 33, 60, 100]) -def test_padded_output_shape(gen, wrapper, seq_len): +def test_padded_output_shape(wrapper, seq_len): """Graph output must be sliced back to actual seq_len, not the bucket size.""" - ids, mask, attn = _make_inputs(seq_len) + ids, mask, cu_seqs = _make_inputs(seq_len) with torch.no_grad(): - graph_out = wrapper(ids, mask, attn) - assert graph_out.shape == (2, NUM_CB, seq_len, VOCAB) + graph_out = wrapper(ids, mask, cu_seqs, 1) + assert graph_out.shape == (NUM_CB, seq_len, VOCAB) @pytest.mark.parametrize("seq_len", [15, 33, 60, 100]) def test_padded_output_matches_eager(gen, wrapper, seq_len): - """Padded graph output must equal eager output at actual positions. - - The synthetic model has no attention across positions, so zero-padding - does not affect non-padded positions — exact match is expected. - """ - ids, mask, attn = _make_inputs(seq_len) + """Padded graph output must equal eager output at actual positions.""" + ids, mask, cu_seqs = _make_inputs(seq_len) with torch.no_grad(): - eager_out = _eager(gen, ids, mask, attn) - graph_out = wrapper(ids, mask, attn) + eager_out = _bucketed_eager(gen, wrapper, ids, mask, cu_seqs, 1) + graph_out = wrapper(ids, mask, cu_seqs, 1) torch.testing.assert_close(graph_out, eager_out, atol=0, rtol=0) # --------------------------------------------------------------------------- -# 3. Oversized inputs → fallback to eager (lazy capture), bit-identical +# 3. Oversized inputs → aligned lazy capture, bit-identical # --------------------------------------------------------------------------- @pytest.mark.parametrize("seq_len", [129, 200, 256]) -def test_fallback_eager_bit_identical(gen, wrapper, seq_len): - """Sequences exceeding the largest bucket fall back to lazy capture → bit-identical.""" - ids, mask, attn = _make_inputs(seq_len) +def test_aligned_lazy_capture_bit_identical(gen, wrapper, seq_len): + """Sequences beyond the static plan use aligned lazy capture and remain bit-identical.""" + ids, mask, cu_seqs = _make_inputs(seq_len) with torch.no_grad(): - eager_out = _eager(gen, ids, mask, attn) - graph_out = wrapper(ids, mask, attn) + eager_out = _bucketed_eager(gen, wrapper, ids, mask, cu_seqs, 1) + graph_out = wrapper(ids, mask, cu_seqs, 1) torch.testing.assert_close(graph_out, eager_out, atol=0, rtol=0) @@ -184,77 +172,96 @@ def test_fallback_eager_bit_identical(gen, wrapper, seq_len): def test_deterministic_across_calls(wrapper): """Same input must produce identical output on repeated CUDA graph replays.""" - ids, mask, attn = _make_inputs(32) + ids, mask, cu_seqs = _make_inputs(32) with torch.no_grad(): - out1 = wrapper(ids, mask, attn) - out2 = wrapper(ids, mask, attn) + out1 = wrapper(ids, mask, cu_seqs, 1).clone() + out2 = wrapper(ids, mask, cu_seqs, 1).clone() torch.testing.assert_close(out1, out2, atol=0, rtol=0) -# --------------------------------------------------------------------------- -# 5. _find_bucket logic (CPU, no CUDA graph) -# --------------------------------------------------------------------------- - - -def test_find_bucket_returns_nearest_bucket(): - """_find_bucket must return the smallest bucket >= seq_len, or None if all are smaller.""" - from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( - _OmniVoiceCUDAGraphForward, - ) +def test_replay_uses_updated_cu_seqs(gen, wrapper): + """The same graph key must honor new sequence boundaries on replay.""" + ids, mask, _ = _make_inputs(48) + cu_seqs_a = torch.tensor([0, 32, 48, 48], dtype=torch.int32, device=DEVICE) + cu_seqs_b = torch.tensor([0, 17, 48, 48], dtype=torch.int32, device=DEVICE) - w = _OmniVoiceCUDAGraphForward.__new__(_OmniVoiceCUDAGraphForward) - w._capture_sizes = [32, 64, 128] - w._graphs = {} - - assert w._find_bucket(1) == 32 - assert w._find_bucket(32) == 32 - assert w._find_bucket(33) == 64 - assert w._find_bucket(64) == 64 - assert w._find_bucket(128) == 128 - assert w._find_bucket(129) is None - assert w._find_bucket(1000) is None + with torch.no_grad(): + eager_a = _bucketed_eager(gen, wrapper, ids, mask, cu_seqs_a, 1) + graph_a = wrapper(ids, mask, cu_seqs_a, 1).clone() + eager_b = _bucketed_eager(gen, wrapper, ids, mask, cu_seqs_b, 1) + graph_b = wrapper(ids, mask, cu_seqs_b, 1).clone() + torch.testing.assert_close(graph_a, eager_a, atol=0, rtol=0) + torch.testing.assert_close(graph_b, eager_b, atol=0, rtol=0) + assert not torch.equal(graph_a, graph_b) -# --------------------------------------------------------------------------- -# 6. enable_cuda_graph=False produces same tokens as enable_cuda_graph=True -# --------------------------------------------------------------------------- +def test_batch_two_cu_seqs_matches_eager(gen, wrapper): + """B=2 uses four real sequence segments plus the fixed tail slot.""" + ids, mask, _ = _make_inputs(48) + cu_seqs = torch.tensor([0, 10, 18, 32, 48, 48], dtype=torch.int32, device=DEVICE) -def test_cuda_graph_disabled_matches_eager_generator(): - """OmniVoiceGenerator with enable_cuda_graph=False must produce the same - _step_forward output as one with enable_cuda_graph=True (after warmup). + assert cu_seqs.numel() == 2 * 2 + 2 + with torch.no_grad(): + eager_out = _bucketed_eager(gen, wrapper, ids, mask, cu_seqs, 2) + graph_out = wrapper(ids, mask, cu_seqs, 2).clone() - Uses a 2-layer config so warmup completes quickly without loading weights. - """ - from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import OmniVoiceGenerator - from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig + torch.testing.assert_close(graph_out, eager_out, atol=0, rtol=0) - cfg = OmniVoiceConfig() - cfg.llm_num_hidden_layers = 2 - cfg.num_hidden_layers = 2 - torch.manual_seed(0) - gen_eager = OmniVoiceGenerator(cfg).to(DEVICE).eval() - gen_eager.config.enable_cuda_graph = False - gen_eager._cuda_graph_fwd = None +def test_lazy_graph_cache_uses_lru_eviction(wrapper): + """A lazy-cache hit must protect that graph from the next eviction.""" + original_limit = wrapper._MAX_LAZY_GRAPHS + wrapper._lazy_graphs.clear() + wrapper._MAX_LAZY_GRAPHS = 2 + try: + # Static B=1 coverage ends at 128. These lengths round to three + # distinct 128-aligned lazy keys: 256, 384, and 512. + for seq_len in (129, 257): + ids, mask, cu_seqs = _make_inputs(seq_len) + with torch.no_grad(): + wrapper(ids, mask, cu_seqs, 1) + assert list(wrapper._lazy_graphs) == [(1, 256), (1, 384)] + + # Refresh key 256, making key 384 the least recently used entry. + ids, mask, cu_seqs = _make_inputs(129) + with torch.no_grad(): + wrapper(ids, mask, cu_seqs, 1) + assert list(wrapper._lazy_graphs) == [(1, 384), (1, 256)] + + # Inserting key 512 must evict key 384, not the recently hit key 256. + ids, mask, cu_seqs = _make_inputs(385) + with torch.no_grad(): + wrapper(ids, mask, cu_seqs, 1) + assert list(wrapper._lazy_graphs) == [(1, 256), (1, 512)] + finally: + wrapper._MAX_LAZY_GRAPHS = original_limit + wrapper._lazy_graphs.clear() - torch.manual_seed(0) - gen_graph = OmniVoiceGenerator(cfg).to(DEVICE).eval() - gen_graph.config.cuda_graph_capture_sizes = [64] - from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import _OmniVoiceCUDAGraphForward - gen_graph._cuda_graph_fwd = _OmniVoiceCUDAGraphForward(gen_graph, capture_sizes=[64]) - gen_graph._cuda_graph_fwd.warmup(DEVICE) +# --------------------------------------------------------------------------- +# 5. _find_bucket logic (CPU, no CUDA graph) +# --------------------------------------------------------------------------- - seq_len = 64 - ids = torch.zeros(2, cfg.num_audio_codebook, seq_len, dtype=torch.long, device=DEVICE) - mask = torch.ones(2, seq_len, dtype=torch.bool, device=DEVICE) - attn = torch.ones(2, 1, seq_len, seq_len, dtype=torch.bool, device=DEVICE) - rope_table = gen_eager._rope_table_for(seq_len, DEVICE, gen_eager.model_dtype) +def test_find_bucket_returns_nearest_bucket(): + """_find_bucket must return the smallest bucket >= seq_len, or None if all are smaller.""" + from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( + _OmniVoiceCUDAGraphForward, + ) - with torch.no_grad(): - eager_logits = gen_eager._step_forward(ids, mask, attn, rope_table) - graph_logits = gen_graph._cuda_graph_fwd(ids, mask, attn) + w = _OmniVoiceCUDAGraphForward.__new__(_OmniVoiceCUDAGraphForward) + w.capture_bucket_sizes_by_batch = { + 1: [32, 64, 128], + 2: [64, 128, 256], + } + w._graphs = {} - torch.testing.assert_close(graph_logits, eager_logits, atol=0, rtol=0) + assert w._find_bucket(1, 1) == 32 + assert w._find_bucket(1, 32) == 32 + assert w._find_bucket(1, 33) == 64 + assert w._find_bucket(1, 128) == 128 + assert w._find_bucket(1, 129) is None + assert w._find_bucket(2, 33) == 64 + assert w._find_bucket(2, 129) == 256 + assert w._find_bucket(3, 64) is None diff --git a/tests/model_executor/models/omnivoice/test_pipeline_batching.py b/tests/model_executor/models/omnivoice/test_pipeline_batching.py deleted file mode 100644 index c7519935de0..00000000000 --- a/tests/model_executor/models/omnivoice/test_pipeline_batching.py +++ /dev/null @@ -1,97 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Lightweight batching invariants for the OmniVoice pipeline.""" - -from types import SimpleNamespace - -import pytest -import torch -import torch.nn as nn - -from vllm_omni.diffusion.models.omnivoice.pipeline_omnivoice import OmniVoicePipeline - -pytestmark = [pytest.mark.core_model, pytest.mark.cpu] - - -def _minimal_pipeline() -> OmniVoicePipeline: - pipeline = OmniVoicePipeline.__new__(OmniVoicePipeline) - nn.Module.__init__(pipeline) - pipeline.config = SimpleNamespace(audio_mask_id=1024) - return pipeline - - -def _state(request_id: str, seq_len: int, fill: int): - return SimpleNamespace( - request_id=request_id, - latents=torch.full((2, 8, seq_len), fill, dtype=torch.long), - extra={ - "audio_mask": torch.ones(2, seq_len, dtype=torch.bool), - "attn_mask": torch.ones(2, 1, seq_len, seq_len, dtype=torch.bool), - }, - ) - - -@pytest.mark.parametrize("batch_size", [1, 2, 3, 4]) -def test_cfg_layout_round_trip(batch_size: int) -> None: - request_major = torch.arange(2 * batch_size).reshape(2 * batch_size, 1) - - cfg_major = OmniVoicePipeline._request_major_to_cfg_major(request_major, batch_size) - restored = OmniVoicePipeline._cfg_major_to_request_major(cfg_major, batch_size) - - expected_cfg = torch.tensor([*range(0, 2 * batch_size, 2), *range(1, 2 * batch_size, 2)]).reshape(2 * batch_size, 1) - assert torch.equal(cfg_major, expected_cfg) - assert torch.equal(restored, request_major) - - -def test_new_longer_request_repads_existing_state_and_masks() -> None: - pipeline = _minimal_pipeline() - old = _state("old", seq_len=5, fill=11) - new = _state("new", seq_len=8, fill=22) - - old_latents = old.latents.clone() - old_audio_mask = old.extra["audio_mask"].clone() - old_attn_mask = old.extra["attn_mask"].clone() - - states = pipeline.prepare_state_batch([old, new], new_request_ids=["new"]) - - assert states[0] is old - assert states[1] is new - assert old.latents.shape == new.latents.shape == (2, 8, 8) - assert old.extra["audio_mask"].shape == new.extra["audio_mask"].shape == (2, 8) - assert old.extra["attn_mask"].shape == new.extra["attn_mask"].shape == (2, 1, 8, 8) - - torch.testing.assert_close(old.latents[..., :5], old_latents) - assert torch.all(old.latents[..., 5:] == 1024) - assert torch.equal(old.extra["audio_mask"][..., :5], old_audio_mask) - assert not torch.any(old.extra["audio_mask"][..., 5:]) - assert torch.equal(old.extra["attn_mask"][..., :5, :5], old_attn_mask) - assert not torch.any(old.extra["attn_mask"][..., 5:, :]) - assert not torch.any(old.extra["attn_mask"][..., :, 5:]) - - -def test_later_short_request_is_padded_to_existing_active_length() -> None: - pipeline = _minimal_pipeline() - first = _state("first", seq_len=8, fill=11) - second = _state("second", seq_len=8, fill=22) - newcomer = _state("new", seq_len=6, fill=33) - - pipeline.prepare_state_batch([first, second, newcomer], new_request_ids=["new"]) - - assert first.latents.shape == second.latents.shape == newcomer.latents.shape == (2, 8, 8) - assert torch.all(newcomer.latents[..., :6] == 33) - assert torch.all(newcomer.latents[..., 6:] == 1024) - assert not torch.any(newcomer.extra["audio_mask"][..., 6:]) - assert not torch.any(newcomer.extra["attn_mask"][..., 6:, :]) - assert not torch.any(newcomer.extra["attn_mask"][..., :, 6:]) - - -def test_audio_outputs_are_trimmed_to_each_target_length() -> None: - audio = torch.arange(2 * 10 * 960, dtype=torch.float32).reshape(2, 1, 10 * 960) - - outputs = OmniVoicePipeline._split_audio_outputs(audio, target_lens=[3, 7]) - - assert len(outputs) == 2 - assert outputs[0].output.shape == (1, 1, 3 * 960) - assert outputs[1].output.shape == (1, 1, 7 * 960) - assert torch.equal(outputs[0].output, audio[0:1, :, : 3 * 960]) - assert torch.equal(outputs[1].output, audio[1:2, :, : 7 * 960]) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 441e5d91590..0c5185d38f3 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -37,6 +37,7 @@ from vllm_omni.model_executor.models.omnivoice.omnivoice_decoder import OmniVoiceDecoder from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( OmniVoiceGenerator, + _build_cu_seqs, _get_time_steps, _gumbel_sample, ) @@ -57,7 +58,7 @@ class _PreparedOmniVoiceRequest: input_ids: torch.Tensor audio_mask: torch.Tensor - attention_mask: torch.Tensor + cond_len: int target_len: int seed: int | None @@ -331,29 +332,16 @@ def _prepare_request_input( ) cond_len = cond_ids.shape[1] uncond_ids = target_ids.clone() - uncond_len = target_len - max_len = max(cond_len, uncond_len) - if uncond_len < max_len: - pad = torch.full( - (num_cb, max_len - uncond_len), - mask_id, - dtype=torch.long, - device=self.device, - ) - uncond_ids = torch.cat([uncond_ids, pad], dim=1) + input_ids = torch.cat([cond_ids, uncond_ids], dim=1).transpose(0, 1).contiguous() - input_ids = torch.stack([cond_ids, uncond_ids]) - audio_mask = torch.zeros(2, max_len, dtype=torch.bool, device=self.device) - audio_mask[0, text_len:cond_len] = True - audio_mask[1, :uncond_len] = True - attention_mask = torch.zeros(2, 1, max_len, max_len, dtype=torch.bool, device=self.device) - attention_mask[0, :, :cond_len, :cond_len] = True - attention_mask[1, :, :uncond_len, :uncond_len] = True + max_len = input_ids.shape[0] + audio_mask = torch.zeros(max_len, dtype=torch.bool, device=self.device) + audio_mask[text_len:] = True return _PreparedOmniVoiceRequest( input_ids=input_ids, audio_mask=audio_mask, - attention_mask=attention_mask, + cond_len=cond_len, target_len=target_len, seed=seed, ) @@ -361,83 +349,24 @@ def _prepare_request_input( def _collate_request_inputs( self, prepared_requests: Sequence[_PreparedOmniVoiceRequest], - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Right-pad requests and arrange rows as all cond followed by all uncond.""" - max_len = max(request.input_ids.shape[-1] for request in prepared_requests) - input_pairs: list[torch.Tensor] = [] - audio_mask_pairs: list[torch.Tensor] = [] - attention_mask_pairs: list[torch.Tensor] = [] - mask_id = self.config.audio_mask_id + ) -> tuple[torch.Tensor, torch.Tensor, list[int]]: + """Pack request-major [cond, uncond] token sequences.""" + input_ids: list[torch.Tensor] = [] + audio_masks: list[torch.Tensor] = [] + cond_lens: list[int] = [] for request in prepared_requests: - pad_len = max_len - request.input_ids.shape[-1] - input_ids = request.input_ids + input_id = request.input_ids audio_mask = request.audio_mask - attention_mask = request.attention_mask - if pad_len: - input_ids = F.pad(input_ids, (0, pad_len), value=mask_id) - audio_mask = F.pad(audio_mask, (0, pad_len), value=False) - attention_mask = F.pad(attention_mask, (0, pad_len, 0, pad_len), value=False) - input_pairs.append(input_ids) - audio_mask_pairs.append(audio_mask) - attention_mask_pairs.append(attention_mask) - - input_pairs_tensor = torch.stack(input_pairs, dim=0) - audio_mask_pairs_tensor = torch.stack(audio_mask_pairs, dim=0) - attention_mask_pairs_tensor = torch.stack(attention_mask_pairs, dim=0) - return ( - torch.cat([input_pairs_tensor[:, 0], input_pairs_tensor[:, 1]], dim=0), - torch.cat([audio_mask_pairs_tensor[:, 0], audio_mask_pairs_tensor[:, 1]], dim=0), - torch.cat([attention_mask_pairs_tensor[:, 0], attention_mask_pairs_tensor[:, 1]], dim=0), - ) - - @staticmethod - def _request_major_to_cfg_major(x: torch.Tensor, batch_size: int) -> torch.Tensor: - """Convert [cond0, uncond0, ...] to [cond0, ..., uncond0, ...].""" - if x.shape[0] != 2 * batch_size: - raise ValueError(f"Expected {2 * batch_size} CFG rows for batch size {batch_size}, got {x.shape[0]}.") - pairs = x.reshape(batch_size, 2, *x.shape[1:]) - return torch.cat([pairs[:, 0], pairs[:, 1]], dim=0) - - @staticmethod - def _cfg_major_to_request_major(x: torch.Tensor, batch_size: int) -> torch.Tensor: - """Convert [cond0, ..., uncond0, ...] to [cond0, uncond0, ...].""" - if x.shape[0] != 2 * batch_size: - raise ValueError(f"Expected {2 * batch_size} CFG rows for batch size {batch_size}, got {x.shape[0]}.") - pairs = torch.stack([x[:batch_size], x[batch_size:]], dim=1) - return pairs.reshape(2 * batch_size, *x.shape[1:]) - - @staticmethod - def _split_audio_outputs(audio: torch.Tensor, target_lens: Sequence[int]) -> list[DiffusionOutput]: - """Split a padded waveform batch into one length-trimmed output per request.""" - return [ - DiffusionOutput(output=audio[i : i + 1, :, : target_len * 960]) for i, target_len in enumerate(target_lens) - ] - - def prepare_state_batch(self, states: list[StepRequestState], new_request_ids: list[str]) -> list[StepRequestState]: - if not new_request_ids: - return states - else: - max_target_len = max([state.latents.shape[-1] for state in states]) - if states[0].extra.get("max_target_len", None) == max_target_len: - states_to_repad = [state for state in states if state.request_id in new_request_ids] - self.repadding_state(states_to_repad, max_target_len) - return states + cond_len = request.cond_len + input_ids.append(input_id) + audio_masks.append(audio_mask) + cond_lens.append(cond_len) - self.repadding_state(states, max_target_len) - return states + input_ids = torch.cat(input_ids, dim=0) + audio_masks = torch.cat(audio_masks, dim=0) - def repadding_state(self, states: list[StepRequestState], max_new_target_len: int) -> list[StepRequestState]: - for state in states: - num_pads = max_new_target_len - state.latents.shape[-1] - state.latents = F.pad( - state.latents, - (0, num_pads), - value=self.config.audio_mask_id, - ) - state.extra["audio_mask"] = F.pad(state.extra["audio_mask"], (0, num_pads), value=False) - state.extra["attn_mask"] = F.pad(state.extra["attn_mask"], (0, num_pads, 0, num_pads), value=False) - return states + return input_ids, audio_masks, cond_lens def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: prompt = state.prompt if state.prompt else "" @@ -447,10 +376,10 @@ def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: return prepared prepared_request = prepared - target_lens = prepared_request.target_len + cond_len = prepared_request.cond_len + target_len = prepared_request.target_len input_ids = prepared_request.input_ids audio_mask = prepared_request.audio_mask - attn_mask = prepared_request.attention_mask seed = prepared_request.seed device = self.device mask_id = self.config.audio_mask_id @@ -461,16 +390,13 @@ def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: t_shift = self.t_shift # Initialize all target tokens as [MASK] - positions = torch.arange(target_lens, device=device).unsqueeze(0) - valid_target_mask = positions < torch.tensor(target_lens, device=device) - tokens = torch.zeros((1, num_codebooks, target_lens), dtype=torch.long, device=device) - tokens.masked_fill_(valid_target_mask.unsqueeze(1), mask_id) + tokens = torch.full((1, num_codebooks, target_len), mask_id, dtype=torch.long, device=device) timesteps = _get_time_steps(0.0, 1.0, num_step + 1, t_shift) # Compute unmasking schedule schedules = [] - total_mask = target_lens * num_codebooks + total_mask = target_len * num_codebooks rem = total_mask sched = [] for step in range(num_step): @@ -496,9 +422,9 @@ def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: state.extra["layer_ids"] = layer_ids state.extra["generator"] = generator state.extra["t_shift"] = t_shift - state.extra["target_len"] = target_lens + state.extra["cond_len"] = cond_len + state.extra["target_len"] = target_len state.extra["audio_mask"] = audio_mask - state.extra["attn_mask"] = attn_mask state.extra["tokens"] = tokens def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestState] | None = None, **kwargs: Any): @@ -506,83 +432,84 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS input_ids = input_batch.latents layer_ids = states[0].extra["layer_ids"] - batch_audio_mask: list[torch.Tesor] = [] - batch_attn_mask: list[torch.Tensor] = [] - batch_target_len: list[int] = [] + audio_masks: list[torch.Tensor] = [] + target_lens: list[int] = [] batch_tokens: list[torch.Tensor] = [] + cond_lens: list[int] = [] steps: list[int] = [] schedules: list[torch.Tensor] = [] generators: list[torch.Generator] = [] + guidance_scales: list[float] = [] for state in states: - batch_audio_mask.append(state.extra.get("audio_mask", None)) - batch_attn_mask.append(state.extra.get("attn_mask", None)) - guidance_scale = state.extra.get("guidance", self.guidance_scale) - batch_target_len.append(state.extra["target_len"]) + audio_masks.append(state.extra.get("audio_mask", None)) + cond_lens.append(state.extra["cond_len"]) + target_lens.append(state.extra["target_len"]) batch_tokens.append(state.extra["tokens"]) + guidance_scales.append(state.extra.get("guidance", self.guidance_scale)) generators.append(state.extra.get("generator", None)) schedules.append(state.extra["schedules"]) steps.append(state.step_index) - if len(batch_audio_mask) == 1: - batch_audio_mask = batch_audio_mask[0] - batch_attn_mask = batch_attn_mask[0] - else: - batch_audio_mask = torch.stack(batch_audio_mask, dim=0) - batch_attn_mask = torch.stack(batch_attn_mask, dim=0) - if len(batch_target_len) > 1: - batch_audio_mask = batch_audio_mask.reshape(-1, *batch_audio_mask.shape[2:]) - batch_attn_mask = batch_attn_mask.reshape(-1, *batch_attn_mask.shape[2:]) + audio_masks = torch.cat(audio_masks, dim=0) - B = len(batch_target_len) - target_lens = batch_target_len - input_ids = self._request_major_to_cfg_major(input_ids, B) - batch_audio_mask = self._request_major_to_cfg_major(batch_audio_mask, B) - batch_attn_mask = self._request_major_to_cfg_major(batch_attn_mask, B) + B = len(target_lens) + cu_seqs = _build_cu_seqs(cond_lens, target_lens, input_ids.device) mask_id = self.config.audio_mask_id position_temperature = self.position_temperature class_temperature = self.class_temperature layer_penalty_factor = self.layer_penalty_factor - c_lens = batch_attn_mask[:B, 0, 0].sum(dim=-1).tolist() - - # Materialize the SDPA float mask once so the captured graph (and eager path) skip per-layer conversion. - sdpa_attn_mask = torch.zeros_like(batch_attn_mask, dtype=torch.float32).masked_fill_( - ~batch_attn_mask, float("-inf") - ) if use_cuda_graph: # Float mask skips per-layer conversion; fp32 cast deferred to the per-item slices below. - batch_logits = self.generator._cuda_graph_fwd(input_ids, batch_audio_mask, sdpa_attn_mask) + batch_logits = self.generator._cuda_graph_fwd(input_ids, audio_masks, cu_seqs, B) else: # Recompute embeddings and RoPE from the current dynamically # padded/reordered batch. - inputs_embeds = self.generator._prepare_embeddings(input_ids, batch_audio_mask) - hidden_states = self.generator._transformer_forward(inputs_embeds, sdpa_attn_mask) + inputs_embeds = self.generator._prepare_embeddings(input_ids, audio_masks) + hidden_states = self.generator._transformer_forward(inputs_embeds, cu_seqs) # fp32 cast deferred to the per-item slices below. batch_logits = self.generator._get_logits(hidden_states) - # batch_logits: [2*B, 8, S, 1025] + # batch_logits: [8, total_seq_len, 1025] + + target_offsets: list[int] = [] + target_offset = 0 + for target_len in target_lens: + target_offsets.append(target_offset) + target_offset += target_len + + sequence_offsets: list[int] = [] + sequence_offset = 0 + for cond_len, target_len in zip(cond_lens, target_lens): + sequence_offsets.append(sequence_offset) + sequence_offset += cond_len + target_len for i in range(B): k = schedules[i][steps[i]] if k <= 0: continue - c_len = c_lens[i] + c_len = cond_lens[i] t_len = target_lens[i] # Extract logits for target region; upcast only the slices we actually consume. - c_logits = batch_logits[i : i + 1, :, c_len - t_len : c_len, :].to(torch.float32) - u_logits = batch_logits[B + i : B + i + 1, :, :t_len, :].to(torch.float32) + request_start = sequence_offsets[i] + cond_end = request_start + c_len + uncond_start = cond_end + + # Extract logits for target region; upcast only the slices we actually consume. + c_logits = batch_logits[:, cond_end - t_len : cond_end, :].unsqueeze(0).to(torch.float32) + u_logits = batch_logits[:, uncond_start : uncond_start + t_len, :].unsqueeze(0).to(torch.float32) # Classifier-free guidance. Fuse the chain: the two inner # log_softmax normalizers are per-position scalars that the final # shift-invariant log_softmax cancels, so guide on the raw logits # with a single softmax: log_softmax((1+s)*c - s*u). Exact. - if guidance_scale != 0: + if guidance_scales[i] != 0: log_probs = F.log_softmax( - (1.0 + guidance_scale) * c_logits - guidance_scale * u_logits, + (1.0 + guidance_scales[i]) * c_logits - guidance_scales[i] * u_logits, dim=-1, ) else: @@ -620,10 +547,15 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS states[i].extra["tokens"] = sample_tokens # Mirror update into both cond and uncond input_ids halves for the next step. - input_ids[i, :, c_len - t_len : c_len] = sample_tokens.squeeze(0) - input_ids[B + i, :, :t_len] = sample_tokens.squeeze(0) + packed_sample_tokens = sample_tokens.squeeze(0).transpose(0, 1) + input_ids[cond_end - t_len : cond_end] = packed_sample_tokens + input_ids[uncond_start : uncond_start + t_len] = packed_sample_tokens - return self._cfg_major_to_request_major(input_ids, B) + # InputBatch reuses its latents buffer across steps. Returning that + # same storage would make the Runner persist per-request views into the + # cached destination; the next make_batch() would then copy overlapping + # source/destination slices. Break the alias at the lifecycle boundary. + return input_ids.clone() def step_scheduler(self, state: StepRequestState, noise_pred: torch.Tensor, **kwargs: Any): state.latents = noise_pred @@ -654,13 +586,13 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: prepared_requests.append(prepared) batch_target_len = [request.target_len for request in prepared_requests] - seed = prepared_requests[-1].seed - batch_input_ids, batch_audio_mask, batch_attn_mask = self._collate_request_inputs(prepared_requests) + batch_seeds = [request.seed for request in prepared_requests] + batch_input_ids, batch_audio_mask, batch_cond_lens = self._collate_request_inputs(prepared_requests) # Run 32-step iterative unmasking tokens = self.generator( input_ids=batch_input_ids, audio_mask=batch_audio_mask, - attention_mask=batch_attn_mask, + cond_lens=batch_cond_lens, target_lens=batch_target_len, num_step=self.num_step, guidance_scale=self.guidance_scale, @@ -668,11 +600,17 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: layer_penalty_factor=self.layer_penalty_factor, position_temperature=self.position_temperature, class_temperature=self.class_temperature, - seed=seed, + seed=batch_seeds, ) - audio = self.decoder(tokens, batch_target_len) # [B, 1, max_target_len * 960] - return self._split_audio_outputs(audio, batch_target_len) + outputs: list[DiffusionOutput] = [] + target_offset = 0 + for target_len in batch_target_len: + request_tokens = tokens[:, :, target_offset : target_offset + target_len] + audio = self.decoder(request_tokens, [target_len]) + outputs.append(DiffusionOutput(output=audio)) + target_offset += target_len + return outputs def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: """Load weights from model directory (not from the iterator). diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py index 01bf3de734e..d680de908b3 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py @@ -6,9 +6,8 @@ Generates 8-codebook audio tokens from text via 32-step non-autoregressive iterative masked prediction with classifier-free guidance. -Uses full bidirectional attention computed directly with PyTorch SDPA -(torch.nn.functional.scaled_dot_product_attention); no auto-selected -FlashAttention/SageAttention/DiffusionAttention backend is used. +Uses backend-dispatched variable-length full bidirectional attention over +packed conditional and unconditional sequences, with an SDPA mask fallback. """ from __future__ import annotations @@ -22,7 +21,11 @@ import torch.nn as nn import torch.nn.functional as F from vllm.logger import init_logger +from vllm.utils.math_utils import round_up +from vllm_omni.diffusion.attention.backends.abstract import AttentionMetadata +from vllm_omni.diffusion.attention.layer import Attention +from vllm_omni.diffusion.data import OmniDiffusionConfig from vllm_omni.model_executor.models.omnivoice.fused_qkv_rope import fused_qkv_norm_rope from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig @@ -261,9 +264,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class OmniVoiceAttention(nn.Module): - """Qwen3-style GQA attention using PyTorch SDPA (full bidirectional).""" + """Qwen3-style GQA using packed, full-bidirectional varlen attention.""" - def __init__(self, config: OmniVoiceConfig): + def __init__(self, config: OmniVoiceConfig, layer_idx: int): super().__init__() self.hidden_size = config.llm_hidden_size self.num_heads = config.llm_num_attention_heads @@ -282,16 +285,27 @@ def __init__(self, config: OmniVoiceConfig): self.k_norm = OmniVoiceRMSNorm(self.head_dim) self.scale = 1.0 / math.sqrt(self.head_dim) + self.attention_op = Attention( + num_heads=self.num_heads, + num_kv_heads=self.num_heads, + head_size=self.head_dim, + causal=False, + softmax_scale=self.scale, + prefix=f"layers.{layer_idx}.self_attn.attention_op", + qkv_layout="BSND", + skip_sequence_parallel=True, + ) + self.use_packed_varlen = self.attention_op.attn_backend.supports_multi_doc_packed_varlen() def forward( self, hidden_states: torch.Tensor, rope_table: torch.Tensor, - attention_mask: torch.Tensor | None = None, + attn_metadata: AttentionMetadata, ) -> torch.Tensor: - batch_size, seq_len, _ = hidden_states.shape + seq_len, _ = hidden_states.shape - qkv = self.qkv_proj(hidden_states).view(batch_size, seq_len, self.num_qkv_heads, self.head_dim) + qkv = self.qkv_proj(hidden_states).view(1, seq_len, self.num_qkv_heads, self.head_dim) # One kernel for the whole prologue: split the packed projection, RMSNorm # Q and K per head, rotate both, broadcast K and V across their query @@ -306,22 +320,16 @@ def forward( self.num_kv_heads, ) - # Caller passes a float mask; materialize float form if a bool slips through. - sdpa_mask = attention_mask - if sdpa_mask is not None and sdpa_mask.dtype == torch.bool: - sdpa_mask = torch.zeros_like(attention_mask, dtype=q.dtype).masked_fill_(~attention_mask, float("-inf")) - - out = F.scaled_dot_product_attention( - q, - k, - v, - attn_mask=sdpa_mask, - scale=self.scale, - ) + # The fused prologue emits [B, N, S, D]; Omni attention backends use + # [B, S, N, D]. Keep the conversion outside the backend. + q, k, v = (tensor.permute(0, 2, 1, 3).contiguous().to(torch.bfloat16) for tensor in (q, k, v)) + if self.use_packed_varlen: + out = self.attention_op(q, k, v, attn_metadata) + else: + out = self.attention_op.sdpa_fallback.forward(q, k, v, attn_metadata) - # Back to (batch, seq, heads * head_dim) - out = out.permute(0, 2, 1, 3).contiguous() - out = out.view(batch_size, seq_len, self.num_heads * self.head_dim) + out = out.squeeze(0).to(hidden_states.dtype) + out = out.reshape(seq_len, self.num_heads * self.head_dim) return self.o_proj(out) @@ -351,19 +359,19 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class OmniVoiceTransformerBlock(nn.Module): - """Single Qwen3 transformer block with PyTorch SDPA attention.""" + """Single Qwen3 transformer block with variable-length attention.""" - def __init__(self, config: OmniVoiceConfig): + def __init__(self, config: OmniVoiceConfig, layer_idx: int): super().__init__() self.input_layernorm = OmniVoiceRMSNorm(config.llm_hidden_size, eps=config.llm_rms_norm_eps) - self.self_attn = OmniVoiceAttention(config) + self.self_attn = OmniVoiceAttention(config, layer_idx) self.post_attention_layernorm = OmniVoiceRMSNorm(config.llm_hidden_size, eps=config.llm_rms_norm_eps) self.mlp = OmniVoiceMLP(config) def forward( self, hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, + attn_metadata: AttentionMetadata, rope_table: torch.Tensor | None = None, residual: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: @@ -391,7 +399,7 @@ def forward( residual = hidden_states hidden_states = self.input_layernorm(hidden_states) - hidden_states = self.self_attn(hidden_states, rope_table, attention_mask=attention_mask) + hidden_states = self.self_attn(hidden_states, rope_table, attn_metadata) if _TRITON_AVAILABLE: # Fused: (attn_out + residual) + RMSNorm in one kernel @@ -432,6 +440,63 @@ def _precompute_rope_table( return torch.cat([freqs.cos(), freqs.sin()], dim=-1) +def _position_ids_from_cu_seqs(cu_seqs: torch.Tensor, seq_len: int) -> torch.Tensor: + """Return each packed token's zero-based position within its sequence.""" + token_ids = torch.arange(seq_len, device=cu_seqs.device, dtype=cu_seqs.dtype) + sequence_ids = torch.searchsorted(cu_seqs[1:], token_ids, right=True) + sequence_starts = cu_seqs.index_select(0, sequence_ids.to(torch.long)) + return (token_ids - sequence_starts).to(torch.long) + + +def _attention_metadata_from_cu_seqs( + cu_seqs: torch.Tensor, + seq_len: int, + *, + needs_sdpa_mask: bool, + max_seqlen: int | None = None, +) -> AttentionMetadata: + """Build packed-varlen metadata and, when needed, an SDPA block mask.""" + # Local benchmarks found no measurable performance difference between the + # exact longest segment and the packed total. Eager uses the exact bound; + # CUDA Graph uses the fixed token-bucket maximum. + kernel_max_seqlen = seq_len if max_seqlen is None else max_seqlen + extra = { + "cu_seqlens_q": cu_seqs, + "cu_seqlens_k": cu_seqs, + "max_seqlen_q": kernel_max_seqlen, + "max_seqlen_k": kernel_max_seqlen, + } + if not needs_sdpa_mask: + return AttentionMetadata(extra=extra) + + token_ids = torch.arange(seq_len, device=cu_seqs.device, dtype=cu_seqs.dtype) + sequence_ids = torch.searchsorted(cu_seqs[1:], token_ids, right=True) + attn_mask = sequence_ids.view(1, 1, seq_len, 1) == sequence_ids.view(1, 1, 1, seq_len) + return AttentionMetadata(attn_mask=attn_mask, extra=extra) + + +def _build_cu_seqs( + cond_lens: list[int], + uncond_lens: list[int], + device: torch.device, + *, + tail_end: int | None = None, +) -> torch.Tensor: + """Build request-major [cond0, uncond0, ...] cumulative offsets.""" + if len(cond_lens) != len(uncond_lens): + raise ValueError(f"Mismatched cond/uncond lengths: {len(cond_lens)} != {len(uncond_lens)}.") + offsets = [0] + for cond_len, uncond_len in zip(cond_lens, uncond_lens): + offsets.append(offsets[-1] + cond_len) + offsets.append(offsets[-1] + uncond_len) + if tail_end is None: + tail_end = offsets[-1] + if offsets[-1] > tail_end: + raise ValueError(f"Packed length {offsets[-1]} exceeds tail end {tail_end}.") + offsets.append(tail_end) + return torch.tensor(offsets, device=device, dtype=torch.int32) + + # --------------------------------------------------------------------------- # TF32 opt-in (process-wide; default off) # --------------------------------------------------------------------------- @@ -461,23 +526,6 @@ def _maybe_enable_tf32() -> None: # --------------------------------------------------------------------------- -def _additive_float_mask(mask: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: - """Convert a boolean attention mask to its additive float form. - - ``True`` (attend) maps to ``0.0`` and ``False`` (masked) to ``-inf``. A bool - mask must never be copied straight into a float buffer: the implicit cast - maps True/False to 1.0/0.0, which leaves masked positions at 0.0 and so - silently *unmasks* them. - - ``dtype`` is required rather than defaulting to float32: SDPA rejects an - additive mask whose dtype differs from the query, so the mask has to follow - the model dtype and a default would just hide that coupling. - """ - if mask.dtype != torch.bool: - return mask - return torch.zeros_like(mask, dtype=dtype).masked_fill_(~mask, float("-inf")) - - class _OmniVoiceCUDAGraphForward: """Pre-captures CUDA graphs for predefined sequence-length buckets. @@ -486,24 +534,47 @@ class _OmniVoiceCUDAGraphForward: (one step at a time) means pool sharing is safe. """ - # Default bucket count is 10; 16 gives modest headroom for edge cases - # (seq_len > max bucket or non-CFG batch) without unbounded GPU growth. - _MAX_LAZY_GRAPHS: int = 16 + _MAX_LAZY_GRAPHS = 16 + _STATIC_CAPTURE_TOKEN_LIMIT = 1024 + _LAZY_CAPTURE_ALIGNMENT = 128 def __init__(self, generator: OmniVoiceGenerator, capture_sizes: list[int]) -> None: self._gen = generator - self._capture_sizes = sorted(capture_sizes) - # Pre-warmed graphs keyed by (two_b, bucket); fixed set, never evicted. + self.capture_batch_sizes = self._derive_capture_batch_size() + self.capture_bucket_sizes_by_batch = self._derive_capture_bucket_sizes(capture_sizes) self._graphs: dict[tuple[int, int], dict] = {} - # Lazy-captured graphs for oversized / non-CFG shapes; capped via LRU. self._lazy_graphs: OrderedDict[tuple[int, int], dict] = OrderedDict() self._lock = threading.Lock() # Per-instance pool handle: isolates OmniVoice CUDA memory from other # vllm modules while still allowing safe re-use across sequential replays. self._pool_handle: int | None = None - def _find_bucket(self, seq_len: int) -> int | None: - for bucket in self._capture_sizes: + def _derive_capture_batch_size(self) -> list[int]: + return list(range(1, self._gen.od_config.max_num_seqs + 1)) + + def _derive_capture_bucket_sizes(self, capture_sizes: list[int]) -> dict[int, list[int]]: + """Build a triangular, batch-aware token-bucket capture plan.""" + base_sizes = sorted(set(capture_sizes)) + if not base_sizes: + return {batch_size: [] for batch_size in self.capture_batch_sizes} + alignment = base_sizes[0] + single_request_cap = min(base_sizes[-1], 512) + max_graph_tokens = min(base_sizes[-1], self._STATIC_CAPTURE_TOKEN_LIMIT) + growth_per_batch = max(alignment, single_request_cap // 2) + candidates = sorted(set(base_sizes) | set(range(alignment, max_graph_tokens + alignment, alignment))) + return { + batch_size: [ + bucket + for bucket in candidates + if alignment * batch_size + <= bucket + <= min(single_request_cap + (batch_size - 1) * growth_per_batch, max_graph_tokens) + ] + for batch_size in self.capture_batch_sizes + } + + def _find_bucket(self, batch_size: int, seq_len: int) -> int | None: + for bucket in self.capture_bucket_sizes_by_batch.get(batch_size, ()): if bucket >= seq_len: return bucket return None @@ -512,57 +583,36 @@ def _pad_inputs( self, input_ids: torch.Tensor, audio_mask: torch.Tensor, - attention_mask: torch.Tensor | None, bucket: int, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: - S = input_ids.shape[-1] - if S == bucket: - return input_ids, audio_mask, attention_mask - - two_b = input_ids.shape[0] - num_cb = input_ids.shape[1] - - ids_padded = torch.zeros(two_b, num_cb, bucket, dtype=input_ids.dtype, device=input_ids.device) - ids_padded[:, :, :S] = input_ids - - mask_padded = torch.zeros(two_b, bucket, dtype=torch.bool, device=audio_mask.device) - mask_padded[:, :S] = audio_mask - - if attention_mask is not None: - # Callers normalize to the additive float form first, so pad with -inf. - attn_padded = torch.full( - (two_b, 1, bucket, bucket), - float("-inf"), - dtype=attention_mask.dtype, - device=attention_mask.device, - ) - attn_padded[:, :, :S, :S] = attention_mask - else: - attn_padded = None - - return ids_padded, mask_padded, attn_padded + ) -> tuple[torch.Tensor, torch.Tensor]: + seq_len = input_ids.shape[0] + if seq_len == bucket: + return input_ids, audio_mask + return ( + F.pad(input_ids, (0, 0, 0, bucket - seq_len), value=0), + F.pad(audio_mask, (0, bucket - seq_len), value=False), + ) def _capture_for_key( self, key: tuple[int, int], input_ids: torch.Tensor, audio_mask: torch.Tensor, - attention_mask: torch.Tensor | None, + cu_seqs: torch.Tensor, ) -> dict: _, bucket = key device = input_ids.device - static_rope_table = self._gen._rope_table_for(bucket, device, self._gen.model_dtype) - static_input_ids = input_ids.clone() static_audio_mask = audio_mask.clone() - static_attn_mask = attention_mask.clone() if attention_mask is not None else None + static_cu_seqs = cu_seqs.clone() + static_rope_table = self._gen._rope_table_for(bucket, device, self._gen.model_dtype) with torch.no_grad(): _ = self._gen._step_forward( static_input_ids, static_audio_mask, - static_attn_mask, + static_cu_seqs, static_rope_table, ) torch.accelerator.synchronize(device) @@ -580,7 +630,7 @@ def _capture_for_key( static_output = self._gen._step_forward( static_input_ids, static_audio_mask, - static_attn_mask, + static_cu_seqs, static_rope_table, ) @@ -588,87 +638,83 @@ def _capture_for_key( "graph": graph, "static_input_ids": static_input_ids, "static_audio_mask": static_audio_mask, - "static_attn_mask": static_attn_mask, + "static_cu_seqs": static_cu_seqs, "static_rope_table": static_rope_table, "static_output": static_output, } logger.info("OmniVoice CUDA Graph captured for key %s", key) return entry + def make_capture_cu_seq(self, batch_size: int, bucket_size: int, device: torch.device) -> torch.Tensor: + num_real_sequences = 2 * batch_size + base, rem = divmod(bucket_size, num_real_sequences) + lengths = torch.full((num_real_sequences,), base, dtype=torch.int32, device=device) + lengths[:rem] += 1 + cu_seq = torch.empty(num_real_sequences + 2, dtype=torch.int32, device=device) + cu_seq[0] = 0 + cu_seq[1:-1] = lengths.cumsum(0) + cu_seq[-1] = bucket_size + return cu_seq + def warmup(self, device: torch.device) -> None: - """Pre-capture graphs for all bucket sizes with B=1 (two_b=2 for CFG).""" + """Pre-capture common request batch sizes for every token bucket.""" if not torch.cuda.is_available(): return logger.info( - "OmniVoice CUDA Graph warmup: capturing %d bucket sizes %s", - len(self._capture_sizes), - self._capture_sizes, + "OmniVoice CUDA Graph warmup: batch-aware capture plan %s", + self.capture_bucket_sizes_by_batch, ) - two_b = 2 num_cb = self._gen.config.num_audio_codebook - for bucket in self._capture_sizes: - key = (two_b, bucket) - dummy_ids = torch.zeros(two_b, num_cb, bucket, dtype=torch.long, device=device) - dummy_mask = torch.zeros(two_b, bucket, dtype=torch.bool, device=device) - # Capture with a float mask to match what forward() feeds at replay time, - # in the model dtype so replay can copy_ into it without a cast. - dummy_attn = torch.zeros(two_b, 1, bucket, bucket, dtype=self._gen.model_dtype, device=device) - self._graphs[key] = self._capture_for_key(key, dummy_ids, dummy_mask, dummy_attn) + for batch_size in self.capture_batch_sizes: + for bucket in self.capture_bucket_sizes_by_batch[batch_size]: + key = (batch_size, bucket) + dummy_ids = torch.zeros(bucket, num_cb, dtype=torch.long, device=device) + dummy_mask = torch.zeros(bucket, dtype=torch.bool, device=device) + dummy_cu_seqs = self.make_capture_cu_seq(batch_size, bucket, device) + self._graphs[key] = self._capture_for_key(key, dummy_ids, dummy_mask, dummy_cu_seqs) logger.info("OmniVoice CUDA Graph warmup complete (%d graphs)", len(self._graphs)) def __call__( self, input_ids: torch.Tensor, audio_mask: torch.Tensor, - attention_mask: torch.Tensor | None, + cu_seqs: torch.Tensor, + batch_size: int, ) -> torch.Tensor: if torch.cuda.is_current_stream_capturing(): - rope_table = self._gen._rope_table_for(input_ids.shape[-1], input_ids.device, self._gen.model_dtype) - return self._gen._step_forward(input_ids, audio_mask, attention_mask, rope_table) - - seq_len = input_ids.shape[-1] - two_b = input_ids.shape[0] - bucket = self._find_bucket(seq_len) if two_b == 2 else None - - # Graphs are captured with (and their static buffers hold) the additive - # float mask, so normalize here, before padding or any copy_ into them. - if attention_mask is not None: - attention_mask = _additive_float_mask(attention_mask, self._gen.model_dtype) + rope_table = self._gen._rope_table_for(input_ids.shape[0], input_ids.device, self._gen.model_dtype) + return self._gen._step_forward(input_ids, audio_mask, cu_seqs, rope_table) + seq_len = input_ids.shape[0] + bucket = self._find_bucket(batch_size, seq_len) + is_lazy = bucket is None if bucket is None: - # Lazy capture: oversized sequence or non-unit batch (no pre-warmed bucket). - # Lock prevents concurrent threads from double-capturing the same key. - # _lazy_graphs is capped at _MAX_LAZY_GRAPHS with LRU eviction to - # prevent unbounded GPU memory growth when seq_len varies widely. - key = (two_b, seq_len) - ids_in, mask_in, attn_in = input_ids, audio_mask, attention_mask - with self._lock: - entry = self._lazy_graphs.get(key) - if entry is None: - entry = self._capture_for_key(key, ids_in, mask_in, attn_in) - if len(self._lazy_graphs) >= self._MAX_LAZY_GRAPHS: - evicted_key, _ = self._lazy_graphs.popitem(last=False) - logger.warning("OmniVoice CUDA Graph lazy cache full; evicted key %s", evicted_key) - self._lazy_graphs[key] = entry - else: - key = (two_b, bucket) - ids_in, mask_in, attn_in = self._pad_inputs(input_ids, audio_mask, attention_mask, bucket) - with self._lock: - entry = self._graphs.get(key) - if entry is None: - entry = self._capture_for_key(key, ids_in, mask_in, attn_in) - self._graphs[key] = entry + bucket = round_up(seq_len, self._LAZY_CAPTURE_ALIGNMENT) + ids_in, mask_in = self._pad_inputs(input_ids, audio_mask, bucket) + runtime_cu_seqs = cu_seqs.clone() + runtime_cu_seqs[-1] = bucket + key = (batch_size, bucket) + cache = self._lazy_graphs if is_lazy else self._graphs + with self._lock: + entry = cache.get(key) + if entry is None: + if is_lazy and len(self._lazy_graphs) >= self._MAX_LAZY_GRAPHS: + evicted_key, _ = self._lazy_graphs.popitem(last=False) + logger.info("Evicted OmniVoice lazy CUDA Graph key %s", evicted_key) + entry = self._capture_for_key(key, ids_in, mask_in, runtime_cu_seqs) + cache[key] = entry + elif is_lazy: + self._lazy_graphs.move_to_end(key) entry["static_input_ids"].copy_(ids_in) entry["static_audio_mask"].copy_(mask_in) - if attn_in is not None and entry["static_attn_mask"] is not None: - entry["static_attn_mask"].copy_(attn_in) + entry["static_cu_seqs"].copy_(runtime_cu_seqs) entry["graph"].replay() output = entry["static_output"] - if bucket is not None and bucket != seq_len: - output = output[:, :, :seq_len, :] + if bucket != seq_len: + output = output[:, :seq_len, :] return output def clear(self) -> None: @@ -692,17 +738,17 @@ class OmniVoiceGenerator(nn.Module): - 32-step iterative unmasking with classifier-free guidance Optimizations: - - Full bidirectional attention via PyTorch SDPA (no auto-selected - FlashAttn/SageAttn/DiffusionAttention backend) + - Packed full-bidirectional varlen attention with SDPA fallback - regionally_compile() compatible for torch.compile on repeated blocks """ # For regionally_compile() support _repeated_blocks = ["layers"] - def __init__(self, config: OmniVoiceConfig): + def __init__(self, config: OmniVoiceConfig, od_config: OmniDiffusionConfig): super().__init__() self.config = config + self.od_config = od_config # Opt-in TF32; must run before any CUDA-graph capture so captured kernels honour it. if getattr(config, "enable_tf32", False): @@ -722,7 +768,10 @@ def __init__(self, config: OmniVoiceConfig): ) # Transformer layers - self.layers = nn.ModuleList([OmniVoiceTransformerBlock(config) for _ in range(config.llm_num_hidden_layers)]) + self.layers = nn.ModuleList( + [OmniVoiceTransformerBlock(config, layer_idx) for layer_idx in range(config.llm_num_hidden_layers)] + ) + self._needs_sdpa_mask = any(not layer.self_attn.use_packed_varlen for layer in self.layers) self.norm = OmniVoiceRMSNorm(config.llm_hidden_size, eps=config.llm_rms_norm_eps) # Prediction head: hidden → 8 * 1025 @@ -775,23 +824,23 @@ def _prepare_embeddings( """Prepare mixed text+audio embeddings. Args: - input_ids: [B, 8, S] - text tokens replicated across codebooks, + input_ids: [T, 8] - text tokens replicated across codebooks, audio positions have per-codebook token IDs - audio_mask: [B, S] - True for audio positions, False for text - text_embeds: optional cached [B, S, H] text-position embeddings - audio_mask_3d: optional cached [B, S, 1] audio_mask.unsqueeze(-1) + audio_mask: [T] - True for audio positions, False for text + text_embeds: optional cached [T, H] text-position embeddings + audio_mask_3d: optional cached [T, 1] audio_mask.unsqueeze(-1) Returns: - embeddings: [B, S, hidden_size] + embeddings: [T, hidden_size] """ # Cached across the denoising loop since text ids don't change. if text_embeds is None: - text_embeds = self.text_embedding(input_ids[:, 0, :]) + text_embeds = self.text_embedding(input_ids[:, 0]) if audio_mask_3d is None: audio_mask_3d = audio_mask.unsqueeze(-1) # Audio embeddings: offset per codebook, then sum across codebooks - shifted_ids = (input_ids * audio_mask.unsqueeze(1)) + self.codebook_layer_offsets.view(1, -1, 1) + shifted_ids = (input_ids * audio_mask.unsqueeze(1)) + self.codebook_layer_offsets.view(1, -1) audio_embeds = self.audio_embeddings(shifted_ids).sum(dim=1) # Merge: audio where audio_mask=True, text elsewhere @@ -800,35 +849,39 @@ def _prepare_embeddings( def _transformer_forward( self, inputs_embeds: torch.Tensor, - attention_mask: torch.Tensor | None = None, + cu_seqs: torch.Tensor, + max_seqlen: int | None = None, rope_table: torch.Tensor | None = None, ) -> torch.Tensor: """Run through transformer layers. Args: - inputs_embeds: [B, S, hidden_size] - attention_mask: [B, 1, S, S] or None - rope_table: optional precomputed [B * S, head_dim] RoPE table + inputs_embeds: [T, hidden_size] + cu_seqs: cumulative boundaries for packed sequences + rope_table: optional base [T, head_dim] RoPE table Returns: - hidden_states: [B, S, hidden_size] + hidden_states: [T, hidden_size] """ hidden_states = inputs_embeds if rope_table is None: - rope_table = self._rope_table_for(inputs_embeds.shape[1], inputs_embeds.device, hidden_states.dtype) - - # Safety: convert bool mask if caller hasn't (e.g. external paths beyond forward()). - if attention_mask is not None and attention_mask.dtype == torch.bool: - attention_mask = torch.zeros_like(attention_mask, dtype=hidden_states.dtype).masked_fill_( - ~attention_mask, float("-inf") - ) + rope_table = self._rope_table_for(inputs_embeds.shape[0], inputs_embeds.device, hidden_states.dtype) + seq_len = inputs_embeds.shape[0] + position_ids = _position_ids_from_cu_seqs(cu_seqs, seq_len) + packed_rope_table = rope_table.index_select(0, position_ids).contiguous() + attn_metadata = _attention_metadata_from_cu_seqs( + cu_seqs, + seq_len, + needs_sdpa_mask=self._needs_sdpa_mask, + max_seqlen=max_seqlen, + ) residual = None for layer in self.layers: hidden_states, residual = layer( hidden_states, - attention_mask=attention_mask, - rope_table=rope_table, + attn_metadata=attn_metadata, + rope_table=packed_rope_table, residual=residual, ) @@ -838,44 +891,39 @@ def _get_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: """Project hidden states to per-codebook logits. Args: - hidden_states: [B, S, hidden_size] + hidden_states: [T, hidden_size] Returns: - logits: [B, 8, S, 1025] + logits: [8, T, 1025] """ - batch_size, seq_len, _ = hidden_states.shape - logits_flat = self.audio_heads(hidden_states) # [B, S, 8*1025] + seq_len, _ = hidden_states.shape + logits_flat = self.audio_heads(hidden_states) # [T, 8*1025] return logits_flat.view( - batch_size, seq_len, self.config.num_audio_codebook, self.config.audio_vocab_size, - ).permute(0, 2, 1, 3) # [B, 8, S, 1025] + ).permute(1, 0, 2) # [8, T, 1025] def _step_forward( self, input_ids: torch.Tensor, audio_mask: torch.Tensor, - attention_mask: torch.Tensor | None, + cu_seqs: torch.Tensor, rope_table: torch.Tensor, ) -> torch.Tensor: """Single unmasking-step forward using a pre-cast RoPE table (CUDA graph safe).""" hidden_states = self._prepare_embeddings(input_ids, audio_mask) - residual = None - for layer in self.layers: - hidden_states, residual = layer( - hidden_states, attention_mask=attention_mask, rope_table=rope_table, residual=residual - ) - return self._get_logits(self.norm(hidden_states + residual)) + hidden_states = self._transformer_forward(hidden_states, cu_seqs, rope_table=rope_table) + return self._get_logits(hidden_states) @torch.inference_mode() def forward( self, input_ids: torch.Tensor, audio_mask: torch.Tensor, - attention_mask: torch.Tensor, + cond_lens: list[int], target_lens: list[int], - seed: int | None = None, + seed: int | list[int | None] | None = None, num_step: int = 32, guidance_scale: float = 2.0, t_shift: float = 0.1, @@ -886,9 +934,9 @@ def forward( """Run the full 32-step iterative unmasking generation. Args: - input_ids: [2*B, 8, S] - conditional (0:B) + unconditional (B:2B) - audio_mask: [2*B, S] - True for audio positions - attention_mask: [2*B, 1, S, S] - attention mask + input_ids: [T, 8] packed as [cond0, uncond0, ...] + audio_mask: [T] - True for audio positions + cond_lens: conditional lengths per request target_lens: List of target audio lengths per batch item num_step: Number of unmasking steps guidance_scale: CFG scale @@ -898,22 +946,33 @@ def forward( class_temperature: Temperature for token prediction (0=greedy) Returns: - tokens: [B, 8, max_target_len] - generated audio tokens + tokens: [1, 8, sum(target_lens)] - generated audio tokens """ B = len(target_lens) device = input_ids.device - max_target_len = max(target_lens) + total_target_lens = sum(target_lens) mask_id = self.config.audio_mask_id num_codebooks = self.config.num_audio_codebook - if seed is None: - seed = random.randint(0, 2**63 - 1) - generator = torch.Generator(device=device).manual_seed(seed) + seeds = seed if isinstance(seed, list) else [seed] * B + generators = [ + torch.Generator(device=device).manual_seed( + request_seed if request_seed is not None else random.randint(0, 2**63 - 1) + ) + for request_seed in seeds + ] # Initialize all target tokens as [MASK] - positions = torch.arange(max_target_len, device=device).unsqueeze(0) - valid_target_mask = positions < torch.tensor(target_lens, device=device).unsqueeze(1) - tokens = torch.zeros((B, num_codebooks, max_target_len), dtype=torch.long, device=device) - tokens.masked_fill_(valid_target_mask.unsqueeze(1), mask_id) + tokens = torch.full((1, num_codebooks, total_target_lens), mask_id, dtype=torch.long, device=device) + target_offsets: list[int] = [] + target_offset = 0 + sequence_offsets: list[int] = [] + sequence_offset = 0 + for cond_len, target_len in zip(cond_lens, target_lens): + target_offsets.append(target_offset) + target_offset += target_len + sequence_offsets.append(sequence_offset) + sequence_offset += cond_len + target_len + cu_seqs = _build_cu_seqs(cond_lens, target_lens, device) # Compute unmasking schedule timesteps = _get_time_steps(0.0, 1.0, num_step + 1, t_shift).tolist() @@ -937,45 +996,46 @@ def forward( layer_ids = torch.arange(num_codebooks, device=device).view(1, -1, 1) - # Single D2H pull for all conditional lengths instead of B per-item .item() syncs. - c_lens = attention_mask[:B, 0, 0].sum(dim=-1).tolist() - - # Materialize the SDPA float mask once so the captured graph (and eager path) skip per-layer conversion. - sdpa_attn_mask = _additive_float_mask(attention_mask, self.model_dtype) - use_cuda_graph = self._cuda_graph_fwd is not None and input_ids.is_cuda if not use_cuda_graph: # Eager-path-only constants (the cuda-graph captures its own). - text_embeds_cached = self.text_embedding(input_ids[:, 0, :]) + text_embeds_cached = self.text_embedding(input_ids[:, 0]) audio_mask_3d = audio_mask.unsqueeze(-1) - rope_table = self._rope_table_for(input_ids.shape[-1], device, text_embeds_cached.dtype) + rope_table = self._rope_table_for(input_ids.shape[0], device, text_embeds_cached.dtype) # Main iterative loop for step in range(num_step): if use_cuda_graph: # Float mask skips per-layer conversion; fp32 cast deferred to the per-item slices below. - batch_logits = self._cuda_graph_fwd(input_ids, audio_mask, sdpa_attn_mask) + batch_logits = self._cuda_graph_fwd(input_ids, audio_mask, cu_seqs, B) else: # Eager fallback reuses hoisted constants (text embeds, sdpa mask, rope table). inputs_embeds = self._prepare_embeddings( input_ids, audio_mask, text_embeds=text_embeds_cached, audio_mask_3d=audio_mask_3d ) - hidden_states = self._transformer_forward(inputs_embeds, sdpa_attn_mask, rope_table=rope_table) + hidden_states = self._transformer_forward( + inputs_embeds, + cu_seqs, + max_seqlen=max(cond_lens), + rope_table=rope_table, + ) # fp32 cast deferred to the per-item slices below. batch_logits = self._get_logits(hidden_states) - # batch_logits: [2*B, 8, S, 1025] + # batch_logits: [8, T, 1025] for i in range(B): k = schedules[i][step] if k <= 0: continue - c_len = c_lens[i] + c_len = cond_lens[i] t_len = target_lens[i] + request_start = sequence_offsets[i] + cond_end = request_start + c_len # Extract logits for target region; upcast only the slices we actually consume. - c_logits = batch_logits[i : i + 1, :, c_len - t_len : c_len, :].to(torch.float32) - u_logits = batch_logits[B + i : B + i + 1, :, :t_len, :].to(torch.float32) + c_logits = batch_logits[:, cond_end - t_len : cond_end, :].unsqueeze(0).to(torch.float32) + u_logits = batch_logits[:, cond_end : cond_end + t_len, :].unsqueeze(0).to(torch.float32) # Classifier-free guidance. Fuse the chain: the two inner # log_softmax normalizers are per-position scalars that the final @@ -994,7 +1054,7 @@ def forward( # Token prediction if class_temperature > 0.0: - pred_tokens = _gumbel_sample(log_probs, class_temperature, generator).argmax(dim=-1) + pred_tokens = _gumbel_sample(log_probs, class_temperature, generators[i]).argmax(dim=-1) else: pred_tokens = log_probs.argmax(dim=-1) # [1, 8, T] @@ -1006,10 +1066,11 @@ def forward( # Gumbel noise for position selection if position_temperature > 0.0: - scores = _gumbel_sample(scores, position_temperature, generator) + scores = _gumbel_sample(scores, position_temperature, generators[i]) # Mask out already unmasked positions - sample_tokens = tokens[i : i + 1, :, :t_len] + target_start = target_offsets[i] + sample_tokens = tokens[:, :, target_start : target_start + t_len] scores.masked_fill_(sample_tokens != mask_id, -float("inf")) # Select top-k positions to unmask. .flatten() on this non-contiguous view already copies. @@ -1019,8 +1080,9 @@ def forward( sample_tokens.copy_(flat_tokens.view_as(sample_tokens)) # Mirror update into both cond and uncond input_ids halves for the next step. - input_ids[i, :, c_len - t_len : c_len] = sample_tokens.squeeze(0) - input_ids[B + i, :, :t_len] = sample_tokens.squeeze(0) + packed_sample_tokens = sample_tokens.squeeze(0).transpose(0, 1) + input_ids[cond_end - t_len : cond_end] = packed_sample_tokens + input_ids[cond_end : cond_end + t_len] = packed_sample_tokens return tokens From 7e0c40aa0c4e88e89023cf54095156acba7d0ff8 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Sun, 23 Aug 2026 20:38:53 +0800 Subject: [PATCH 06/17] update comments Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../models/omnivoice/pipeline_omnivoice.py | 5 ++-- .../models/omnivoice/omnivoice_generator.py | 30 +++++++++++-------- 2 files changed, 20 insertions(+), 15 deletions(-) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 0c5185d38f3..30f8f84dee9 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -463,11 +463,10 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS layer_penalty_factor = self.layer_penalty_factor if use_cuda_graph: - # Float mask skips per-layer conversion; fp32 cast deferred to the per-item slices below. + # Replay a fixed packed-token bucket with dynamic varlen metadata. batch_logits = self.generator._cuda_graph_fwd(input_ids, audio_masks, cu_seqs, B) else: - # Recompute embeddings and RoPE from the current dynamically - # padded/reordered batch. + # Run packed eager attention for the current active requests. inputs_embeds = self.generator._prepare_embeddings(input_ids, audio_masks) hidden_states = self.generator._transformer_forward(inputs_embeds, cu_seqs) # fp32 cast deferred to the per-item slices below. diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py index d680de908b3..ba23030cf80 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py @@ -244,7 +244,7 @@ def _gumbel_sample(logits: torch.Tensor, temperature: float, generator: torch.Ge # --------------------------------------------------------------------------- -# Qwen3-style transformer blocks using PyTorch SDPA +# Qwen3-style transformer blocks using PyTorch variable-length attention # --------------------------------------------------------------------------- @@ -934,19 +934,25 @@ def forward( """Run the full 32-step iterative unmasking generation. Args: - input_ids: [T, 8] packed as [cond0, uncond0, ...] - audio_mask: [T] - True for audio positions - cond_lens: conditional lengths per request - target_lens: List of target audio lengths per batch item - num_step: Number of unmasking steps - guidance_scale: CFG scale - t_shift: Time shift for schedule - layer_penalty_factor: Penalty for later codebooks - position_temperature: Gumbel temperature for position selection - class_temperature: Temperature for token prediction (0=greedy) + input_ids: Packed token IDs with shape ``[total_seq_len, 8]`` in + request-major ``[cond0, uncond0, ...]`` order. + audio_mask: Boolean audio-position mask with shape + ``[total_seq_len]``. + cond_lens: Conditional sequence length for each request. + target_lens: Target length for each request; also the corresponding + unconditional sequence length. + seed: One seed per request, a shared scalar seed, or ``None``. + num_step: Number of iterative unmasking steps. + guidance_scale: Classifier-free guidance scale. + t_shift: Time shift used to construct the unmasking schedule. + layer_penalty_factor: Penalty applied to later codebooks. + position_temperature: Gumbel temperature for position selection. + class_temperature: Token sampling temperature; zero selects greedy + decoding. Returns: - tokens: [1, 8, sum(target_lens)] - generated audio tokens + Packed generated audio tokens with shape + ``[1, 8, sum(target_lens)]``. """ B = len(target_lens) device = input_ids.device From a3ce527ee99bd928e7cc1489da3bb03c2f4e5ce1 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Mon, 24 Aug 2026 01:46:35 +0800 Subject: [PATCH 07/17] revert batched vae decode commit Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../models/omnivoice/test_decoder.py | 56 -------------- .../models/omnivoice/pipeline_omnivoice.py | 2 +- .../models/omnivoice/omnivoice_decoder.py | 73 ++----------------- 3 files changed, 9 insertions(+), 122 deletions(-) diff --git a/tests/model_executor/models/omnivoice/test_decoder.py b/tests/model_executor/models/omnivoice/test_decoder.py index acfb8a4a062..db8e7328ea1 100644 --- a/tests/model_executor/models/omnivoice/test_decoder.py +++ b/tests/model_executor/models/omnivoice/test_decoder.py @@ -212,59 +212,3 @@ def test_acoustic_decoder_weights_are_float32(): assert param.dtype == torch.float32, ( f"acoustic_decoder.{name} is {param.dtype} — must be float32 (see load_weights fix)" ) - - -# --------------------------------------------------------------------------- -# 4. Bit-Exact when Input is padded -# --------------------------------------------------------------------------- - - -def _build_padded_decoder() -> OmniVoiceDecoder: - from transformers import DacConfig, DacModel - - torch.manual_seed(42) - - decoder = OmniVoiceDecoder(OmniVoiceConfig()) - - decoder.quantizer = HiggsAudioRVQ( - num_quantizers=8, - codebook_size=1024, - codebook_dim=64, - hidden_size=1024, - ).to(DEVICE) - - decoder.fc2 = nn.Linear(1024, 256).to(DEVICE).float() - - dac_config = DacConfig( - hidden_size=256, - decoder_hidden_size=64, - upsampling_ratios=[2, 2], - ) - decoder.acoustic_decoder = DacModel(dac_config).decoder.to(DEVICE).float().eval() - - decoder._adjust_output_padding(decoder.acoustic_decoder) - decoder.acoustic_decoder.tanh = nn.Identity() - - decoder._loaded = True - return decoder - - -def test_output_bit_exact_when_frames_are_padded(): - decoder = _build_padded_decoder() - exact_inputs = torch.randint( - 0, - 1024, - (1, 8, T_FRAMES), - device=DEVICE, - ) - - padded_inputs = torch.nn.functional.pad( - exact_inputs, - (0, 1), - value=0, - ) - padded_length = [T_FRAMES] - exact_output = decoder(exact_inputs) - padded_output = decoder(padded_inputs, padded_length) - padded_output = padded_output[:, :, : T_FRAMES * UPSAMPLE] - torch.testing.assert_close(exact_output, padded_output) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 30f8f84dee9..be4fba4eb8d 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -606,7 +606,7 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: target_offset = 0 for target_len in batch_target_len: request_tokens = tokens[:, :, target_offset : target_offset + target_len] - audio = self.decoder(request_tokens, [target_len]) + audio = self.decoder(request_tokens) outputs.append(DiffusionOutput(output=audio)) target_offset += target_len return outputs diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py index d40a67be3f4..efb6d9dc7af 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_decoder.py @@ -86,47 +86,8 @@ def __init__(self, config: OmniVoiceConfig): self.fc2 = None self.acoustic_decoder = None - @staticmethod - def _mask_by_lengths(hidden_states: torch.Tensor, lengths: torch.Tensor) -> torch.Tensor: - positions = torch.arange(hidden_states.shape[-1], device=hidden_states.device) - valid = positions.unsqueeze(0) < lengths.unsqueeze(1) - return hidden_states.masked_fill(~valid.unsqueeze(1), 0.0) - - def _decode_acoustic_length_aware( - self, - hidden_states: torch.Tensor, - lengths: torch.Tensor, - ) -> torch.Tensor: - """Run the non-causal DAC decoder while zeroing padded activations.""" - decoder = self.acoustic_decoder - hidden_states = self._mask_by_lengths(hidden_states, lengths) - hidden_states = decoder.conv1(hidden_states) - hidden_states = self._mask_by_lengths(hidden_states, lengths) - - for block in decoder.block: - hidden_states = block.snake1(hidden_states) - hidden_states = block.conv_t1(hidden_states) - stride = block.conv_t1.stride[0] - lengths = lengths * stride - hidden_states = self._mask_by_lengths(hidden_states, lengths) - - for residual_unit in (block.res_unit1, block.res_unit2, block.res_unit3): - hidden_states = residual_unit(hidden_states) - hidden_states = self._mask_by_lengths(hidden_states, lengths) - - hidden_states = decoder.snake1(hidden_states) - hidden_states = self._mask_by_lengths(hidden_states, lengths) - hidden_states = decoder.conv2(hidden_states) - hidden_states = self._mask_by_lengths(hidden_states, lengths) - hidden_states = decoder.tanh(hidden_states) - return self._mask_by_lengths(hidden_states, lengths) - @torch.inference_mode() - def forward( - self, - audio_codes: torch.Tensor, - target_lens: list[int] | torch.Tensor | None = None, - ) -> torch.Tensor: + def forward(self, audio_codes: torch.Tensor) -> torch.Tensor: """Decode audio tokens to waveform. Args: @@ -139,25 +100,6 @@ def forward( raise RuntimeError("Decoder not loaded. Call load_weights() first.") device = audio_codes.device - if target_lens is None: - lengths = torch.full( - (audio_codes.shape[0],), - audio_codes.shape[-1], - dtype=torch.long, - device=device, - ) - elif isinstance(target_lens, torch.Tensor): - lengths = target_lens.to(device=device, dtype=torch.long) - else: - lengths = torch.tensor(target_lens, device=device, dtype=torch.long) - - if lengths.shape != (audio_codes.shape[0],): - raise ValueError( - f"Expected one target length per request, got shape {tuple(lengths.shape)} " - f"for batch size {audio_codes.shape[0]}." - ) - if torch.any(lengths <= 0) or torch.any(lengths > audio_codes.shape[-1]): - raise ValueError(f"Target lengths must be in [1, {audio_codes.shape[-1]}], got {lengths.tolist()}.") # Transpose: [B, 8, T] → [8, B, T] codes = audio_codes.transpose(0, 1).long() @@ -165,18 +107,19 @@ def forward( # RVQ decode: sum codebook embeddings → [B, 1024, T] quantized = self.quantizer.decode(codes) - # Project: [B, 1024, T] → fc2 → [B, 256, T]. Keep the acoustic - # decoder in float32 to avoid ConvTranspose1d intermediate overflow. + # Project: [B, 1024, T] → fc2 → [B, 256, T] + # Cast to fc2 weight dtype (may be fp16 when checkpoint stores weights as fp16), + # then upcast back to float32 — acoustic decoder ConvTranspose1d upsampling + # produces intermediate values that exceed the fp16 range (~65504), causing NaN. quantized = self.fc2(quantized.transpose(1, 2).to(self.fc2.weight.dtype)).transpose(1, 2).float() # Acoustic decoder: [B, 256, T] → [B, 1, T*960] - if all(hasattr(self.acoustic_decoder, name) for name in ("conv1", "block", "snake1", "conv2", "tanh")): - audio = self._decode_acoustic_length_aware(quantized, lengths) - else: - audio = self.acoustic_decoder(quantized) + audio = self.acoustic_decoder(quantized) + # Ensure [B, 1, samples] if audio.dim() == 2: audio = audio.unsqueeze(1) + return audio.to(device) def _adjust_output_padding(self, decoder: nn.Module): From 1e6dfe4757afa4c9459017ac1db609ffee9e46c5 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 27 Aug 2026 02:04:13 +0800 Subject: [PATCH 08/17] resolve comments, varlen_attn riskies for non-CUDA devices still pending Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- docs/serving/speech_api.md | 4 + tests/diffusion/test_diffusion_scheduler.py | 31 +++ .../test_omnivoice_expansion.py | 61 ++++-- .../online_serving/test_omnivoice_parity.py | 63 ++++-- .../omnivoice/test_pipeline_batching.py | 191 ++++++++++++++++++ .../models/omnivoice/pipeline_omnivoice.py | 111 +++++----- vllm_omni/diffusion/worker/utils.py | 2 + .../entrypoints/openai/serving_speech.py | 16 +- .../models/omnivoice/omnivoice_generator.py | 38 ++++ 9 files changed, 421 insertions(+), 96 deletions(-) create mode 100644 tests/model_executor/models/omnivoice/test_pipeline_batching.py diff --git a/docs/serving/speech_api.md b/docs/serving/speech_api.md index d4871ce15fa..a515e655768 100644 --- a/docs/serving/speech_api.md +++ b/docs/serving/speech_api.md @@ -752,6 +752,10 @@ Fish Speech uses `ref_audio` and `ref_text` for voice cloning (no `task_type` ne | ------- | ------------- | | `k2-fsa/OmniVoice` | Pure-diffusion TTS. Supports voice cloning via `ref_audio` (with optional `ref_text`); no built-in voice presets. | +OmniVoice uses packed variable-length attention for batched generator execution. The attention operator accepts FP16 and BF16 inputs, so its query, +key, and value tensors are evaluated in BF16 even when the stage is configured with `dtype: float32`; the attention output is converted back to the model's +hidden-state dtype before the output projection. Consequently, a float32 stage configuration does not imply FP32 attention arithmetic. + ### VoxCPM2 | Model | Description | diff --git a/tests/diffusion/test_diffusion_scheduler.py b/tests/diffusion/test_diffusion_scheduler.py index 0a317f1fb55..71f1a6c81ad 100644 --- a/tests/diffusion/test_diffusion_scheduler.py +++ b/tests/diffusion/test_diffusion_scheduler.py @@ -272,6 +272,7 @@ class TestGetRequestBatchSamplingParamsKey: def _make( *, num_inference_steps: int = 2, + guidance_scale: float | None = None, seed: int | None = 123, generator: torch.Generator | None = None, extra_args: dict | None = None, @@ -279,6 +280,7 @@ def _make( ) -> OmniDiffusionRequest: sp = OmniDiffusionSamplingParams( num_inference_steps=num_inference_steps, + guidance_scale=guidance_scale, seed=seed, generator=generator, extra_args=extra_args or {}, @@ -839,6 +841,35 @@ def test_batches_incompatible_request_sampling_params_separately(self) -> None: assert first.num_running_reqs == 1 assert first.num_waiting_reqs == 1 + def test_batches_different_guidance_scales_separately(self) -> None: + scheduler = RequestScheduler() + scheduler.initialize(SimpleNamespace(max_num_seqs=2)) + + first_id = scheduler.add_request( + _make_step_request( + "guidance-2", + sampling_params=OmniDiffusionSamplingParams( + num_inference_steps=2, + guidance_scale=2.0, + ), + ) + ) + scheduler.add_request( + _make_step_request( + "guidance-7", + sampling_params=OmniDiffusionSamplingParams( + num_inference_steps=2, + guidance_scale=7.0, + ), + ) + ) + + first = scheduler.schedule() + + assert _new_ids(first) == [first_id] + assert first.num_running_reqs == 1 + assert first.num_waiting_reqs == 1 + def test_batches_different_quality_levels_separately(self) -> None: scheduler = RequestScheduler() scheduler.initialize(SimpleNamespace(max_num_seqs=2)) diff --git a/tests/e2e/online_serving/test_omnivoice_expansion.py b/tests/e2e/online_serving/test_omnivoice_expansion.py index 1a17d11d45d..d77405a735a 100644 --- a/tests/e2e/online_serving/test_omnivoice_expansion.py +++ b/tests/e2e/online_serving/test_omnivoice_expansion.py @@ -7,7 +7,10 @@ accessed through the standard OpenAI-compatible speech API. """ +import io import os +import wave +from collections.abc import Sequence os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn" @@ -15,7 +18,7 @@ from tests.helpers.mark import hardware_test from tests.helpers.media import get_asset_path -from tests.helpers.runtime import OmniServerParams +from tests.helpers.runtime import OmniResponse, OmniServerParams from tests.helpers.stage_config import get_deploy_config_path from vllm_omni.entrypoints.openai.serving_speech import _DEFAULT_VOICE_NAME @@ -31,7 +34,14 @@ MODEL = "k2-fsa/OmniVoice" STAGE_CONFIG = get_deploy_config_path("omnivoice.yaml") -EXTRA_ARGS = ["--trust-remote-code", "--disable-log-stats", "--max-num-seqs", "8"] +EXTRA_ARGS = [ + "--trust-remote-code", + "--disable-log-stats", + "--max-num-seqs", + "8", + "--request-batch-max-wait-ms", + "50", +] TEST_PARAMS = [ OmniServerParams( model=MODEL, @@ -39,11 +49,7 @@ server_args=EXTRA_ARGS, ) ] -STEP_EXECUTION_ARGS = [ - "--trust-remote-code", - "--enforce-eager", - "--step-execution", -] +STEP_EXECUTION_ARGS = EXTRA_ARGS + ["--step-execution"] STEP_EXECUTION_PARAMS = [ OmniServerParams( model=MODEL, @@ -66,6 +72,33 @@ def get_prompt(prompt_type="text"): return prompts.get(prompt_type, prompts["text"]) +def _assert_valid_omnivoice_wav_responses(responses: Sequence[OmniResponse]) -> None: + """Validate that every concurrent response contains plausible OmniVoice audio.""" + audio_shapes: list[tuple[int, int, int, int]] = [] + for response in responses: + audio_bytes = response.audio_bytes + assert audio_bytes is not None + assert audio_bytes.startswith(b"RIFF") + with wave.open(io.BytesIO(audio_bytes), "rb") as wav_file: + num_channels = wav_file.getnchannels() + sample_width = wav_file.getsampwidth() + sample_rate = wav_file.getframerate() + num_frames = wav_file.getnframes() + assert num_channels == 1 + assert sample_width == 2 + assert sample_rate == 24000 + assert num_frames > 0 + audio_shapes.append((num_channels, sample_width, sample_rate, num_frames)) + + duration_s = num_frames / sample_rate + assert 0.1 <= duration_s <= 30.0 + + # Identical prompts use the same estimated target length. Their samples may + # differ numerically across batch layouts, but the decoded audio shape must + # remain consistent across concurrent requests. + assert len(set(audio_shapes)) == 1 + + @pytest.mark.parametrize("omni_server", TEST_PARAMS, indirect=True) class TestOmniVoiceTTS: """E2E tests for OmniVoice TTS model.""" @@ -84,7 +117,7 @@ def test_speech_auto_voice(self, omni_server, online_client) -> None: @hardware_test(res={"cuda": "L4"}, num_cards=1) def test_speech_auto_voice_batch(self, omni_server, openai_client) -> None: - """The same seeded request must match in request-batch and per-request execution.""" + """Test concurrent request-batch TTS generation.""" batch_request_config = { "model": omni_server.model, "input": get_prompt("text"), @@ -92,12 +125,9 @@ def test_speech_auto_voice_batch(self, omni_server, openai_client) -> None: "seed": 42, "min_audio_bytes": _DEFAULT_MIN_AUDIO_BYTES, } - single_request_config = dict(batch_request_config) batch_r = openai_client.send_audio_speech_request(batch_request_config, request_num=2) - single_r = openai_client.send_audio_speech_request(single_request_config)[0] assert len(batch_r) == 2 - assert batch_r[0].audio_bytes == single_r.audio_bytes - assert batch_r[1].audio_bytes == single_r.audio_bytes + _assert_valid_omnivoice_wav_responses(batch_r) @pytest.mark.parametrize("omni_server", STEP_EXECUTION_PARAMS, indirect=True) @@ -121,7 +151,7 @@ def test_speech_auto_voice_step_execution(self, omni_server, openai_client) -> N @hardware_test(res={"cuda": "L4"}, num_cards=1) def test_speech_auto_voice_batch_step_execution(self, omni_server, openai_client) -> None: - """The same seeded request must match in step-execution-batch and per-request execution.""" + """Test concurrent step-execution TTS generation.""" batch_request_config = { "model": omni_server.model, "input": get_prompt("text"), @@ -132,12 +162,9 @@ def test_speech_auto_voice_batch_step_execution(self, omni_server, openai_client "num_inference_steps": 32, }, } - single_request_config = dict(batch_request_config) batch_r = openai_client.send_audio_speech_request(batch_request_config, request_num=2) - single_r = openai_client.send_audio_speech_request(single_request_config)[0] assert len(batch_r) == 2 - assert batch_r[0].audio_bytes == single_r.audio_bytes - assert batch_r[1].audio_bytes == single_r.audio_bytes + _assert_valid_omnivoice_wav_responses(batch_r) @pytest.mark.parametrize("omni_server", TEST_PARAMS, indirect=True) diff --git a/tests/e2e/online_serving/test_omnivoice_parity.py b/tests/e2e/online_serving/test_omnivoice_parity.py index f06ce6b7c15..e9b267f749d 100644 --- a/tests/e2e/online_serving/test_omnivoice_parity.py +++ b/tests/e2e/online_serving/test_omnivoice_parity.py @@ -19,16 +19,17 @@ STAGE_CONFIG = get_deploy_config_path("omnivoice.yaml") PROMPT = "The weather is nice today, perfect for a walk in the park." +payload = { + "model": MODEL, + "input": PROMPT, + "language": "English", + "seed": 42, + "response_format": "wav", + "extra_params": {"num_inference_steps": 32}, +} -def _generate_once(server_args: list[str]) -> bytes: - payload = { - "model": MODEL, - "input": PROMPT, - "language": "English", - "seed": 42, - "response_format": "wav", - "extra_params": {"num_inference_steps": 32}, - } + +def _generate_without_graph(server_args: list[str]) -> bytes: with OmniServer( MODEL, server_args, @@ -45,10 +46,25 @@ def _generate_once(server_args: list[str]) -> bytes: return response.content -@hardware_test(res={"cuda": "L4"}, num_cards=1) -def test_request_mode_and_step_execution_b1_parity() -> None: - """B=1 request mode and step execution must produce identical seeded WAV bytes.""" - common_args = [ +def _generate_with_graph(server_args: list[str]) -> bytes: + with OmniServer( + MODEL, + server_args, + use_omni=True, + ) as server: + response = requests.post( + f"http://{server.host}:{server.port}/v1/audio/speech", + json=payload, + timeout=600, + ) + response.raise_for_status() + assert response.content.startswith(b"RIFF") + return response.content + + +def _common_args() -> list[str]: + """Return the shared B=1 server arguments for parity tests.""" + return [ "--trust-remote-code", "--disable-log-stats", "--deploy-config", @@ -57,7 +73,24 @@ def test_request_mode_and_step_execution_b1_parity() -> None: "1", ] - request_audio = _generate_once(common_args) - step_audio = _generate_once([*common_args, "--step-execution", "--enforce-eager"]) + +@hardware_test(res={"cuda": "L4"}, num_cards=1) +def test_request_mode_and_step_execution_b1_parity_without_graph() -> None: + """B=1 eager request and step modes must produce identical seeded WAV bytes.""" + common_args = _common_args() + + request_audio = _generate_without_graph(common_args) + step_audio = _generate_without_graph([*common_args, "--step-execution", "--enforce-eager"]) + + assert request_audio == step_audio + + +@hardware_test(res={"cuda": "L4"}, num_cards=1) +def test_request_mode_and_step_execution_b1_parity_with_graph() -> None: + """B=1 Graph request and step modes must produce identical seeded WAV bytes.""" + common_args = _common_args() + + request_audio = _generate_with_graph(common_args) + step_audio = _generate_with_graph([*common_args, "--step-execution", "--enforce-eager"]) assert request_audio == step_audio diff --git a/tests/model_executor/models/omnivoice/test_pipeline_batching.py b/tests/model_executor/models/omnivoice/test_pipeline_batching.py new file mode 100644 index 00000000000..2890999b209 --- /dev/null +++ b/tests/model_executor/models/omnivoice/test_pipeline_batching.py @@ -0,0 +1,191 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace + +import pytest +import torch + +from vllm_omni.diffusion.data import DiffusionOutput +from vllm_omni.diffusion.models.omnivoice.pipeline_omnivoice import ( + OmniVoicePipeline, + _PreparedOmniVoiceRequest, +) +from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch +from vllm_omni.inputs.data import OmniDiffusionSamplingParams + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + + +def test_request_batch_prepare_error_does_not_skip_valid_request() -> None: + """A malformed request must not prevent another request from completing.""" + prepared = _PreparedOmniVoiceRequest( + input_ids=torch.zeros(5, 8, dtype=torch.long), + audio_mask=torch.ones(5, dtype=torch.bool), + cond_len=3, + target_len=2, + seed=42, + ) + + def prepare_request(prompt, extra): + del extra + if prompt == "bad": + return DiffusionOutput(error="invalid prompt") + return prepared + + def collate_requests(requests): + assert requests == [prepared] + return prepared.input_ids, prepared.audio_mask, [prepared.cond_len] + + generator_calls = [] + + def generate_tokens(**kwargs): + generator_calls.append(kwargs) + assert kwargs["target_lens"] == [prepared.target_len] + return torch.ones(1, 8, prepared.target_len, dtype=torch.long) + + pipeline = SimpleNamespace( + _prepare_request_input=prepare_request, + _collate_request_inputs=collate_requests, + generator=generate_tokens, + decoder=lambda tokens: tokens.float(), + num_step=32, + guidance_scale=7.0, + t_shift=1.0, + layer_penalty_factor=0.0, + position_temperature=0.0, + class_temperature=0.0, + ) + sampling = OmniDiffusionSamplingParams(num_inference_steps=32, guidance_scale=7.0) + batch = DiffusionRequestBatch( + requests=[ + OmniDiffusionRequest(prompt="bad", sampling_params=sampling, request_id="bad"), + OmniDiffusionRequest( + prompt="good", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=32, guidance_scale=7.0), + request_id="good", + ), + ] + ) + + outputs = OmniVoicePipeline.forward(pipeline, batch) + + assert len(outputs) == 2 + assert outputs[0].error == "invalid prompt" + assert outputs[1].error is None + torch.testing.assert_close(outputs[1].output, torch.ones(1, 8, 2)) + assert generator_calls[0]["guidance_scale"] == 7.0 + + +def test_request_batch_honors_explicit_sampling_overrides() -> None: + """Request steps and guidance must override the OmniVoice defaults.""" + prepared = _PreparedOmniVoiceRequest( + input_ids=torch.zeros(5, 8, dtype=torch.long), + audio_mask=torch.ones(5, dtype=torch.bool), + cond_len=3, + target_len=2, + seed=42, + ) + captured_sampling = [] + + def generate_tokens(**kwargs): + captured_sampling.append((kwargs["num_step"], kwargs["guidance_scale"])) + return torch.ones(1, 8, prepared.target_len, dtype=torch.long) + + pipeline = SimpleNamespace( + _prepare_request_input=lambda prompt, extra: prepared, + _collate_request_inputs=lambda requests: (prepared.input_ids, prepared.audio_mask, [prepared.cond_len]), + generator=generate_tokens, + decoder=lambda tokens: tokens.float(), + num_step=32, + guidance_scale=2.0, + t_shift=1.0, + layer_penalty_factor=0.0, + position_temperature=0.0, + class_temperature=0.0, + ) + batch = DiffusionRequestBatch( + requests=[ + OmniDiffusionRequest( + prompt="hello", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=12, guidance_scale=6.5), + request_id="request", + ) + ] + ) + + outputs = OmniVoicePipeline.forward(pipeline, batch) + + assert outputs[0].error is None + assert captured_sampling == [(12, 6.5)] + + +def test_request_batch_prepare_error_preserves_surrounding_output_indices() -> None: + """A middle prepare error must not shift valid outputs into the wrong slots.""" + prepared_before = _PreparedOmniVoiceRequest( + input_ids=torch.zeros(3, 8, dtype=torch.long), + audio_mask=torch.ones(3, dtype=torch.bool), + cond_len=2, + target_len=1, + seed=1, + ) + prepared_after = _PreparedOmniVoiceRequest( + input_ids=torch.zeros(5, 8, dtype=torch.long), + audio_mask=torch.ones(5, dtype=torch.bool), + cond_len=3, + target_len=2, + seed=2, + ) + + def prepare_request(prompt, extra): + del extra + if prompt == "before": + return prepared_before + if prompt == "after": + return prepared_after + return DiffusionOutput(error="invalid middle request") + + def collate_requests(requests): + assert requests == [prepared_before, prepared_after] + return ( + torch.cat([prepared_before.input_ids, prepared_after.input_ids]), + torch.cat([prepared_before.audio_mask, prepared_after.audio_mask]), + [prepared_before.cond_len, prepared_after.cond_len], + ) + + def generate_tokens(**kwargs): + assert kwargs["target_lens"] == [1, 2] + before = torch.full((1, 8, 1), 11, dtype=torch.long) + after = torch.full((1, 8, 2), 22, dtype=torch.long) + return torch.cat([before, after], dim=-1) + + pipeline = SimpleNamespace( + _prepare_request_input=prepare_request, + _collate_request_inputs=collate_requests, + generator=generate_tokens, + decoder=lambda tokens: tokens.float(), + num_step=32, + guidance_scale=2.0, + t_shift=1.0, + layer_penalty_factor=0.0, + position_temperature=0.0, + class_temperature=0.0, + ) + batch = DiffusionRequestBatch( + requests=[ + OmniDiffusionRequest( + prompt=prompt, + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=32), + request_id=prompt, + ) + for prompt in ("before", "bad", "after") + ] + ) + + outputs = OmniVoicePipeline.forward(pipeline, batch) + + assert len(outputs) == 3 + torch.testing.assert_close(outputs[0].output, torch.full((1, 8, 1), 11.0)) + assert outputs[1].error == "invalid middle request" + torch.testing.assert_close(outputs[2].output, torch.full((1, 8, 2), 22.0)) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index be4fba4eb8d..865974a0e0f 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -22,7 +22,6 @@ import numpy as np import torch -import torch.nn.functional as F from tokenizers import Tokenizer as HFTokenizer from torch import nn from vllm.logger import init_logger @@ -39,7 +38,6 @@ OmniVoiceGenerator, _build_cu_seqs, _get_time_steps, - _gumbel_sample, ) from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig from vllm_omni.utils.speaker_cache import get_speaker_cache @@ -174,7 +172,7 @@ def __init__(self, *, od_config: OmniDiffusionConfig, prefix: str = ""): self.config = OmniVoiceConfig(**hf_config) # Build generator and decoder - self.generator = OmniVoiceGenerator(self.config) + self.generator = OmniVoiceGenerator(self.config, od_config) self.decoder = OmniVoiceDecoder(self.config) # Tokenizer (low-level, avoids HF tokenizer extra_special_tokens issue) @@ -368,14 +366,15 @@ def _collate_request_inputs( return input_ids, audio_masks, cond_lens - def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: + def prepare_encode(self, state: StepRequestState) -> StepRequestState: prompt = state.prompt if state.prompt else "" extra = state.sampling.extra_args or {} prepared = self._prepare_request_input(prompt, extra) if isinstance(prepared, DiffusionOutput): - return prepared - prepared_request = prepared + state.error = prepared.error + return state + prepared_request = prepared cond_len = prepared_request.cond_len target_len = prepared_request.target_len input_ids = prepared_request.input_ids @@ -386,7 +385,10 @@ def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: num_codebooks = self.config.num_audio_codebook if seed is None: seed = random.randint(0, 2**63 - 1) - num_step = self.num_step + num_step = ( + state.sampling.num_inference_steps if state.sampling.num_inference_steps is not None else self.num_step + ) + t_shift = self.t_shift # Initialize all target tokens as [MASK] @@ -415,9 +417,12 @@ def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: layer_ids = torch.arange(num_codebooks, device=device).view(1, -1, 1) generator = torch.Generator(device=device).manual_seed(seed) + guidance_scale = ( + state.sampling.guidance_scale if state.sampling.guidance_scale is not None else self.guidance_scale + ) state.latents = input_ids state.timesteps = schedules - state.guidance = self.guidance_scale + state.guidance = guidance_scale state.extra["schedules"] = schedules state.extra["layer_ids"] = layer_ids state.extra["generator"] = generator @@ -426,10 +431,11 @@ def prepare_encode(self, state: StepRequestState) -> DiffusionRequestBatch: state.extra["target_len"] = target_len state.extra["audio_mask"] = audio_mask state.extra["tokens"] = tokens + return state def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestState] | None = None, **kwargs: Any): - use_cuda_graph = self.generator._cuda_graph_fwd is not None input_ids = input_batch.latents + use_cuda_graph = self.generator._cuda_graph_fwd is not None and input_ids.is_cuda layer_ids = states[0].extra["layer_ids"] audio_masks: list[torch.Tensor] = [] @@ -447,7 +453,7 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS cond_lens.append(state.extra["cond_len"]) target_lens.append(state.extra["target_len"]) batch_tokens.append(state.extra["tokens"]) - guidance_scales.append(state.extra.get("guidance", self.guidance_scale)) + guidance_scales.append(state.guidance) generators.append(state.extra.get("generator", None)) schedules.append(state.extra["schedules"]) steps.append(state.step_index) @@ -457,18 +463,20 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS B = len(target_lens) cu_seqs = _build_cu_seqs(cond_lens, target_lens, input_ids.device) - mask_id = self.config.audio_mask_id position_temperature = self.position_temperature class_temperature = self.class_temperature layer_penalty_factor = self.layer_penalty_factor - if use_cuda_graph: # Replay a fixed packed-token bucket with dynamic varlen metadata. batch_logits = self.generator._cuda_graph_fwd(input_ids, audio_masks, cu_seqs, B) else: # Run packed eager attention for the current active requests. inputs_embeds = self.generator._prepare_embeddings(input_ids, audio_masks) - hidden_states = self.generator._transformer_forward(inputs_embeds, cu_seqs) + hidden_states = self.generator._transformer_forward( + inputs_embeds, + cu_seqs, + max_seqlen=max(cond_lens), + ) # fp32 cast deferred to the per-item slices below. batch_logits = self.generator._get_logits(hidden_states) # batch_logits: [8, total_seq_len, 1025] @@ -501,54 +509,26 @@ def denoise_step(self, input_batch: InputBatch, *, states: Sequence[StepRequestS # Extract logits for target region; upcast only the slices we actually consume. c_logits = batch_logits[:, cond_end - t_len : cond_end, :].unsqueeze(0).to(torch.float32) u_logits = batch_logits[:, uncond_start : uncond_start + t_len, :].unsqueeze(0).to(torch.float32) - - # Classifier-free guidance. Fuse the chain: the two inner - # log_softmax normalizers are per-position scalars that the final - # shift-invariant log_softmax cancels, so guide on the raw logits - # with a single softmax: log_softmax((1+s)*c - s*u). Exact. - if guidance_scales[i] != 0: - log_probs = F.log_softmax( - (1.0 + guidance_scales[i]) * c_logits - guidance_scales[i] * u_logits, - dim=-1, - ) - else: - log_probs = F.log_softmax(c_logits, dim=-1) - - # Prevent predicting [MASK] - log_probs[..., mask_id] = -float("inf") - - # Token prediction - if class_temperature > 0.0: - pred_tokens = _gumbel_sample(log_probs, class_temperature, generators[i]).argmax(dim=-1) - else: - pred_tokens = log_probs.argmax(dim=-1) # [1, 8, T] - - # Confidence scores - scores = log_probs.max(dim=-1)[0] # [1, 8, T] - - # Layer penalty (earlier codebooks get higher priority) - scores = scores - (layer_ids * layer_penalty_factor) - - # Gumbel noise for position selection - if position_temperature > 0.0: - scores = _gumbel_sample(scores, position_temperature, generators[i]) - - # Mask out already unmasked positions sample = batch_tokens[i] sample_tokens = sample[..., :t_len] - scores.masked_fill_(sample_tokens != mask_id, -float("inf")) - - # Select top-k positions to unmask. .flatten() on this non-contiguous view already copies. - _, topk_idx = torch.topk(scores.flatten(), k) - flat_tokens = sample_tokens.flatten() - flat_tokens[topk_idx] = pred_tokens.flatten()[topk_idx] - sample_tokens.copy_(flat_tokens.view_as(sample_tokens)) - states[i].extra["tokens"] = sample_tokens + self.generator._unmask_one_request( + c_logits, + u_logits, + sample_tokens, + num_to_unmask=k, + guidance_scale=guidance_scales[i], + generator=generators[i], + class_temperature=class_temperature, + position_temperature=position_temperature, + layer_penalty_factor=layer_penalty_factor, + layer_ids=layer_ids, + ) # Mirror update into both cond and uncond input_ids halves for the next step. packed_sample_tokens = sample_tokens.squeeze(0).transpose(0, 1) input_ids[cond_end - t_len : cond_end] = packed_sample_tokens input_ids[uncond_start : uncond_start + t_len] = packed_sample_tokens + states[i].extra["tokens"] = sample_tokens # InputBatch reuses its latents buffer across steps. Returning that # same storage would make the Runner persist per-request views into the @@ -576,25 +556,35 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: "lang": "...", "instruct": "..."} """ prepared_requests: list[_PreparedOmniVoiceRequest] = [] - for request in req.requests: + outputs = [None] * len(req.requests) + prepared_indices: list[int] = [] + for i, request in enumerate(req.requests): prompt = request.prompt if request.prompt else "" extra = request.sampling_params.extra_args or {} prepared = self._prepare_request_input(prompt, extra) if isinstance(prepared, DiffusionOutput): - return prepared + outputs[i] = prepared + continue + prepared_indices.append(i) prepared_requests.append(prepared) + if not prepared_requests: + return outputs + batch_target_len = [request.target_len for request in prepared_requests] batch_seeds = [request.seed for request in prepared_requests] batch_input_ids, batch_audio_mask, batch_cond_lens = self._collate_request_inputs(prepared_requests) # Run 32-step iterative unmasking + sampling = req.requests[0].sampling_params + num_step = sampling.num_inference_steps if sampling.num_inference_steps is not None else self.num_step + guidance_scale = sampling.guidance_scale if sampling.guidance_scale is not None else self.guidance_scale tokens = self.generator( input_ids=batch_input_ids, audio_mask=batch_audio_mask, cond_lens=batch_cond_lens, target_lens=batch_target_len, - num_step=self.num_step, - guidance_scale=self.guidance_scale, + num_step=num_step, + guidance_scale=guidance_scale, t_shift=self.t_shift, layer_penalty_factor=self.layer_penalty_factor, position_temperature=self.position_temperature, @@ -602,12 +592,11 @@ def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: seed=batch_seeds, ) - outputs: list[DiffusionOutput] = [] target_offset = 0 - for target_len in batch_target_len: + for i, target_len in enumerate(batch_target_len): request_tokens = tokens[:, :, target_offset : target_offset + target_len] audio = self.decoder(request_tokens) - outputs.append(DiffusionOutput(output=audio)) + outputs[prepared_indices[i]] = DiffusionOutput(output=audio) target_offset += target_len return outputs diff --git a/vllm_omni/diffusion/worker/utils.py b/vllm_omni/diffusion/worker/utils.py index 5158dbfbf37..a3bc81f9dc6 100644 --- a/vllm_omni/diffusion/worker/utils.py +++ b/vllm_omni/diffusion/worker/utils.py @@ -131,6 +131,8 @@ class StepRequestState: # Peak device memory observed while this request is active in step mode. peak_memory_mb: float = 0.0 + error: str | None = None + # ── Properties ── @property diff --git a/vllm_omni/entrypoints/openai/serving_speech.py b/vllm_omni/entrypoints/openai/serving_speech.py index 767077381b6..6512290bdd6 100644 --- a/vllm_omni/entrypoints/openai/serving_speech.py +++ b/vllm_omni/entrypoints/openai/serving_speech.py @@ -3508,12 +3508,22 @@ async def _create_diffusion_speech( sampling = sampling_params_list[0] + # This change allows StepScheduler read total_steps from upper + # sampling.num_inference_steps, check diffusion/sched/step_scheduler:_get_total_steps if "num_inference_steps" in extra: - sampling.num_inference_steps = int(extra["num_inference_steps"]) + value = extra["num_inference_steps"] + try: + sampling.num_inference_steps = int(value) + except (TypeError, ValueError) as exc: + raise ValueError("num_inference_steps must be an integer") from exc if "guidance_scale" in extra: - sampling.guidance_scale = float(extra["guidance_scale"]) - sampling.guidance_scale_provided = True + value = extra["guidance_scale"] + try: + sampling.guidance_scale = float(value) + except (TypeError, ValueError) as exc: + raise ValueError("guidance_scale must be a number") from exc + logger.info("Applied extra_params to diffusion: %s", extra) generator = self._diffusion_engine.generate( diff --git a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py index ba23030cf80..965d021b82f 100644 --- a/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py +++ b/vllm_omni/model_executor/models/omnivoice/omnivoice_generator.py @@ -916,6 +916,44 @@ def _step_forward( hidden_states = self._transformer_forward(hidden_states, cu_seqs, rope_table=rope_table) return self._get_logits(hidden_states) + def _unmask_one_request( + self, + c_logits: torch.Tensor, + u_logits: torch.Tensor, + sample_tokens: torch.Tensor, + *, + num_to_unmask: int | torch.Tensor, + guidance_scale: float, + generator: torch.Generator, + class_temperature: float, + position_temperature: float, + layer_penalty_factor: float, + layer_ids: torch.Tensor, + ) -> None: + """Sample and unmask one request in place from FP32 target logits.""" + mask_id = self.config.audio_mask_id + if guidance_scale != 0: + log_probs = F.log_softmax( + (1.0 + guidance_scale) * c_logits - guidance_scale * u_logits, + dim=-1, + ) + else: + log_probs = F.log_softmax(c_logits, dim=-1) + log_probs[..., mask_id] = -float("inf") + if class_temperature > 0.0: + pred_tokens = _gumbel_sample(log_probs, class_temperature, generator).argmax(dim=-1) + else: + pred_tokens = log_probs.argmax(dim=-1) + scores = log_probs.max(dim=-1)[0] + scores = scores - (layer_ids * layer_penalty_factor) + if position_temperature > 0.0: + scores = _gumbel_sample(scores, position_temperature, generator) + scores.masked_fill_(sample_tokens != mask_id, -float("inf")) + _, topk_idx = torch.topk(scores.flatten(), num_to_unmask) + flat_tokens = sample_tokens.flatten() + flat_tokens[topk_idx] = pred_tokens.flatten()[topk_idx] + sample_tokens.copy_(flat_tokens.view_as(sample_tokens)) + @torch.inference_mode() def forward( self, From ad94581538ed04a46d92e2caed9b8c201b00f3c6 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 27 Aug 2026 02:36:22 +0800 Subject: [PATCH 09/17] remove reduntant params and fix pre-commit Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- tests/diffusion/test_diffusion_scheduler.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/tests/diffusion/test_diffusion_scheduler.py b/tests/diffusion/test_diffusion_scheduler.py index 71f1a6c81ad..2c500faf5b6 100644 --- a/tests/diffusion/test_diffusion_scheduler.py +++ b/tests/diffusion/test_diffusion_scheduler.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project import asyncio import queue @@ -141,7 +141,7 @@ def __init__(self, request: OmniDiffusionRequest, output) -> None: self._output = output self.initialized_with = None self._request_id = request.request_id - self._state = None + self._state: SimpleNamespace | None = None self._scheduled = False self.max_num_running_reqs = 1 @@ -272,7 +272,6 @@ class TestGetRequestBatchSamplingParamsKey: def _make( *, num_inference_steps: int = 2, - guidance_scale: float | None = None, seed: int | None = 123, generator: torch.Generator | None = None, extra_args: dict | None = None, @@ -280,7 +279,6 @@ def _make( ) -> OmniDiffusionRequest: sp = OmniDiffusionSamplingParams( num_inference_steps=num_inference_steps, - guidance_scale=guidance_scale, seed=seed, generator=generator, extra_args=extra_args or {}, From 7c5b030045d4eec58c378b5cb5dc53a51f8e830d Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 27 Aug 2026 09:45:30 +0800 Subject: [PATCH 10/17] use in-tree Attention for SDPA fallback and prevent version mismatch Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../omnivoice/test_pipeline_batching.py | 25 ++++++++++++++++++- 1 file changed, 24 insertions(+), 1 deletion(-) diff --git a/tests/model_executor/models/omnivoice/test_pipeline_batching.py b/tests/model_executor/models/omnivoice/test_pipeline_batching.py index 2890999b209..013cb106e4c 100644 --- a/tests/model_executor/models/omnivoice/test_pipeline_batching.py +++ b/tests/model_executor/models/omnivoice/test_pipeline_batching.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project from types import SimpleNamespace @@ -14,10 +14,33 @@ from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniDiffusionSamplingParams +from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( + _attention_metadata_from_cu_seqs, +) pytestmark = [pytest.mark.core_model, pytest.mark.cpu] +def test_sdpa_fallback_mask_preserves_packed_sequence_boundaries() -> None: + cu_seqs = torch.tensor([0, 2, 3, 6, 6], dtype=torch.int32) + + metadata = _attention_metadata_from_cu_seqs(cu_seqs, 6, needs_sdpa_mask=True) + + assert metadata.attn_mask is not None + expected = torch.tensor( + [ + [1, 1, 0, 0, 0, 0], + [1, 1, 0, 0, 0, 0], + [0, 0, 1, 0, 0, 0], + [0, 0, 0, 1, 1, 1], + [0, 0, 0, 1, 1, 1], + [0, 0, 0, 1, 1, 1], + ], + dtype=torch.bool, + ) + torch.testing.assert_close(metadata.attn_mask[0, 0], expected) + + def test_request_batch_prepare_error_does_not_skip_valid_request() -> None: """A malformed request must not prevent another request from completing.""" prepared = _PreparedOmniVoiceRequest( From c259759605d2a5babca5170302402e6a39a5e590 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 27 Aug 2026 10:22:03 +0800 Subject: [PATCH 11/17] set default max_num_seqs=8 in yaml Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- vllm_omni/deploy/omnivoice.yaml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/vllm_omni/deploy/omnivoice.yaml b/vllm_omni/deploy/omnivoice.yaml index 6769dd5e995..df9e6a5daf0 100644 --- a/vllm_omni/deploy/omnivoice.yaml +++ b/vllm_omni/deploy/omnivoice.yaml @@ -11,3 +11,5 @@ stages: trust_remote_code: true distributed_executor_backend: "mp" dtype: "float32" + max_num_seqs: 8 + request_batch_max_wait_ms: 10 From 9b2943f76263f8d74b25343af4c76c82ebdb2bbc Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 27 Aug 2026 11:50:21 +0800 Subject: [PATCH 12/17] add full pass regression test for requestscheduler and sampling extra_args Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- tests/diffusion/test_diffusion_scheduler.py | 33 ++------------------- 1 file changed, 2 insertions(+), 31 deletions(-) diff --git a/tests/diffusion/test_diffusion_scheduler.py b/tests/diffusion/test_diffusion_scheduler.py index 2c500faf5b6..0a317f1fb55 100644 --- a/tests/diffusion/test_diffusion_scheduler.py +++ b/tests/diffusion/test_diffusion_scheduler.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project import asyncio import queue @@ -141,7 +141,7 @@ def __init__(self, request: OmniDiffusionRequest, output) -> None: self._output = output self.initialized_with = None self._request_id = request.request_id - self._state: SimpleNamespace | None = None + self._state = None self._scheduled = False self.max_num_running_reqs = 1 @@ -839,35 +839,6 @@ def test_batches_incompatible_request_sampling_params_separately(self) -> None: assert first.num_running_reqs == 1 assert first.num_waiting_reqs == 1 - def test_batches_different_guidance_scales_separately(self) -> None: - scheduler = RequestScheduler() - scheduler.initialize(SimpleNamespace(max_num_seqs=2)) - - first_id = scheduler.add_request( - _make_step_request( - "guidance-2", - sampling_params=OmniDiffusionSamplingParams( - num_inference_steps=2, - guidance_scale=2.0, - ), - ) - ) - scheduler.add_request( - _make_step_request( - "guidance-7", - sampling_params=OmniDiffusionSamplingParams( - num_inference_steps=2, - guidance_scale=7.0, - ), - ) - ) - - first = scheduler.schedule() - - assert _new_ids(first) == [first_id] - assert first.num_running_reqs == 1 - assert first.num_waiting_reqs == 1 - def test_batches_different_quality_levels_separately(self) -> None: scheduler = RequestScheduler() scheduler.initialize(SimpleNamespace(max_num_seqs=2)) From f01406b7bc0a3c16d3c17d4bc6d9408085239371 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Thu, 27 Aug 2026 13:24:07 +0800 Subject: [PATCH 13/17] fix test bugs Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- tests/helpers/client.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/helpers/client.py b/tests/helpers/client.py index 0f834620073..43111397f16 100644 --- a/tests/helpers/client.py +++ b/tests/helpers/client.py @@ -1462,6 +1462,7 @@ def send_audio_speech_request(self, request_config: dict[str, Any], request_num: "instructions", "speed", "sample_rate", + "extra_params", "stream_format", "x_vector_only_mode", ): From 817cf68bcbaa6c386336876aacf877c5838b21fc Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Sun, 30 Aug 2026 21:07:41 +0800 Subject: [PATCH 14/17] support exact shape in eager path, add comments for graph paths' metadata Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../models/omnivoice/test_pipeline_batching.py | 14 ++++++++++++++ .../models/omnivoice/pipeline_omnivoice.py | 2 +- 2 files changed, 15 insertions(+), 1 deletion(-) diff --git a/tests/model_executor/models/omnivoice/test_pipeline_batching.py b/tests/model_executor/models/omnivoice/test_pipeline_batching.py index 013cb106e4c..6e611f2ff09 100644 --- a/tests/model_executor/models/omnivoice/test_pipeline_batching.py +++ b/tests/model_executor/models/omnivoice/test_pipeline_batching.py @@ -41,6 +41,20 @@ def test_sdpa_fallback_mask_preserves_packed_sequence_boundaries() -> None: torch.testing.assert_close(metadata.attn_mask[0, 0], expected) +def test_eager_attention_metadata_honors_exact_max_seqlen() -> None: + cu_seqs = torch.tensor([0, 2, 3, 6, 6], dtype=torch.int32) + + metadata = _attention_metadata_from_cu_seqs( + cu_seqs, + 6, + needs_sdpa_mask=False, + max_seqlen=3, + ) + + assert metadata.extra["max_seqlen_q"] == 3 + assert metadata.extra["max_seqlen_k"] == 3 + + def test_request_batch_prepare_error_does_not_skip_valid_request() -> None: """A malformed request must not prevent another request from completing.""" prepared = _PreparedOmniVoiceRequest( diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 865974a0e0f..59ac8b9e18b 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project """ OmniVoice TTS Pipeline for vLLM-Omni diffusion engine. From b66d35e6a21b6da464f280fc69405bb9b067ee86 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Sat, 5 Sep 2026 09:53:55 +0800 Subject: [PATCH 15/17] rebase Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py | 4 ++-- vllm_omni/diffusion/worker/utils.py | 4 +--- 2 files changed, 3 insertions(+), 5 deletions(-) diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 59ac8b9e18b..cdf3d0060d7 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -32,6 +32,7 @@ from vllm_omni.diffusion.worker.input_batch import InputBatch from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.diffusion.worker.utils import StepRequestState +from vllm_omni.errors import OmniClientError from vllm_omni.model_executor.models.omnivoice.duration import RuleDurationEstimator from vllm_omni.model_executor.models.omnivoice.omnivoice_decoder import OmniVoiceDecoder from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( @@ -371,8 +372,7 @@ def prepare_encode(self, state: StepRequestState) -> StepRequestState: extra = state.sampling.extra_args or {} prepared = self._prepare_request_input(prompt, extra) if isinstance(prepared, DiffusionOutput): - state.error = prepared.error - return state + raise OmniClientError(prepared.error or "OmniVoice request preparation failed") prepared_request = prepared cond_len = prepared_request.cond_len diff --git a/vllm_omni/diffusion/worker/utils.py b/vllm_omni/diffusion/worker/utils.py index a3bc81f9dc6..efc935a5dbd 100644 --- a/vllm_omni/diffusion/worker/utils.py +++ b/vllm_omni/diffusion/worker/utils.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project """Per-request mutable state for step-wise diffusion execution.""" from __future__ import annotations @@ -131,8 +131,6 @@ class StepRequestState: # Peak device memory observed while this request is active in step mode. peak_memory_mb: float = 0.0 - error: str | None = None - # ── Properties ── @property From 9a67cb722ad7778222e387b1c0aa49e5b2ffc436 Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Sat, 5 Sep 2026 10:50:44 +0800 Subject: [PATCH 16/17] resolve rebase errors, restore omitted test Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../openai_api/test_serving_speech.py | 112 +++++++++- .../models/omnivoice/test_mask_dtype.py | 196 ++++-------------- 2 files changed, 153 insertions(+), 155 deletions(-) diff --git a/tests/entrypoints/openai_api/test_serving_speech.py b/tests/entrypoints/openai_api/test_serving_speech.py index 47c11d7405c..f6348384a37 100644 --- a/tests/entrypoints/openai_api/test_serving_speech.py +++ b/tests/entrypoints/openai_api/test_serving_speech.py @@ -27,6 +27,9 @@ from pytest_mock import MockerFixture from vllm.entrypoints.openai.engine.protocol import ErrorInfo, ErrorResponse +from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.sched.request_scheduler import RequestScheduler +from vllm_omni.diffusion.sched.step_scheduler import StepScheduler from vllm_omni.entrypoints.omni_base import OmniEngineDeadError from vllm_omni.entrypoints.openai import api_server as api_server_module from vllm_omni.entrypoints.openai import serving_speech as serving_speech_module @@ -55,6 +58,7 @@ from vllm_omni.entrypoints.openai.tts_adapters.ming_tts import MingTTSAdapter from vllm_omni.entrypoints.openai.tts_adapters.qwen3_tts import Qwen3TTSCodecLimitError from vllm_omni.entrypoints.openai.tts_adapters.voxtral import VoxtralTTSAdapter +from vllm_omni.inputs.data import OmniDiffusionSamplingParams from vllm_omni.model_executor.models.fish_speech.prompt_utils import ( FISH_TEXT_ONLY_SYSTEM_PROMPT, build_fish_voice_clone_prompt_ids, @@ -753,13 +757,12 @@ async def test_diffusion_create_speech_with_unknown_voice(self, mocker: MockerFi @pytest.mark.asyncio async def test_create_diffusion_speech_extra_params(self, mocker: MockerFixture): - """Test public diffusion speech success and extra_params propagation.""" + """Test diffusion parameters reach StepScheduler as standard fields.""" # Mock the engine client mock_engine = mocker.MagicMock() # Mock default sampling params - mock_sampling_param = mocker.MagicMock() - mock_sampling_param.extra_args = {"existing_arg": "value"} + mock_sampling_param = OmniDiffusionSamplingParams(extra_args={"existing_arg": "value"}) mock_engine.default_sampling_params_list = [mock_sampling_param] # Mock generate to yield a valid OmniRequestOutput @@ -775,7 +778,15 @@ async def mock_generate(*args, **kwargs): server, "create_audio", return_value=mocker.MagicMock(audio_data=b"dummy", media_type="audio/wav") ) - req = OpenAICreateSpeechRequest(input="Hello", extra_params={"new_arg": 123, "existing_arg": "new_value"}) + req = OpenAICreateSpeechRequest( + input="Hello", + extra_params={ + "new_arg": 123, + "existing_arg": "new_value", + "num_inference_steps": 12, + "guidance_scale": 7.0, + }, + ) response = await server.create_speech(req) @@ -792,7 +803,98 @@ async def mock_generate(*args, **kwargs): # Verify it was deepcopied and updated assert passed_params is not mock_engine.default_sampling_params_list - assert passed_params[0].extra_args == {"existing_arg": "new_value", "new_arg": 123} + assert passed_params[0].extra_args == { + "existing_arg": "new_value", + "new_arg": 123, + "num_inference_steps": 12, + "guidance_scale": 7.0, + } + assert passed_params[0].num_inference_steps == 12 + assert passed_params[0].guidance_scale == 7.0 + + # Regression: StepScheduler.add_request() used to receive + # num_inference_steps=None and fail while converting it to int. + scheduler = StepScheduler() + scheduler.add_request( + OmniDiffusionRequest( + prompt="Hello", + sampling_params=passed_params[0], + request_id="speech-test", + ) + ) + assert scheduler._request_progress["speech-test"].total_steps == 12 + + @pytest.mark.asyncio + async def test_diffusion_speech_guidance_promotion_controls_request_batch_admission( + self, + mocker: MockerFixture, + ) -> None: + """Different request guidance values must not enter one request batch.""" + mock_engine = mocker.MagicMock() + mock_engine.default_sampling_params_list = [OmniDiffusionSamplingParams(num_inference_steps=12)] + passed_sampling_params = [] + + async def mock_generate(*args, **kwargs): + passed_sampling_params.append(kwargs["sampling_params_list"][0]) + yield create_mock_audio_output_for_test() + + mock_engine.generate = mocker.MagicMock(side_effect=mock_generate) + server = OmniOpenAIServingSpeech.for_diffusion(diffusion_engine=mock_engine, model_name="test-model") + mocker.patch.object( + server, + "create_audio", + return_value=mocker.MagicMock(audio_data=b"dummy", media_type="audio/wav"), + ) + + for guidance_scale in (2.0, 7.0): + response = await server.create_speech( + OpenAICreateSpeechRequest( + input="Hello", + extra_params={"guidance_scale": guidance_scale}, + ) + ) + assert response.status_code == 200 + + scheduler = RequestScheduler() + scheduler.initialize(SimpleNamespace(max_num_seqs=2)) + for index, sampling_params in enumerate(passed_sampling_params): + scheduler.add_request( + OmniDiffusionRequest( + prompt="Hello", + sampling_params=sampling_params, + request_id=f"speech-{index}", + ) + ) + + first = scheduler.schedule() + + assert [request.request_id for request in first.scheduled_new_reqs] == ["speech-0"] + assert first.num_running_reqs == 1 + assert first.num_waiting_reqs == 1 + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("extra_params", "expected_message"), + [ + ({"num_inference_steps": "invalid"}, "num_inference_steps must be an integer"), + ({"guidance_scale": "invalid"}, "guidance_scale must be a number"), + ], + ) + async def test_create_diffusion_speech_rejects_invalid_scheduler_params( + self, + mocker: MockerFixture, + extra_params: dict[str, str], + expected_message: str, + ) -> None: + mock_engine = mocker.MagicMock() + mock_engine.default_sampling_params_list = [OmniDiffusionSamplingParams()] + server = OmniOpenAIServingSpeech.for_diffusion(diffusion_engine=mock_engine, model_name="test-model") + + response = await server.create_speech(OpenAICreateSpeechRequest(input="Hello", extra_params=extra_params)) + + assert response.status_code == 400 + assert expected_message in response.body.decode() + mock_engine.generate.assert_not_called() class TestTTSMethods: diff --git a/tests/model_executor/models/omnivoice/test_mask_dtype.py b/tests/model_executor/models/omnivoice/test_mask_dtype.py index ff899e2a2f1..ff8dda49d7e 100644 --- a/tests/model_executor/models/omnivoice/test_mask_dtype.py +++ b/tests/model_executor/models/omnivoice/test_mask_dtype.py @@ -1,183 +1,79 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""OmniVoice attention masks must follow the model dtype, not a hardcoded float32. +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +"""OmniVoice SDPA fallback masks must not depend on the model dtype. -SDPA requires the additive attention mask to carry the same dtype as the query. -The generator used to materialize that mask as float32 unconditionally, in three -places (the shared per-forward mask, the CUDA-graph capture buffer, and the -replay-time normalization), so every half-precision deployment died on the first -attention call with:: +The previous padded attention path used an additive floating-point mask. SDPA +requires such a mask to have the same dtype as the query, and constructing it +as float32 caused half-precision OmniVoice inference to fail on CUDA. - RuntimeError: invalid dtype for bias - should match query's dtype - -(the exact wording depends on which SDPA backend is selected for the shape; the -math backend words the same rejection as ``attn_mask.dtype``.) - -That left OmniVoice servable only in float32, even though the upstream k2-fsa -implementation runs it in float16. - -Note the split between the CPU and CUDA tests below: the CPU SDPA backend -silently accepts the mismatched mask, so only the CUDA tests can pin the actual -crash. The CPU tests cover the mask contract itself, which is what the fix -changes and what a future regression would break first. +The packed-varlen path no longer needs an additive mask. When SDPA fallback is +required, ``_attention_metadata_from_cu_seqs`` builds a boolean block mask. +Boolean SDPA masks have no dtype coupling with the query, so no mask coercion +is required for float16 or bfloat16 inference. """ from __future__ import annotations import pytest import torch +import torch.nn.functional as F from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( - OmniVoiceAttention, - OmniVoiceGenerator, - _additive_float_mask, + _attention_metadata_from_cu_seqs, ) -from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig HALF_DTYPES = [torch.float16, torch.bfloat16] -ALL_DTYPES = HALF_DTYPES + [torch.float32] cpu_test = pytest.mark.core_model cuda_test = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") -def _tiny_config() -> OmniVoiceConfig: - """A real OmniVoiceConfig sized so the module fits in a unit test.""" - return OmniVoiceConfig( - llm_config={ - "hidden_size": 32, - "num_hidden_layers": 2, - "num_attention_heads": 4, - "num_key_value_heads": 2, - "head_dim": 8, - "intermediate_size": 64, - "vocab_size": 64, - "max_position_embeddings": 128, - }, - enable_cuda_graph=False, - ) - - -def _inputs(dtype: torch.dtype, device: torch.device) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - hidden = torch.randn(2, 6, 32, device=device, dtype=dtype) - bool_mask = torch.ones(2, 1, 6, 6, dtype=torch.bool, device=device) - bool_mask[:, :, :, 4:] = False - rope_table = torch.zeros(6, 8, device=device, dtype=dtype) - rope_table[:, :4] = 1.0 # cos = 1, sin = 0: no rotation, so the test is about the mask - return hidden, bool_mask, rope_table - - -# -------------------------------------------------------------------------- -# The mask contract (CPU) -# -------------------------------------------------------------------------- - - -@cpu_test -@pytest.mark.cpu -@pytest.mark.parametrize("dtype", ALL_DTYPES) -def test_additive_float_mask_uses_requested_dtype(dtype: torch.dtype) -> None: - out = _additive_float_mask(torch.tensor([[True, False]]), dtype) - assert out.dtype == dtype - - -@cpu_test -@pytest.mark.cpu -def test_additive_float_mask_keeps_mask_semantics() -> None: - """Attend positions stay at 0.0 and masked ones at -inf, in half precision too.""" - out = _additive_float_mask(torch.tensor([[True, False]]), torch.float16) - assert out[0, 0].item() == 0.0 - assert out[0, 1].item() == float("-inf") - - -@cpu_test -@pytest.mark.cpu -def test_additive_float_mask_requires_an_explicit_dtype() -> None: - """The float32 default is the bug; keeping the parameter required is the fix.""" - with pytest.raises(TypeError): - _additive_float_mask(torch.tensor([[True]])) # type: ignore[call-arg] - - -@cpu_test -@pytest.mark.cpu -@pytest.mark.parametrize("dtype", ALL_DTYPES) -def test_generator_model_dtype_tracks_its_weights(dtype: torch.dtype) -> None: - """The single source of truth the three former float32 literals now defer to.""" - assert OmniVoiceGenerator(_tiny_config()).to(dtype).model_dtype == dtype - - @cpu_test @pytest.mark.cpu -@pytest.mark.parametrize("dtype", HALF_DTYPES) -def test_attention_accepts_a_model_dtype_mask(dtype: torch.dtype) -> None: - attn = OmniVoiceAttention(_tiny_config()).to(dtype).eval() - hidden, bool_mask, rope_table = _inputs(dtype, torch.device("cpu")) - - with torch.inference_mode(): - out = attn(hidden, rope_table, attention_mask=_additive_float_mask(bool_mask, dtype)) - - assert out.shape == hidden.shape - assert out.dtype == dtype - assert torch.isfinite(out).all() - - -# -------------------------------------------------------------------------- -# The crash itself (CUDA only — the CPU backend does not reject the mismatch) -# -------------------------------------------------------------------------- - - -@cuda_test -@pytest.mark.parametrize("dtype", HALF_DTYPES) -def test_float32_mask_is_what_broke_half_precision(dtype: torch.dtype) -> None: - """Pin the original failure so a float32 default cannot come back unnoticed.""" - attn = OmniVoiceAttention(_tiny_config()).to(device="cuda:0", dtype=dtype).eval() - hidden, bool_mask, rope_table = _inputs(dtype, torch.device("cuda:0")) +def test_sdpa_fallback_mask_is_dtype_independent() -> None: + """The fallback mask is boolean rather than an additive model-dtype mask.""" + cu_seqs = torch.tensor([0, 2, 3, 6, 6], dtype=torch.int32) + + metadata = _attention_metadata_from_cu_seqs( + cu_seqs, + 6, + needs_sdpa_mask=True, + ) - with pytest.raises(RuntimeError, match=r"(invalid dtype for bias|attn_mask)"), torch.inference_mode(): - attn(hidden, rope_table, attention_mask=_additive_float_mask(bool_mask, torch.float32)) + assert metadata.attn_mask is not None + assert metadata.attn_mask.dtype == torch.bool @cuda_test @pytest.mark.parametrize("dtype", HALF_DTYPES) -def test_real_path_forward_runs_in_half_precision(dtype: torch.dtype) -> None: - """The regression, on the production path with the Triton norm kernels live.""" - attn = OmniVoiceAttention(_tiny_config()).to(device="cuda:0", dtype=dtype).eval() - hidden, bool_mask, rope_table = _inputs(dtype, torch.device("cuda:0")) - - with torch.inference_mode(): - out = attn(hidden, rope_table, attention_mask=_additive_float_mask(bool_mask, dtype)) - - assert out.dtype == dtype - assert torch.isfinite(out).all() - +def test_sdpa_fallback_mask_needs_no_half_precision_coercion( + dtype: torch.dtype, +) -> None: + """Half-precision SDPA accepts the production bool mask without casting.""" + device = torch.device("cuda:0") + seq_len = 6 + cu_seqs = torch.tensor([0, 2, 3, 6, 6], dtype=torch.int32, device=device) -@cuda_test -@pytest.mark.parametrize("dtype", HALF_DTYPES) -def test_generator_forward_runs_in_half_precision(dtype: torch.dtype) -> None: - """End-to-end over the real iterative loop, which is where the float32 mask was built. + metadata = _attention_metadata_from_cu_seqs( + cu_seqs, + seq_len, + needs_sdpa_mask=True, + ) - This is the test that fails on an unfixed tree: ``forward`` materialized the - shared SDPA mask as float32 regardless of the weights' dtype. - """ - device = torch.device("cuda:0") - config = _tiny_config() - generator = OmniVoiceGenerator(config).to(device=device, dtype=dtype).eval() + assert metadata.attn_mask is not None + assert metadata.attn_mask.dtype == torch.bool - seq_len, target_len = 12, 4 - input_ids = torch.zeros(2, config.num_audio_codebook, seq_len, dtype=torch.long, device=device) - input_ids[:, 1:, :] = config.audio_mask_id - audio_mask = torch.zeros(2, seq_len, dtype=torch.bool, device=device) - audio_mask[:, seq_len - target_len :] = True - attention_mask = torch.ones(2, 1, seq_len, seq_len, dtype=torch.bool, device=device) + query = torch.randn(1, 2, seq_len, 8, device=device, dtype=dtype) + key = torch.randn_like(query) + value = torch.randn_like(query) with torch.inference_mode(): - tokens = generator( - input_ids, - audio_mask, - attention_mask, - target_lens=[target_len], - seed=0, - num_step=2, + output = F.scaled_dot_product_attention( + query, + key, + value, + attn_mask=metadata.attn_mask, ) - assert tokens.shape == (1, config.num_audio_codebook, target_len) - assert tokens.dtype == torch.long + assert output.dtype == dtype + assert torch.isfinite(output).all() From 5557903a3b1697626fc27ddd962798deb44fdaab Mon Sep 17 00:00:00 2001 From: boatman <109857087+sphinxkkkbc@users.noreply.github.com> Date: Sat, 5 Sep 2026 11:13:08 +0800 Subject: [PATCH 17/17] fix bugs Signed-off-by: boatman <109857087+sphinxkkkbc@users.noreply.github.com> --- .../omnivoice/test_fused_projection_load.py | 27 ++++++++++++------- 1 file changed, 18 insertions(+), 9 deletions(-) diff --git a/tests/model_executor/models/omnivoice/test_fused_projection_load.py b/tests/model_executor/models/omnivoice/test_fused_projection_load.py index fc5db01c7ab..c80f0e65bd2 100644 --- a/tests/model_executor/models/omnivoice/test_fused_projection_load.py +++ b/tests/model_executor/models/omnivoice/test_fused_projection_load.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project """The OmniVoice generator packs q/k/v and gate/up into fused projections. The HF checkpoint stores those five tensors separately, so packing them is a @@ -14,6 +14,8 @@ from __future__ import annotations +from types import SimpleNamespace + import pytest import torch @@ -49,6 +51,13 @@ def _config() -> OmniVoiceConfig: ) +def _generator() -> OmniVoiceGenerator: + return OmniVoiceGenerator( + _config(), + SimpleNamespace(max_num_seqs=1), + ) + + def _checkpoint_shards() -> dict[str, torch.Tensor]: """The five per-layer tensors an OmniVoice checkpoint actually stores.""" torch.manual_seed(0) @@ -64,7 +73,7 @@ def _checkpoint_shards() -> dict[str, torch.Tensor]: def test_every_fused_parameter_is_written() -> None: - generator = OmniVoiceGenerator(_config()) + generator = _generator() shards = _checkpoint_shards() loaded = generator._load_fused_projections(shards) @@ -77,7 +86,7 @@ def test_every_fused_parameter_is_written() -> None: def test_packed_weight_matches_the_shards_it_came_from() -> None: """Packing order is q,k,v and gate,up -- the split in forward() assumes it.""" - generator = OmniVoiceGenerator(_config()) + generator = _generator() shards = _checkpoint_shards() generator._load_fused_projections(shards) @@ -101,7 +110,7 @@ def test_packed_weight_matches_the_shards_it_came_from() -> None: def test_no_fused_parameter_keeps_its_random_init() -> None: """The corruption this guards against: a fused param never written at all.""" - generator = OmniVoiceGenerator(_config()) + generator = _generator() before = { name: param.detach().clone() for name, param in generator.named_parameters() @@ -121,7 +130,7 @@ def test_no_fused_parameter_keeps_its_random_init() -> None: ["self_attn.k_proj", "self_attn.v_proj", "self_attn.q_proj", "mlp.gate_proj", "mlp.up_proj"], ) def test_a_missing_shard_raises_instead_of_loading_partially(dropped: str) -> None: - generator = OmniVoiceGenerator(_config()) + generator = _generator() shards = _checkpoint_shards() del shards[f"llm.layers.0.{dropped}.weight"] @@ -130,7 +139,7 @@ def test_a_missing_shard_raises_instead_of_loading_partially(dropped: str) -> No def test_a_wrong_shaped_shard_raises() -> None: - generator = OmniVoiceGenerator(_config()) + generator = _generator() shards = _checkpoint_shards() shards["llm.layers.0.self_attn.q_proj.weight"] = torch.randn(NUM_HEADS * HEAD_DIM + 1, HIDDEN) @@ -140,7 +149,7 @@ def test_a_wrong_shaped_shard_raises() -> None: def test_a_checkpoint_missing_a_whole_layer_raises() -> None: """Half-loaded is the dangerous state: it neither errors nor works.""" - generator = OmniVoiceGenerator(_config()) + generator = _generator() shards = {k: v for k, v in _checkpoint_shards().items() if not k.startswith("llm.layers.1.")} with pytest.raises(ValueError, match="fused projections"): @@ -149,7 +158,7 @@ def test_a_checkpoint_missing_a_whole_layer_raises() -> None: def test_a_checkpoint_with_no_fused_shards_is_left_alone() -> None: """An unrelated state_dict must not trip the completeness check.""" - generator = OmniVoiceGenerator(_config()) + generator = _generator() assert generator._load_fused_projections({"audio_heads.weight": torch.randn(4, 4)}) == set() @@ -162,7 +171,7 @@ def test_load_weights_wires_the_packing_in(tmp_path, monkeypatch: pytest.MonkeyP """ import safetensors.torch - generator = OmniVoiceGenerator(_config()) + generator = _generator() shards = _checkpoint_shards() before = generator.layers[0].self_attn.qkv_proj.weight.detach().clone()