perf(breeze): keep sampled depth tokens on the device (builds on #966) - #988
Draft
EdwardGong wants to merge 3 commits into
Draft
EdwardGong wants to merge 3 commits into
EdwardGong wants to merge 3 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.
test_cached_walk_reproduces_prefix_recompute_logits failed on an M5 Max (mlx 0.32.2): the cached walk and the prefix recompute differed by ~3e-3. The single-query cached step and the full-prefix pass use different GPU attention kernels; on the CPU the same comparison agrees to ~1e-6, so the cache arithmetic is right and only the tolerance was hardware-dependent. Run this exact-logits check on the CPU; the token-level tests still exercise the GPU. Co-Authored-By: Warp <agent@warp.dev>
_depth_tokens read every sampled codebook back to the host with .item() before building the next step, stalling the GPU 15 times per frame. Sample into device arrays instead (new _sample_array; _sample wraps it with .item()), feed each token straight into the next step's embedding, async_eval as we go, and read the frame back once. The RNG is drawn in the same order, so seeded output is unchanged: a new test compares seeded stochastic frames with the per-token loop, with and without CFG. On an M5 Max, clone mode, 15 s of audio: bf16 RTF 0.54 -> 0.48, 8-bit 0.41 -> 0.34, with identical tokens. Co-Authored-By: Warp <agent@warp.dev>
|
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. |
1 of 3 tasks
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Context
Follow-up to #966, and draft until #966 lands. With #966's per-frame KV cache in place, the Breeze depth decoder is still about 70% of frame time, and part of that is waiting: the loop reads every sampled codebook back to the host with
.item()before it builds the next step, so the GPU stalls 15 times per 80 ms frame. Removing those stalls brings bf16 to RTF 0.47 and 8-bit to 0.33 on an M5 Max.Only the last commit is new (dff8654). The first two are #966's commit, unchanged (095499f, author jolionlands), and a test fix for #966 that I've proposed there (0a71625). I'll rebase once #966 merges.
Description
_sampleis split into_sample_array, which returns the sampled ids as a device array, and_sample, which wraps it with.item()for the backbone's first-codebook/EOS decision._depth_tokensfeeds each sampled token straight into the next step's embedding, callsmx.async_evalon it so the GPU starts while the next step's graph is built, and reads the frame back once at the end.embed_codebook_token/step_logitsaccept a one-element array as well as an int.The RNG is drawn in the same order as before, so seeded output doesn't change. On the real checkpoints (clone mode, 15 s of audio, seed 7), bf16 and 8-bit produced token-identical frames compared with the per-token loop.
Changes in the codebase
mlx_audio/tts/models/breeze_tts/breeze_tts.py: the changes above.mlx_audio/tts/tests/test_breeze_depth_cache.py:test_depth_tokens_stay_on_device_and_keep_the_rng_order, with and without CFG. It compares seeded stochastic frames (temperature 1) against a per-token reference loop, checks that the seeds really pick different frames, and fails if_depth_tokenscalls the syncing_sample.mlx_audio/tts/tests/test_breeze_tts.py: the two CFG tests now stub_sample_arrayinstead of_sample.Changes outside the codebase
None.
Additional information
M5 Max, mlx 0.32.2, clone mode, about 15 s of audio, best of 2 runs. These numbers include the bf16 dtype fix from #987; without it, bf16 is about 3.4x slower (RTF ~2.1, see #986).
Per frame (bf16), the depth decoder goes from 34.0 ms to 28.6 ms.
pytest mlx_audio/tts/tests/test_breeze_*.py: 43 passed; black 26.3.1 and isort 5.13.2 are clean.Checklist
Co-Authored-By: Warp agent@warp.dev