perf(breeze): cache the depth decoder KV per frame (RTF 1.63 -> 0.62) - #966
jolionlands wants to merge 2 commits into
Conversation
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.
|
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. |
116eb84 to
3fbebc1
Compare
|
@jolionlands Looks good but this needs signed commits to be merged. |
|
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: I've also opened #988 as a draft on top of this PR. It removes the per-token |
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
_DepthModelgainsmake_cache(),embed_codebook_token()andstep(): a per-frame cache, a single-position forward, and an embed helper that keeps the codebook-specific vocabulary offset used by__call__._DepthDecodergainsstart_frame()(seed position zero with the backbone hidden state) andstep_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 asksstep_logitsfor one step instead ofnext_logitsfor the whole prefix.next_logitsis untouched and still valid public API.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_logitsreads the hidden state of the last prefix position and multiplies it bycodebooks_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 -firstfor 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 existingtiny_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_tokensequality, 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_codebookandtest_depth_cfg_masks_reserved_tokens_at_every_stepmonkeypatcheddepth_decoder.next_logits, which the cached loop no longer calls. They now patchstart_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:Equivalence, stated plainly. The cached walk is the same arithmetic with a different matmul shape, so in fp32 the two are not bit-identical:
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_logitsmakes 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_CodebooksHeadis zero-initialised, so an un-randomised fixture compares all-zero logits and passes no matter what).next_logitsis unchanged and still covered by the existingtest_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_codebookis 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/testsgoes 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 instttests. The two failures are pre-existingModuleNotFoundErrors inTestSparkTTSModel/TestIndexTTS. The new file also passes withMLX_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 fromstream=True, streaming_interval=2.0and timing the first yielded chunk. Happy to add a benchmark script toexamples/if you want one in-tree.Checklist