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/e2e/online_serving/test_omnivoice_expansion.py b/tests/e2e/online_serving/test_omnivoice_expansion.py index 3a042082e63..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 @@ -34,6 +37,10 @@ EXTRA_ARGS = [ "--trust-remote-code", "--disable-log-stats", + "--max-num-seqs", + "8", + "--request-batch-max-wait-ms", + "50", ] TEST_PARAMS = [ OmniServerParams( @@ -42,7 +49,14 @@ server_args=EXTRA_ARGS, ) ] - +STEP_EXECUTION_ARGS = EXTRA_ARGS + ["--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 @@ -58,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.""" @@ -74,6 +115,57 @@ 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: + """Test concurrent request-batch TTS generation.""" + batch_request_config = { + "model": omni_server.model, + "input": get_prompt("text"), + "response_format": "wav", + "seed": 42, + "min_audio_bytes": _DEFAULT_MIN_AUDIO_BYTES, + } + batch_r = openai_client.send_audio_speech_request(batch_request_config, request_num=2) + assert len(batch_r) == 2 + _assert_valid_omnivoice_wav_responses(batch_r) + + +@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) + + @hardware_test(res={"cuda": "L4"}, num_cards=1) + def test_speech_auto_voice_batch_step_execution(self, omni_server, openai_client) -> None: + """Test concurrent step-execution TTS generation.""" + 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, + }, + } + batch_r = openai_client.send_audio_speech_request(batch_request_config, request_num=2) + assert len(batch_r) == 2 + _assert_valid_omnivoice_wav_responses(batch_r) + @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..e9b267f749d --- /dev/null +++ b/tests/e2e/online_serving/test_omnivoice_parity.py @@ -0,0 +1,96 @@ +# 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." + +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, + 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 + + +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", + STAGE_CONFIG, + "--max-num-seqs", + "1", + ] + + +@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/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/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", ): 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_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() 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() 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..6e611f2ff09 --- /dev/null +++ b/tests/model_executor/models/omnivoice/test_pipeline_batching.py @@ -0,0 +1,228 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni 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 +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_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( + 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/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 diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 62ac027ca8b..cdf3d0060d7 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. @@ -12,10 +12,13 @@ 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 dataclasses import dataclass +from typing import Any, ClassVar import numpy as np import torch @@ -26,10 +29,17 @@ 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.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 OmniVoiceGenerator +from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import ( + OmniVoiceGenerator, + _build_cu_seqs, + _get_time_steps, +) from vllm_omni.transformers_utils.configs.omnivoice import OmniVoiceConfig from vllm_omni.utils.speaker_cache import get_speaker_cache @@ -43,6 +53,15 @@ logger = init_logger(__name__) +@dataclass +class _PreparedOmniVoiceRequest: + input_ids: torch.Tensor + audio_mask: torch.Tensor + cond_len: int + 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.""" @@ -132,6 +151,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__() @@ -152,7 +173,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) @@ -205,50 +226,36 @@ def _encode_ref_audio(self, audio_signal: torch.Tensor, sr: int) -> torch.Tensor tokens = tokens.squeeze(0) # [8, T_ref] return tokens - @torch.inference_mode() - def forward(self, req: DiffusionRequestBatch) -> 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 "" + 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" - extra = req.sampling_params.extra_args or {} + voice_name = None 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 + error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" ) else: audio_field = None @@ -264,7 +271,6 @@ def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: 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" @@ -274,105 +280,325 @@ def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: if not text: return DiffusionOutput(error="Empty text prompt") - 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] + text_tokens_tensor = torch.tensor(encoding_ids, dtype=torch.long, device=self.device) + text_len = text_tokens_tensor.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 + cache_key = None if voice_name: - _cache_key = self._speaker_cache.make_cache_key( + 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) + 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 + 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(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()}) + 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) - # 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) + 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=device, + input_ids = torch.cat([cond_ids, uncond_ids], dim=1).transpose(0, 1).contiguous() + + 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, + cond_len=cond_len, + target_len=target_len, + seed=seed, + ) + + def _collate_request_inputs( + self, + prepared_requests: Sequence[_PreparedOmniVoiceRequest], + ) -> 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: + input_id = request.input_ids + audio_mask = request.audio_mask + cond_len = request.cond_len + input_ids.append(input_id) + audio_masks.append(audio_mask) + cond_lens.append(cond_len) + + input_ids = torch.cat(input_ids, dim=0) + audio_masks = torch.cat(audio_masks, dim=0) + + return input_ids, audio_masks, cond_lens + + 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): + raise OmniClientError(prepared.error or "OmniVoice request preparation failed") + + prepared_request = prepared + cond_len = prepared_request.cond_len + target_len = prepared_request.target_len + input_ids = prepared_request.input_ids + audio_mask = prepared_request.audio_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: + seed = random.randint(0, 2**63 - 1) + 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] + 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_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, + ) ) - uncond_ids = torch.cat([uncond_ids, pad], dim=1) + sched.append(int(num)) + rem -= int(num) + schedules = torch.tensor(sched, dtype=torch.long, device=device) - batch_input_ids = torch.stack([cond_ids, uncond_ids]) + layer_ids = torch.arange(num_codebooks, device=device).view(1, -1, 1) + generator = torch.Generator(device=device).manual_seed(seed) - 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 + 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 = guidance_scale + state.extra["schedules"] = schedules + state.extra["layer_ids"] = layer_ids + state.extra["generator"] = generator + state.extra["t_shift"] = t_shift + state.extra["cond_len"] = cond_len + 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): + 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] = [] + 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: + 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.guidance) + generators.append(state.extra.get("generator", None)) + schedules.append(state.extra["schedules"]) + steps.append(state.step_index) + + audio_masks = torch.cat(audio_masks, dim=0) + + B = len(target_lens) + cu_seqs = _build_cu_seqs(cond_lens, target_lens, input_ids.device) + + 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, + 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] + + 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 = cond_lens[i] + t_len = target_lens[i] + + # Extract logits for target region; upcast only the slices we actually consume. + 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) + sample = batch_tokens[i] + sample_tokens = sample[..., :t_len] + 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 + # 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 + state.step_index += 1 + + def post_decode(self, state: StepRequestState, **kwargs: Any): + tokens = state.extra["tokens"] + if tokens.dim() == 2: + tokens = tokens.unsqueeze(0) + audio = self.decoder(tokens) + return DiffusionOutput(output=audio) - 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 + @torch.inference_mode() + 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": "..."} + """ + prepared_requests: list[_PreparedOmniVoiceRequest] = [] + 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): + 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, - attention_mask=batch_attn_mask, - target_lens=[target_len], - num_step=self.num_step, - guidance_scale=self.guidance_scale, + cond_lens=batch_cond_lens, + target_lens=batch_target_len, + 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, class_temperature=self.class_temperature, - seed=seed, + seed=batch_seeds, ) - # Decode tokens to audio - audio = self.decoder(tokens) # [1, 1, samples] - return DiffusionOutput(output=audio) + target_offset = 0 + for i, target_len in enumerate(batch_target_len): + request_tokens = tokens[:, :, target_offset : target_offset + target_len] + audio = self.decoder(request_tokens) + outputs[prepared_indices[i]] = 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/diffusion/worker/utils.py b/vllm_omni/diffusion/worker/utils.py index 5158dbfbf37..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 diff --git a/vllm_omni/entrypoints/openai/serving_speech.py b/vllm_omni/entrypoints/openai/serving_speech.py index 813d916feb7..6512290bdd6 100644 --- a/vllm_omni/entrypoints/openai/serving_speech.py +++ b/vllm_omni/entrypoints/openai/serving_speech.py @@ -3505,6 +3505,25 @@ 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] + + # 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: + 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: + 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 a75c4e6f7fc..965d021b82f 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 @@ -241,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 # --------------------------------------------------------------------------- @@ -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,77 @@ 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 + 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, ) - return self._get_logits(self.norm(hidden_states + residual)) + 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, 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,36 +972,51 @@ 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 - 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: [B, 8, max_target_len] - generated audio tokens + Packed generated audio tokens with shape + ``[1, 8, sum(target_lens)]``. """ 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] - tokens = torch.full( - (B, num_codebooks, max_target_len), - mask_id, - dtype=torch.long, - device=device, - ) + 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() @@ -939,45 +1040,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 @@ -996,7 +1098,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] @@ -1008,10 +1110,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. @@ -1021,8 +1124,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