Skip to content

[Perf][TTS] Add per-row generator support to MOSS-TTS talker to eliminate serial fallback for seeded batched requests - #7922

Merged
NickCao merged 3 commits into
vllm-project:mainfrom
jingchengtian:perf/moss-tts-per-row-generators
Sep 29, 2026
Merged

NickCao merged 3 commits into
vllm-project:mainfrom
jingchengtian:perf/moss-tts-per-row-generators

Conversation

@jingchengtian

@jingchengtian jingchengtian commented Sep 21, 2026 •

Copy link
Copy Markdown
Contributor

1. What This PR Does

1.1 Background: MOSS-TTS Inference Pipeline

The full MOSS-TTS inference pipeline consists of three stages:

User request → backbone (Qwen3 talker, generates text + semantic hidden states)
             → depth transformer (1-layer GPT2, generates 12 audio codes per frame)
             → codec (MOSS-Audio-Tokenizer-v2, decodes audio codes into waveform)

The depth transformer is a 1-layer GPT2-style decoder responsible for generating 12 codebook audio codes (n_vq=12) per audio frame. It works autoregressively: first generate codebook 0, then use the result as input to generate codebook 1, and so on through codebook 11. So each audio frame requires 12 depth transformer forward passes.

1.2 The seed Parameter and torch.Generator

When a user sends a TTS request, they can include a seed parameter for reproducible speech synthesis — the same seed + same text should produce the same audio. vLLM-Omni's OmniGPUModelRunner creates a dedicated torch.Generator object for each seeded request. This object is a deterministic random number stream; when torch.multinomial uses it for sampling, the result is entirely determined by the generator's state.

1.3 The Core Problem: Serial Fallback for B>1 Seeded Requests

When vLLM's scheduler batches B requests in the same step and multiple of them carry seeds, the runner faces a problem:

torch.multinomial only accepts a single generator parameter — you cannot pass B generators at once. If you use one generator for B rows, each row's result would depend on the presence or absence of other rows, breaking per-row reproducibility.

The runner's code (gpu_model_runner.py:1911-1936) handles this as follows:

if (decode_batch_size > 1
    and any(generator is not None for generator in row_generators)
    and not getattr(self.model, "talker_mtp_accepts_per_row_generators", False)):
    # Serial fallback: B separate calls, each processing only 1 row
    for row, req_id in enumerate(decode_req_ids):
        self._talker_mtp_forward([req_id], inputs_embeds, row_offsets)

That is, when talker_mtp_accepts_per_row_generators = False (the original default for MOSS-TTS), the runner splits B rows into B separate _talker_mtp_forward calls. Each call runs 12 depth transformer iterations, so the total is:

B requests × 12 iterations = B×12 serial depth transformer forward passes

Each forward pass processes only 1 row (batch_size=1), severely underutilizing the NPU's matrix multiply units.

1.4 Qwen3-TTS Already Solved This

Qwen3-TTS's talker (qwen3_tts_talker.py:351) already sets talker_mtp_accepts_per_row_generators = True and threads the generators list through its code_predictor. When the runner sees this flag is True, it skips the serial fallback and passes generators=row_generators directly to talker_mtp — one batched call processing all B rows at once.

MOSS-TTS Local never implemented this.

1.5 Unaffected Scenarios

  • Unseeded requests: The runner does not create generators, generators=None, so the original batched multinomial path is used — completely unaffected
  • B=1 seeded requests: The runner uses the scalar generator path (not the generators list), which does not trigger the serial fallback or this PR's logic
  • Standard vllm bench serve: Does not pass seed parameters, so T1 has no effect on standard benchmarks

2. Testing Methodology

2.1 Test Environment

  • Hardware: Ascend 910B2C (single NPU, device 0)
  • Model: MOSS-TTS-Local-Transformer-v1.5
  • Deploy config: moss_tts_local.yaml
  • Baseline code: vllm-omni upstream/main latest master

2.2 Strict Variable Control

