Conversation
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>
|
Thanks for the PR. The direction is reasonable. Preserving MATH for 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 Before merging, please:
We can help validate the rebased version on an NVIDIA H200 (Hopper, compute capability 9.0 / |
Root cause
decode_n_tokens()unconditionally forcedSDPBackend.MATHfor 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
compileflag throughgenerate_long()andgenerate()todecode_n_tokens(). Eager decoding usesnullcontext(), while compiled decoding still creates a fresh MATH-onlysdpa_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).SDPBackend.MATHSDPBackend.EFFICIENT_ATTENTIONmax 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 iterationpython -c "import fish_speech.models.text2semantic.inference"ruff check fish_speech/models/text2semantic/inference.py tests/test_decode_sdpa_backend.py— no new diagnostics versusmain; the full file retains pre-existing findingsruff 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 frommain🤖 Generated with Claude Code