Skip to content

perf(breeze): keep sampled depth tokens on the device (builds on #966) - #988

Draft
EdwardGong wants to merge 3 commits into
Blaizzy:mainfrom
EdwardGong:feat/breeze-depth-no-sync
Draft

EdwardGong wants to merge 3 commits into
Blaizzy:mainfrom
EdwardGong:feat/breeze-depth-no-sync

Conversation

@EdwardGong

@EdwardGong EdwardGong commented Sep 30, 2026 •

Copy link
Copy Markdown

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

  • _sample is 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_tokens feeds each sampled token straight into the next step's embedding, calls mx.async_eval on 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_logits accept 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_tokens calls the syncing _sample.
  • mlx_audio/tts/tests/test_breeze_tts.py: the two CFG tests now stub _sample_array instead 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).

Checkpoint #966 alone + this PR
bf16 0.541 0.474
8-bit 0.414 0.333

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

jolionlands and others added 3 commits September 30, 2026 16:13
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>
@github-actions

Copy link
Copy Markdown

⚠️ GitHub does not mark 3 commits 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.

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