Fixed Randomness

  • Each request's seed = 42 + i (i is the request index)
  • Generators reset to the same seeds before each round
  • Same 64 texts and same reference audio (common_voice_en_125386.wav) used throughout

Warmup Before Measurement

  • Single-module: 10 warmup iterations + 30 timed iterations per round, with torch.npu.synchronize() before and after each call
  • E2E: 5 rounds per config, discard Round 1 (R1, AscendCompiler cold compilation), report median of R2-R5

Multiple Rounds Averaged

  • 5 rounds, discard R1, report median of R2-R5 (more robust to outliers than mean)
  • Server health checked between rounds

2.3 Single-Module Benchmark

Results

B BEFORE serial (ms) AFTER batched (ms) Speedup
1 6.982 7.436 0.94x
8 54.121 8.537 6.34x
16 110.368 11.809 9.35x
32 217.807 19.265 11.31x
64 440.204 34.577 12.73x

2.4 E2E Benchmark

Results

C BEFORE serial (s) AFTER batched (s) Speedup
1 1.20 1.22 0.98x (neutral)
8 5.11 2.65 1.93x
16 9.09 3.83 2.37x
32 16.69 6.19 2.70x
64 31.73 9.36 3.39x

Note on B=1 Single-Module Result

The single-module benchmark shows a 6.5% slowdown at B=1 (7.436ms vs 6.982ms). This is not a regression caused by this optimization's code logic — it is an artifact of the benchmark harness. The benchmark script forces the B=1 AFTER case to pass generators=[g] (a list), which routes _sample_token into the per-row multinomial branch (Python loop + torch.cat for 1 element). In real serving, the runner passes a scalar generator for B=1 (gpu_model_runner.py:2000-2002), so _sample_token takes the original else branch (direct multinomial, no list overhead). The E2E C=1 result confirms this: 1.22s vs 1.20s (1.7%, within noise) — the per-round values overlap between BEFORE and AFTER, proving no real difference at B=1.

@vllm-omni-review-bot

Copy link
Copy Markdown

This PR appears to belong to: docs/design/module/model_integration.md, docs/design/module/ar_runtime.md.

Module owners: @gcanlin @Sy0307 @tzhouam

Routing: @gcanlin via module of the changed files, CODEOWNERS; @Sy0307 via module of the changed files, CODEOWNERS; @tzhouam via module of the changed files, CODEOWNERS

@jingchengtian, please review your own changes and leave a short self-review comment describing what you checked. PRs without author self-review may not be assigned a reviewer.

Please take a look when you have a chance. If you would like an automated review, mention @vllm-omni-review-bot in a comment.

@jingchengtian
jingchengtian force-pushed the perf/moss-tts-per-row-generators branch 2 times, most recently from be6d1c0 to d138827 Compare September 21, 2026 08:55
@jingchengtian

Copy link
Copy Markdown
Contributor Author

Self-Review (precheck-pr, full mode)

Branch: perf/moss-tts-per-row-generators (based on upstream/main 1b87115)
Type: Performance | Diff: 3 files, +31/-3 lines

Summary

Dimension Result
PR title format ✓ [Perf][TTS] prefix with model identifier
Code quality (5 patterns) ✓ 0 ⚠, 0 ✗ — no kwargs plumbing, broad except, Any hints, hot-path clone, or event-loop blocking
Examples policy ✓ No new Python examples
Simplification ✓ Per-row multinomial loop is minimal — torch.multinomial only accepts single generator
Dead code ✓ else branch preserves backward compat for unseeded/B=1 path
Import hygiene ✓ Only collections.abc.Sequence (stdlib)
SPDX headers ✓ All 3 files already have SPDX headers, no new files
Forbidden imports ✓ No pickle/re/base64/Triton/HF Hub
torch.cuda ✓ No new torch.cuda.* call sites
DCO ✓ Signed-off-by present
Accuracy ✓ Per-row reproducibility verified (same seed → same total_bytes)
Benchmark ✓ Single-module (B=1,8,16,32,64) + E2E (C=1,8,16,32,64), 5 rounds, median of R2-R5
Warmup ✓ R1 discarded, 10-iter warmup for single-module
Hardware specified ✓ Ascend 910B2C
No unexplained regressions ✓ B=1/C=1 neutral — runner passes scalar generator for B=1
Pre-commit/mypy ⚠ Not run locally (no pre-commit installed; CI will run)
VRAM ⚠ Not measured (compute-only change, no new allocations)

