Skip to content

perf: don't force MATH SDPA backend outside torch.compile - #1337

Open
Ray0907 wants to merge 1 commit into
fishaudio:mainfrom
Ray0907:perf/eager-sdpa-backend
Open

Ray0907 wants to merge 1 commit into
fishaudio:mainfrom
Ray0907:perf/eager-sdpa-backend

Conversation

@Ray0907

@Ray0907 Ray0907 commented Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

Root cause

decode_n_tokens() unconditionally forced SDPBackend.MATH for every autoregressive step, including eager inference even though the CLI defaults to --no-compile. The override was added for Inductor code generation, so eager mode unnecessarily prevented PyTorch from selecting a faster fused backend.

This threads the existing compile flag through generate_long() and generate() to decode_n_tokens(). Eager decoding uses nullcontext(), while compiled decoding still creates a fresh MATH-only sdpa_kernel() context manager inside every loop iteration.

Isolated SDPA-kernel benchmark

Apple MPS, PyTorch 2.11.0, no CUDA. This uses the S2-Pro text model's real attention dimensions (n_head=32, head_dim=128, bfloat16) and the current decode mask key width (max_seq_len=32768).

Backend Time per SDPA call Projected attention time per token (36 layers)
SDPBackend.MATH 17.64 ms 635.18 ms
SDPBackend.EFFICIENT_ATTENTION 6.78 ms 243.98 ms
  • Speedup: 2.6x
  • Projected saving: ~391 ms per generated token across 36 attention layers
  • Output comparison: bit-exact at this scale (max abs diff: 0.0)

This is an isolated attention-kernel microbenchmark, not a full end-to-end model generation benchmark; no ~11 GB checkpoint was downloaded. A maintainer with CUDA hardware is invited to confirm the backend impact there as well.

Scope

This does not change forward_generate() mask slicing or KV-cache width. It is orthogonal to the cache-sizing work in #1327 and #1312.

Tests

  • pytest -q tests/test_decode_sdpa_backend.py — 2 passed; verifies eager mode does not call the MATH-only context and compiled mode does so once per iteration
  • python -c "import fish_speech.models.text2semantic.inference"
  • ruff check fish_speech/models/text2semantic/inference.py tests/test_decode_sdpa_backend.py — no new diagnostics versus main; the full file retains pre-existing findings
  • ruff format --check fish_speech/models/text2semantic/inference.py tests/test_decode_sdpa_backend.py — the new test is formatted; the full file retains one pre-existing unrelated formatting diff from main

🤖 Generated with Claude Code

Preserve the MATH backend override only for torch.compile, where Inductor benefits from code-generating attention.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
@Whale-Dolphin

Copy link
Copy Markdown
Member

Thanks for the PR. The direction is reasonable. Preserving MATH for torch.compile is an important distinction from the earlier change and retains the existing compiled-path behavior.

However, this repository has already had a closely related regression:

This PR may avoid the original issue because it only changes eager decoding, but #926 did not clearly establish that the regression was compile-only.

Also, #1312 has now been merged, and this PR currently conflicts with main. Since #1312 passes only the active KV prefix to SDPA, the original fixed K=32768 microbenchmark is not sufficient to represent typical decoding performance. That remains a useful long-context test case, but we also need results at representative shorter prefix lengths.

Before merging, please:

  1. Rebase onto the latest main and preserve the active-prefix logic and regression tests from perf(text2semantic): pass only active KV cache prefix to SDPA #1312.
  2. Run real-model end-to-end A/B tests using the repository-pinned PyTorch 2.8.0:
    • eager + forced MATH;
    • eager + automatic backend selection;
    • compiled + forced MATH, to confirm the existing compiled path does not regress.
  3. Report hardware/software versions, exact commands and generation parameters, warm tokens/s, cold/first-run latency, and peak GPU memory. If possible, also identify the SDPA backend actually selected.
  4. Include generation-quality checks: finite logits/probabilities, numerical differences with explicit tolerances, argmax/greedy-token comparisons, and a small set of generated audio samples or listening results. We do not expect bitwise-identical outputs across backends; the goal is to check for meaningful numerical or audio-quality regressions.
  5. Confirm that eager decoding no longer forces MATH, compiled decoding still uses MATH, and the active-prefix progression remains correct through the real generation call chain.

We can help validate the rebased version on an NVIDIA H200 (Hopper, compute capability 9.0 / sm_90). However, we currently do not have access to Blackwell hardware, so we cannot validate the RTX 5090 configuration locally. Since the historical regression was reported on both RTX 4090 (Ada) and RTX 5090 (Blackwell), community testing on those GPUs would be particularly valuable. If that hardware coverage is unavailable, we can keep the PR open while maintainers and community members help with validation. This is not a rejection of the optimization; we just want to avoid repeating the historical #859/#926 regression.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants