Skip to content

fix(breeze): keep bf16 prompt embeddings in bf16 - #987

Merged
lucasnewman merged 4 commits into
Blaizzy:mainfrom
EdwardGong:fix/breeze-bf16-dtype
Oct 3, 2026
Merged

lucasnewman merged 4 commits into
Blaizzy:mainfrom
EdwardGong:fix/breeze-bf16-dtype

Conversation

@EdwardGong

Copy link
Copy Markdown
Contributor

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 by mx.array(sqrt(hidden_size)). mx.array of 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 of mx.array(...).
  • mlx_audio/tts/tests/test_breeze_tts.py: test_bf16_prompt_embeddings_keep_the_weight_dtype casts the tiny fixture model to bf16 and checks that _prompt_embeddings returns 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:

  • backbone step: 29.1 ms -> 7.7 ms
  • depth-decoder step: 9.0 ms -> 2.3 ms
  • real-time factor: 2.1 -> 0.61
  • the 8-bit checkpoint is unchanged within noise

Checked on the real checkpoint: model.text_encoder.embed_tokens(...) returns float32 on main and 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

  • Tests added/updated
  • Documentation updated (no user-facing change)
  • Issue referenced (e.g., "Closes #...")

Co-Authored-By: Warp agent@warp.dev

lucasnewman and others added 2 commits October 3, 2026 08:09
_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
lucasnewman force-pushed the fix/breeze-bf16-dtype branch from 5a260a3 to 7fe1c17 Compare October 3, 2026 15:10
@lucasnewman
lucasnewman merged commit e1b19b9 into Blaizzy:main Oct 3, 2026
2 checks passed
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.

Breeze TTS: bf16 checkpoint runs ~3.4x slower because prompt embeddings are promoted to float32

2 participants