Verdict: 0 blocking | 2 warnings (environment limitations) | ready for review

@vllm-omni-review-bot

vllm-omni-review-bot commented Sep 21, 2026 •

Copy link
Copy Markdown

Omni ReviewBot triage note

Automated triage of commit f34ae58565d2 produced:

  • Priority: high. Prompt maintainer attention is suggested.

These are automated triage suggestions only — the final decision belongs to the maintainers.

@jingchengtian
jingchengtian force-pushed the perf/moss-tts-per-row-generators branch from d138827 to f496647 Compare September 21, 2026 09:07

@NickCao NickCao left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

On CUDA, the model is marked graph-safe:

self.talker_mtp_graph_safe = not current_omni_platform.is_npu()

That causes the runner to wrap talker_mtp in a CUDA graph.

The sequence is:

  1. During startup, CUDA captures talker_mtp using dummy inputs and no generators.
  2. A seeded request later passes:
    generators=[generator_for_request_1, generator_for_request_2]
  3. The CUDA graph wrapper replays the already-captured GPU graph. It does not call the Python talker_mtp function again, so it never reads those generators.

Thus is not working with cuda graph enabled.

@hsliuustc0106 hsliuustc0106 added enhancement New feature or request tts code related to tts models labels Sep 22, 2026
@jingchengtian

jingchengtian commented Sep 22, 2026 •

Copy link
Copy Markdown
Contributor Author

@NickCao thanks for the review - you are right.

When a batch carries explicit per-row seeds, the runner can no longer replay a graph that was captured without generators. I pushed a follow-up commit (76ea369) that fixes this at the runner level instead of disabling the graph.

What changed (_talker_mtp_forward in gpu_model_runner.py):

  • Compute the per-row generators before resolving the cudagraph mode.
  • If any row of the decode batch carries an explicit generator (tts_local_seed), force CUDAGraphMode.NONE (eager) with unpadded tokens, so talker_mtp actually runs as Python and consumes the generators.
  • Generator-free batches keep the pre-captured graph replay, so unseeded throughput is unchanged.
  • It's a runner-level, per-batch decision, so it applies to CUDA and Ascend (where talker_mtp is ACL-graph wrapped too).

Why this is better than the Qwen3-TTS approach:
Qwen3-TTS gates the capability in the model (talker_mtp_accepts_per_row_generators = not graph_wrapped, qwen3_tts_talker.py:354): under full cudagraphs it simply drops seeding (silently unreproducible), which its own TODO acknowledges. My fix keeps both: seeded batches stay reproducible and batched (no serial fallback), while unseeded batches still get graph replay.

The pushed function body is byte-identical to the version E2E-verified on Ascend 910B2C:

  • no-seed request: 200, 11 talker_mtp graphs still captured/replayed
  • 3 concurrent seeded requests (different seeds): all 200 in ~1s (parallel)
  • same seed repeated: byte-identical output; different seed: different output

I'd love your take on whether this per-batch eager gating is acceptable vs. keeping the model-side not is_npu() guard.

@jingchengtian

Copy link
Copy Markdown
Contributor Author

Runner fix moved to #8009

Following NickCao's review, the runner-side change that was previously on this branch (force eager talker_mtp when a batch carries explicit seeds) has been split out into a standalone PR: #8009

This PR now contains only the MOSS-TTS model-side per-row generator support (local sampler, local_depth, talker). The runner fix in #8009 is generic: once merged, MOSS-TTS and Qwen3-TTS both honor per-row generators on graph-wrapped platforms. Latest main (PR #7781) already sets talker_mtp_accepts_per_row_generators = True for Qwen3-TTS, so that model needs no further change here.

