Skip to content

perf(breeze): cache the depth decoder KV per frame (RTF 1.63 -> 0.62) - #966

Open
jolionlands wants to merge 2 commits into
Blaizzy:mainfrom
jolionlands:breeze-cached-depth-decoder
Open

jolionlands wants to merge 2 commits into
Blaizzy:mainfrom
jolionlands:breeze-cached-depth-decoder

Conversation

@jolionlands

Copy link
Copy Markdown

Context

Breeze TTS 2 is slower than real time on Apple silicon today: RTF 1.63-1.66 on an M4 Pro with the 8-bit MLX weights and a fixed text. Profiling attributes almost all of it to the depth decoder, which re-runs the whole depth stack over the growing prefix at every one of the 15 codebook steps of a frame - 118.8 ms of a 130 ms frame, against 0.4 ms for the backbone that already uses a KV cache. Giving the depth decoder the same treatment is what takes the model from "cannot be served in real time" to RTF 0.62 and 1.28 s to first audio.

Description

  • _DepthModel gains make_cache(), embed_codebook_token() and step(): a per-frame cache, a single-position forward, and an embed helper that keeps the codebook-specific vocabulary offset used by __call__.
  • _DepthDecoder gains start_frame() (seed position zero with the backbone hidden state) and step_logits() (feed one token, read the position just written).
  • Model._depth_tokens() keeps its loop, its reserved-id masking and its sampler calls exactly as before; it now asks step_logits for one step instead of next_logits for the whole prefix. next_logits is untouched and still valid public API.
  • The CFG path (instruct + cfg_scale) gets its own cache, so voice design speeds up too: RTF 2.99 -> 1.00.

On indexing, since that is the whole risk of this patch: next_logits reads the hidden state of the last prefix position and multiplies it by codebooks_head.weight[m-1], where the prefix at step m is [backbone_state, first, s1 .. s_{m-1}]. Step m must therefore feed the token occupying that final position - first for m=1, otherwise the previous step's sample - and read the position just written. Getting this off by one still runs and still looks fast: an earlier attempt of mine reported RTF 0.60 while emitting 42 s of babble for a 30-word sentence, with the top logit landing on the reserved pad id. That is why the new tests compare against the recompute arithmetic rather than celebrating a speed number.

Changes in the codebase

  • mlx_audio/tts/models/breeze_tts/breeze_tts.py: +70/-6.
  • mlx_audio/tts/tests/test_breeze_depth_cache.py (new, 218 lines): six download-free tests built on the existing tiny_config() pattern, with a randomised output head and a non-vacuity guard - cached logits vs the prefix-recompute form (atol 1e-5 plus identical argmax at every step), token equality across 4 seeds x 3 first-codebooks, _depth_tokens equality, CFG-path equality, one cache position per step, and frame independence (no state carried between frames, the classic cache-bug class). Reintroducing the off-by-one makes 3 of these fail, so they are not decorative.
  • mlx_audio/tts/tests/test_breeze_tts.py (+14/-5): two existing tests had to change, and this is the part to look at hardest. test_depth_cfg_applies_to_every_remaining_codebook and test_depth_cfg_masks_reserved_tokens_at_every_step monkeypatched depth_decoder.next_logits, which the cached loop no longer calls. They now patch start_frame/step_logits; their assertions are unchanged (CFG applied at every step = 3.0; reserved ids masked to -inf at every step). Their intent - that CFG and masking happen at every depth step - is exactly what they still verify.

Changes outside the codebase

None.

Additional information

Measured on a Mac mini M4 Pro, mlx-community/Breeze-TTS-2-mlx-8bit, fixed text, warmups discarded, 3 runs:

metric before after
RTF, plain 1.63-1.66 0.61-0.62
RTF, instruct + cfg_scale 2.0 2.99 0.99-1.01
time to first audio (streaming, 2 s interval) 3.28-3.69 s 1.23-1.28 s
peak memory 10.2 GB 6.4-8.4 GB
ASR round-trip WER (both modes) 0.000 0.000

Equivalence, stated plainly. The cached walk is the same arithmetic with a different matmul shape, so in fp32 the two are not bit-identical:

4-bit 8-bit bf16
max per-step logit delta 1.45e-2 1.14e-2 1.42e-2
argmax flips, greedy, all steps 0/16 0/16 0/16

