From df761b32278aa21da3a950fb025efe804eef7ec2 Mon Sep 17 00:00:00 2001 From: Lucas Newman Date: Wed, 23 Sep 2026 19:00:14 -0700 Subject: [PATCH] fix(breeze): keep bf16 prompt embeddings in bf16 _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 Co-authored-by: Edward Gong --- mlx_audio/tts/models/breeze_tts/breeze_tts.py | 2 +- mlx_audio/tts/tests/test_breeze_tts.py | 16 ++++++++++++++++ 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/mlx_audio/tts/models/breeze_tts/breeze_tts.py b/mlx_audio/tts/models/breeze_tts/breeze_tts.py index 3d5288f60..6b9fbbd95 100644 --- a/mlx_audio/tts/models/breeze_tts/breeze_tts.py +++ b/mlx_audio/tts/models/breeze_tts/breeze_tts.py @@ -201,7 +201,7 @@ def __init__(self, vocab_size: int, hidden_size: int, eoi_token_index: int): self.eoi_token_index = eoi_token_index def __call__(self, input_ids: mx.array) -> mx.array: - embeds = self.weight[input_ids] * mx.array(self.weight.shape[-1] ** 0.5) + embeds = self.weight[input_ids] * (self.weight.shape[-1] ** 0.5) return mx.where( (input_ids == self.eoi_token_index)[..., None], self.eoi_embedding, embeds ) diff --git a/mlx_audio/tts/tests/test_breeze_tts.py b/mlx_audio/tts/tests/test_breeze_tts.py index 319d47e60..10eb32971 100644 --- a/mlx_audio/tts/tests/test_breeze_tts.py +++ b/mlx_audio/tts/tests/test_breeze_tts.py @@ -267,6 +267,22 @@ def text_ids(value): assert prompts == ["[S2]reference text", "[S2]warmtarget"] +def test_bf16_prompt_embeddings_keep_the_weight_dtype(monkeypatch): + # A float32 scale in the text embedding used to promote every prompt + # activation, and with it the backbone KV cache, to float32, which made + # bf16 checkpoints ~3.4x slower than necessary. + model = Model(tiny_config()) + model.set_dtype(mx.bfloat16) + monkeypatch.setattr( + model, "_text_ids", lambda _text: mx.array([1, 2, 3], dtype=mx.int32) + ) + embeds = model._prompt_embeddings( + "target", voice=None, instruct=None, ref_audio=None, ref_text=None + ) + assert model.text_encoder.embed_tokens.weight.dtype == mx.bfloat16 + assert embeds.dtype == mx.bfloat16 + + def test_stream_flushes_at_exact_interval_and_resets_state(monkeypatch): model = Model(tiny_config())