@jingchengtian

Copy link
Copy Markdown
Contributor Author

@NickCao thanks for the review. I've reworked the PR exactly as you suggested:

What changed: the runner-side graph-replay fix is no longer in this PR. It's been split out into a standalone generic runner PR #8009. This PR now contains only the MOSS-TTS model-side per-row generator support (MossTTSLocalGPTQ._sample_token, MossTTSLocalDepthGPTQ.generate_frame, MossTTSTalker property + generator plumbing).

Relationship between the two PRs:

Both are rebased on latest main. #8009's CI is green (pre-commit, build, DCO). End-to-end verification for MOSS-TTS is in the PR description (seeded batches drop from ~3s serial to ~1s parallel, same-seed is byte-identical, no-seed graph replay unchanged).

Could you take another look and merge if you're happy? Happy to adjust anything else.

@NickCao NickCao left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM now that #8009 has landed, one nitpick: the length of generators should be validated, as in

def _normalize_generators(
generator: _GeneratorLike, batch_size: int
) -> torch.Generator | list[torch.Generator | None] | None:
if generator is None or isinstance(generator, torch.Generator):
return generator
row_generators = list(generator)
if len(row_generators) != batch_size:
raise ValueError(f"Expected {batch_size} per-row generators, but got {len(row_generators)}.")
return row_generators

@jingchengtian

Copy link
Copy Markdown
Contributor Author

@NickCao thanks for the nitpick — fixed in 4efd513.

What changed

generators is now length-checked against the batch it samples for, exactly like Qwen3CodePredictor._normalize_generators (qwen3_code_predictor.py:830):

  • new _normalize_generators(generators, batch_size) in modeling_moss_tts_local.py, raising ValueError(f"Expected {batch_size} per-row generators, but got {len(row_generators)}.") on a mismatch;
  • _sample_token calls it instead of the old generators[row] if row < len(generators) else None fallback, so a short list can no longer quietly sample the tail rows from the global RNG;
  • talker_mtp runs the same check on entry (bsz is known there), so a mis-sized batch fails before the 12-step depth loop rather than after it.

The validation sits before the any(gen is not None …) branch, so it also fires for an all-None list; that list still falls through to the single batched multinomial, leaving the unseeded and B=1 paths byte-for-byte unchanged.

Tests — new tests/model_executor/models/moss_tts/test_per_row_generators.py (CPU, core_model): length mismatch raises for both short and long lists, per-row reproducibility, a row's codes are unaffected by its batch neighbours, the all-None list matches the scalar batched path, and a single seeded row stays reproducible. tests/model_executor/models/moss_tts/, tests/worker/test_omni_gpu_model_runner.py, tests/worker/test_tts_sampling_seed.py, tests/entrypoints/openai/test_audio8_tts_adapter.py: 119 passed (the cuda-marked cases need a CUDA build and were deselected).