That this is accumulation order and not a systematic difference is not an assumption. On a tiny fixture evaluated in float64 on the CPU the two walks agree exactly - 0.00e+00 at every step - and in float32 they differ by 2e-7 to 7e-7 against a top-2 logit gap of 0.21. The arithmetic is therefore identical and only the summation order changes; the ~1.4e-2 seen on the real checkpoint is that rounding acting on a much larger, quantised weight matrix. bf16 matching the 8-bit and 4-bit builds also rules quantisation out as a cause. Under greedy sampling the walks agree at every step (0/16 argmax flips on all three builds and on the CFG path); under the default sampler, trajectories can diverge after ~20 frames because the loop amplifies those deltas, so the honest claim is "same text, WER 0.000", not "identical audio", and the sampled path is deliberately not seed-stable against the old code.

Which test protects what, since the tolerances differ by design: the 1e-5 logits comparison is a supporting check on a fixture where fp32 error is ~1e-6. The checks that actually stand between you and the off-by-one class of bug are the token-equality test (cached picks against the reference walk, 4 seeds x 3 first-codebooks) and the frame-independence test. I verified those can fail rather than assuming it: reintroducing the off-by-one in step_logits makes 3 of the 6 tests fail, and a deliberately zeroed head trips the non-vacuity guard (the fixture randomises the head for exactly that reason - the real _CodebooksHead is zero-initialised, so an un-randomised fixture compares all-zero logits and passes no matter what).

next_logits is unchanged and still covered by the existing test_depth_decoder_predicts_each_remaining_codebook; the two re-pointed tests moved to the new seam because that is where the loop now calls in. The cached walk is per-frame by construction - first_codebook is a scalar and every frame gets a fresh cache - so a batched call cannot reach it.

Suite impact, same machine and venv: mlx_audio/tts/tests goes from 716 passed / 2 failed to 722 passed / 2 failed, and package collection from 1671 to 1677 tests with the same 4 pre-existing errors in stt tests. The two failures are pre-existing ModuleNotFoundErrors in TestSparkTTSModel/TestIndexTTS. The new file also passes with MLX_DEFAULT_DEVICE=cpu, so it is not skipped on Metal-less CI.

Repro is three lines: load the 8-bit model, run generate() on a fixed text, and divide wall time by audio duration; the streaming number comes from stream=True, streaming_interval=2.0 and timing the first yielded chunk. Happy to add a benchmark script to examples/ if you want one in-tree.

Checklist

  • Tests added/updated
  • Documentation updated
  • Issue referenced - I did not find an existing issue for Breeze speed; happy to open one if you would rather track it that way.

The depth decoder re-ran its full stack over the growing prefix at every one
of the 15 codebook steps of a frame. Profiling the 8-bit checkpoint puts that
at ~91% of generation wall time (118.8 ms of a 130 ms frame, against 0.4 ms
for the already-cached backbone).

It now walks one position per step against a per-frame KV cache: RTF 1.63 ->
0.62 and first audio 3.4 s -> 1.28 s on an M4 Pro, with the CFG path at 1.00
(from 2.99). next_logits is unchanged; the loop uses new start_frame and
step_logits primitives, and test_breeze_depth_cache.py pins their arithmetic
against the prefix-recompute form, including frame independence.
@github-actions

Copy link
Copy Markdown

⚠️ GitHub does not mark 1 commit in this PR as Verified.

Please sign every commit, then update the PR. You can review the commits on the commits tab and follow GitHub's commit-signing guide if needed.

@jolionlands
jolionlands force-pushed the breeze-cached-depth-decoder branch 2 times, most recently from 116eb84 to 3fbebc1 Compare September 19, 2026 19:17
@lucasnewman

Copy link
Copy Markdown
Collaborator

@jolionlands Looks good but this needs signed commits to be merged.

@EdwardGong

Copy link
Copy Markdown

Thanks for this. Some numbers from an M5 Max (mlx 0.32.2), Breeze clone mode, about 15 s of audio:

One test fails on this machine: test_cached_walk_reproduces_prefix_recompute_logits. The cached walk and the prefix recompute differ by up to ~3.2e-3 on the GPU, against atol=1e-5. On the CPU the same comparison agrees to ~7e-7, so the cache arithmetic is right. The single-query cached step and the full-prefix pass just go through different GPU attention kernels. A fix that keeps the check strict is to run that one comparison on the CPU (with mx.stream(mx.cpu):): 0a71625. Feel free to cherry-pick it. The token-level tests pass on the GPU as they are.

I've also opened #988 as a draft on top of this PR. It removes the per-token .item() sync in the depth loop, giving bf16 0.54 -> 0.47 and 8-bit 0.41 -> 0.33 with identical tokens. I'll rebase it once this merges.

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.

4 participants