fix(breeze): keep bf16 prompt embeddings in bf16 - #987
Merged
Merged
Conversation
This was referenced Sep 30, 2026
_TextEmbedding scaled the embeddings by mx.array(sqrt(hidden)), which is float32. That promoted every prompt activation, and through it the backbone KV cache and all later steps, to float32, so the bf16 checkpoint ran ~3.4x slower than it needs to (backbone step 29.1 -> 7.7 ms, depth step 9.0 -> 2.3 ms on an M5 Max; RTF 2.1 -> 0.6). Multiplying by the Python scalar lets MLX keep the weight dtype. Co-Authored-By: Warp <agent@warp.dev> Co-authored-by: Edward Gong <me@edgong.com>
lucasnewman
force-pushed
the
fix/breeze-bf16-dtype
branch
from
October 3, 2026 15:10
5a260a3 to
7fe1c17
Compare
lucasnewman
approved these changes
Oct 3, 2026
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
The bf16 Breeze TTS 2 checkpoint (
mlx-community/Breeze-TTS-2-mlx) runs about 3.4x slower than it needs to, and slower than real time: a float32 scale in the text embedding promotes the whole generation to float32. With the fix, bf16 runs faster than real time. Closes #986.Description
_TextEmbedding.__call__multiplied the embeddings bymx.array(sqrt(hidden_size)).mx.arrayof a Python float is float32, so for a bf16 checkpoint the prompt embeddings, and through them the backbone activations, the KV cache and every later step, were computed in float32. Multiplying by the plain Python scalar lets MLX keep the weight dtype. Float32 checkpoints behave exactly as before.Changes in the codebase
mlx_audio/tts/models/breeze_tts/breeze_tts.py: scale by the Python scalar instead ofmx.array(...).mlx_audio/tts/tests/test_breeze_tts.py:test_bf16_prompt_embeddings_keep_the_weight_dtypecasts the tiny fixture model to bf16 and checks that_prompt_embeddingsreturns bf16. It fails without the fix (the result is float32).Changes outside the codebase
None.
Additional information
Measured on an M5 Max (mlx 0.32.2), voice-clone mode:
Checked on the real checkpoint:
model.text_encoder.embed_tokens(...)returns float32 onmainand bfloat16 with this change.pytest mlx_audio/tts/tests/test_breeze_*.py: 36 passed; black 26.3.1 and isort 5.13.2 (the pre-commit pins) are clean.Checklist
Co-Authored-By: Warp agent@warp.dev