End-to-end on Ascend 910B2C (MOSS-TTS-Local-Transformer-v1.5, real weights, on top of #8009):

  • a mis-sized list surfaces through the real path: generate_frame with 7 generators for B=8 → ValueError: Expected 8 per-row generators, but got 7.;
  • the guarantee itself, checked directly on-device with the real MossTTSLocalDepthTransformer: for 8 consecutive frames, a B=8 batched call with per-row generators produces codes bit-identical to running each row alone with its own generator (10/10 checks);
  • served batches: C=4/8/16 seeded requests all return 200 with valid audio; wall time vs. the serial fallback (same build with talker_mtp_accepts_per_row_generators = False) is 1.61s vs 2.40s at C=4 and 2.81s vs 7.53s at C=16, i.e. the serial path is gone.

One thing I did not claim: that a batched seeded request is byte-identical to serving the same request alone. That does not hold end-to-end on any build, with or without this PR — the Qwen3 backbone's batched prefill logits differ at the last bit from the B=1 ones, and the codes diverge from there. The depth-transformer sampling this PR touches is reproducible (bit-identical per row, as measured above); the upstream hidden state it consumes is what varies.

The branch is still based on the pre-#8009 merge-base and merges cleanly into current main. Happy to rebase if you prefer it on top.

Could you take another look and merge if you're happy? Happy to adjust anything else.

@jingchengtian

Copy link
Copy Markdown
Contributor Author

a02498d1 fixes the two pre-commit failures on 4efd51307 — the SPDX header on modeling_moss_tts_local.py / the new test file, and a mis-sized typo in a comment. Headers and comments only, no functional change. DCO, pre-commit and the 3.11/3.12 builds are green on it now.

@NickCao NickCao added the ready label to trigger buildkite CI label Sep 29, 2026
@NickCao
NickCao enabled auto-merge (squash) September 29, 2026 15:20
@NickCao

NickCao commented Sep 29, 2026

Copy link
Copy Markdown
Collaborator

Oops, need a rebase.

jingchengtian and others added 2 commits September 29, 2026 11:50
…nate serial fallback for seeded batched requests

When B>1 seeded requests are batched, the runner serially processes each row
(B×12 serial depth-transformer forward passes) unless the model sets
talker_mtp_accepts_per_row_generators. This PR sets that flag and threads
generators through _sample_token → generate_frame → talker_mtp, matching
the Qwen3-TTS pattern (qwen3_tts_talker.py:351).

Each row uses its own torch.Generator for multinomial, preserving per-row
reproducibility. The expensive topk/softmax/nucleus filtering stays batched.

AI assistance: Used opencode to draft the code, benchmark scripts, and PR
description. I reviewed all changes and verified benchmark results from
actual NPU runs. Validation: single-module + E2E benchmarks at B/C=1,8,16,32,64.

Signed-off-by: jingchengtian <jingchengtian@users.noreply.github.com>
Signed-off-by: jingchengtian <tjc1995@126.com>
A ``generators`` list shorter than the batch silently sampled the extra rows
from the global RNG, so a mis-sized batch would quietly lose per-row
reproducibility. Validate the length up front instead, mirroring
``Qwen3CodePredictor._normalize_generators``: ``_sample_token`` raises
``ValueError`` on a mismatch and ``talker_mtp`` fails on entry (before the
n_vq-step depth loop) when the batch does not match.

The all-``None`` case still falls through to the single batched
``multinomial``, so unseeded and B=1 requests keep the original path.

Tests: new tests/model_executor/models/moss_tts/test_per_row_generators.py
(CPU, 6 cases) covers the length check plus per-row reproducibility, batch
composition independence and the unseeded fall-through. Validated on Ascend
910B2C with MOSS-TTS-Local-Transformer-v1.5: seeded batches of 4/8/16 return
200 with valid audio, C=4 1.61s vs 2.40s and C=16 2.81s vs 7.53s against the
serial fallback, and a batched generate_frame reproduces per-row codes
bit-identically to running each row alone.

Signed-off-by: jingchengtian <jingchengtian@users.noreply.github.com>
Signed-off-by: jingchengtian <tjc1995@126.com>
…tion

Two hooks rejected the previous commit:

- check-spdx-header: modeling_moss_tts_local.py carried no SPDX header and
  the new test file copied the stale "vLLM project" copyright line. The hook
  rewrites both to the vLLM-Omni lines.
- typos: "mis-sized" in the talker_mtp comment is not a word it knows.

Comment and header only, no functional change.

Signed-off-by: jingchengtian <jingchengtian@users.noreply.github.com>
@NickCao
NickCao force-pushed the perf/moss-tts-per-row-generators branch from a49e095 to f34ae58 Compare September 29, 2026 15:50
@NickCao NickCao added ready label to trigger buildkite CI and removed ready label to trigger buildkite CI labels Sep 29, 2026
@NickCao
NickCao disabled auto-merge September 29, 2026 17:13
@NickCao
NickCao merged commit 43ef1e6 into vllm-project:main Sep 29, 2026
8 of 9 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request ready label to trigger buildkite CI tts code related to tts models

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants