diff --git a/mlx_audio/tts/models/breeze_tts/breeze_tts.py b/mlx_audio/tts/models/breeze_tts/breeze_tts.py index 3d5288f60..f10ad63d7 100644 --- a/mlx_audio/tts/models/breeze_tts/breeze_tts.py +++ b/mlx_audio/tts/models/breeze_tts/breeze_tts.py @@ -530,6 +530,33 @@ def __init__(self, config: ModelConfig): ] self.norm = nn.RMSNorm(self.hidden_size, eps=args.rms_norm_eps) + def make_cache(self) -> list[KVCache]: + """A fresh KV cache for one frame of depth decoding.""" + return [KVCache() for _ in self.layers] + + def embed_codebook_token( + self, token_id: Union[int, mx.array], codebook_idx: int + ) -> mx.array: + """Embed one codebook token at its codebook-specific vocabulary offset. + + Mirrors the offset scheme in :meth:`__call__`, where the token at + sequence position ``i`` carries codebook index ``i - 1``. ``token_id`` + may be a one-element array, so a sampled token can be fed back without + reading it to the host. + """ + if not isinstance(token_id, mx.array): + token_id = mx.array(token_id) + ids = token_id.astype(mx.int32).reshape(1, 1) + codebook_idx * self.vocab_size + return self.embed_tokens(ids) + + def step(self, embeds: mx.array, cache: list[KVCache]) -> mx.array: + """Run one position through the stack, extending ``cache`` in place.""" + hidden = self.inputs_embeds_projector(embeds) + mask = create_attention_mask(hidden, cache[0]) + for layer, layer_cache in zip(self.layers, cache): + hidden = layer(hidden, mask, layer_cache) + return self.norm(hidden) + def __call__( self, token_ids: mx.array, backbone_hidden_state: mx.array ) -> mx.array: @@ -588,6 +615,33 @@ def next_logits( hidden = self.model(token_ids, backbone_hidden_state)[:, -1, :] return hidden @ self.codebooks_head.weight[head_idx] + def start_frame( + self, backbone_hidden_state: mx.array, cache: list[KVCache] + ) -> None: + """Seed a frame's cache with position zero: the backbone hidden state.""" + model = self.model + if model.backbone_hidden_state_projector is not None: + backbone_hidden_state = model.backbone_hidden_state_projector( + backbone_hidden_state + ) + model.step(backbone_hidden_state[:, None, :], cache) + + def step_logits( + self, cache: list[KVCache], *, head_idx: int, token_id: Union[int, mx.array] + ) -> mx.array: + """Logits for ``head_idx`` after feeding that step's input token. + + :meth:`next_logits` reads the hidden state of the last prefix position, + where the prefix at step ``head_idx`` is + ``[backbone_state, first, s1 .. s_{head_idx - 1}]``. Feeding the token + that occupies that final position - ``first`` for step zero, otherwise + the previous step's sample - and reading the position just written + reproduces it exactly, without re-running the prefix. + """ + embeds = self.model.embed_codebook_token(token_id, head_idx) + hidden = self.model.step(embeds, cache)[:, -1, :] + return hidden @ self.codebooks_head.weight[head_idx] + class _CodebooksHead(nn.Module): """Position-specific output projections for depth-decoded codebooks.""" @@ -857,6 +911,26 @@ def _sample( top_k: int, allow_eos: bool = False, ) -> int: + """Sample one token id and read it back to the host.""" + token = self._sample_array( + logits, + temperature=temperature, + top_p=top_p, + top_k=top_k, + allow_eos=allow_eos, + ) + return int(token.item()) + + def _sample_array( + self, + logits: mx.array, + *, + temperature: float, + top_p: float, + top_k: int, + allow_eos: bool = False, + ) -> mx.array: + """Sample token ids from ``[batch, vocab]`` logits, leaving them on device.""" if temperature < 0: raise ValueError("temperature must be non-negative") if not 0 <= top_p <= 1: @@ -880,8 +954,7 @@ def _sample( if effective_top_k == valid: effective_top_k = 0 sampler = make_sampler(temp=temperature, top_p=top_p, top_k=effective_top_k) - token = sampler(nn.log_softmax(logits, axis=-1)) - return int(token.item()) + return sampler(nn.log_softmax(logits, axis=-1)) def _mask_reserved_codec_logits(self, logits: mx.array) -> mx.array: """Mask ids outside the codec codebook in a logits tensor. @@ -929,25 +1002,51 @@ def _depth_tokens( top_p: float, top_k: int, ) -> list[int]: - tokens = [0, first_codebook] - for _ in range(self.num_codebooks - 1): - token_ids = mx.array(tokens, dtype=mx.int32)[None, :] - logits = self.depth_decoder.next_logits(token_ids, conditional_hidden) - if unconditional_hidden is not None: - unconditional_logits = self.depth_decoder.next_logits( - token_ids, unconditional_hidden + """Sample the remaining codebooks of one frame. + + The depth stack is walked one position at a time against a per-frame KV + cache. Re-running it over the growing prefix at every step instead costs + ~91% of generation wall time (118.8 ms of 130 ms per frame, against + 0.4 ms for the already-cached backbone), so the cached walk is 2.7x + faster end to end at unchanged logits. + + Each sampled codebook stays on the device and feeds the next step + directly; the frame is read back to the host once, at the end. Reading + every token with ``.item()`` would stall the GPU 15 times per frame. + The RNG is drawn in the same order either way, so seeded output is + unchanged. + """ + depth = self.depth_decoder + cond_cache = depth.model.make_cache() + depth.start_frame(conditional_hidden, cond_cache) + uncond_cache = None + if unconditional_hidden is not None: + uncond_cache = depth.model.make_cache() + depth.start_frame(unconditional_hidden, uncond_cache) + token: Union[int, mx.array] = first_codebook + sampled: list[mx.array] = [] + for head_idx in range(self.num_codebooks - 1): + logits = depth.step_logits(cond_cache, head_idx=head_idx, token_id=token) + if uncond_cache is not None: + unconditional_logits = depth.step_logits( + uncond_cache, head_idx=head_idx, token_id=token ) logits = unconditional_logits + cfg_scale * ( logits - unconditional_logits ) # Apply the reserved-id mask before handing logits to the sampler. - # Keeping it here (as well as in ``_sample``) means custom samplers - # and deterministic test doubles observe the same official flow. + # Keeping it here (as well as in ``_sample_array``) means custom + # samplers and deterministic test doubles observe the same flow. logits = self._mask_reserved_codec_logits(logits) - tokens.append( - self._sample(logits, temperature=temperature, top_p=top_p, top_k=top_k) - ) - return tokens[1:] + token = self._sample_array( + logits, temperature=temperature, top_p=top_p, top_k=top_k + ).reshape(1) + # Start this step on the GPU while the next one is being built. + mx.async_eval(token) + sampled.append(token) + if not sampled: + return [first_codebook] + return [first_codebook, *mx.concatenate(sampled).tolist()] @staticmethod def _audio_vector(audio: Any) -> mx.array: diff --git a/mlx_audio/tts/tests/test_breeze_depth_cache.py b/mlx_audio/tts/tests/test_breeze_depth_cache.py new file mode 100644 index 000000000..68afc661b --- /dev/null +++ b/mlx_audio/tts/tests/test_breeze_depth_cache.py @@ -0,0 +1,282 @@ +"""Depth-cache regressions for Breeze TTS 2. + +The depth decoder used to re-run its full stack over the growing prefix at every +one of the 15 codebook steps of a frame. It now walks one position per step +against a per-frame KV cache, which is ~2.7x faster end to end; these tests pin +the arithmetic that made the change safe. The reference walk below keeps the +original prefix-recompute form, so both live in the same test file. +""" + +import pytest + +try: + import mlx.core as mx +except (ImportError, RuntimeError) as exc: # pragma: no cover - CI without Metal + pytest.skip(f"MLX device unavailable: {exc}", allow_module_level=True) + +from mlx_audio.tts.models.breeze_tts.breeze_tts import Model +from mlx_audio.tts.models.breeze_tts.config import ModelConfig + + +def _tiny_config(**overrides): + values = dict( + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=2, + num_key_value_heads=1, + head_dim=8, + num_codebooks=4, + vocab_size=8, + text_vocab_size=32, + text_encoder_config={ + "hidden_size": 12, + "num_hidden_layers": 1, + "intermediate_size": 24, + "num_attention_heads": 2, + "num_key_value_heads": 1, + "head_dim": 6, + "rms_norm_eps": 1e-6, + "vocab_size": 32, + "layer_types": ["full_attention"], + }, + depth_decoder_config={ + "hidden_size": 12, + "num_hidden_layers": 1, + "intermediate_size": 24, + "num_attention_heads": 2, + "num_key_value_heads": 1, + "head_dim": 6, + "rms_norm_eps": 1e-5, + "num_codebooks": 4, + "vocab_size": 8, + "audio_embed_size": 16, + }, + ) + values.update(overrides) + return ModelConfig(**values) + + +def _randomized_model(seed: int = 0) -> Model: + """A tiny model with a randomised output head. + + ``_CodebooksHead`` is initialised to zeros, so a fixture that leaves it alone + produces all-zero logits: every comparison would pass and every greedy pick + would return the same id, which is a test that cannot fail. Randomising the + head is what gives these tests something to disagree about. + """ + model = Model(_tiny_config()) + mx.random.seed(seed) + head = model.depth_decoder.codebooks_head + head.weight = mx.random.normal(head.weight.shape) * 0.5 + return model + + +def _next_token(model: Model, logits: mx.array) -> int: + """Greedy pick through the same reserved-id mask the real loop applies.""" + return int(mx.argmax(model._mask_reserved_codec_logits(logits))) + + +def _reference_walk(model: Model, hidden: mx.array, first_codebook: int): + """The original arithmetic: one full-prefix forward per codebook step.""" + decoder = model.depth_decoder + logits_list = [] + tokens = [0, first_codebook] + for _ in range(model.num_codebooks - 1): + token_ids = mx.array(tokens, dtype=mx.int32)[None, :] + step = decoder.next_logits(token_ids, hidden) + mx.eval(step) + logits_list.append(step) + tokens.append(_next_token(model, step)) + return logits_list, tokens[1:] + + +def _cached_walk(model: Model, hidden: mx.array, first_codebook: int): + """The cached walk: one cache position per codebook step.""" + decoder = model.depth_decoder + cache = decoder.model.make_cache() + decoder.start_frame(hidden, cache) + logits_list = [] + tokens = [0, first_codebook] + for head_idx in range(model.num_codebooks - 1): + step = decoder.step_logits(cache, head_idx=head_idx, token_id=tokens[-1]) + mx.eval(step) + logits_list.append(step) + tokens.append(_next_token(model, step)) + return logits_list, tokens[1:] + + +def _synced_walk( + model, hidden, first_codebook, *, unconditional, cfg_scale, **sampling +): + """The per-token loop: each sampled codebook is read back to the host.""" + decoder = model.depth_decoder + cond_cache = decoder.model.make_cache() + decoder.start_frame(hidden, cond_cache) + uncond_cache = None + if unconditional is not None: + uncond_cache = decoder.model.make_cache() + decoder.start_frame(unconditional, uncond_cache) + tokens = [0, first_codebook] + for head_idx in range(model.num_codebooks - 1): + logits = decoder.step_logits(cond_cache, head_idx=head_idx, token_id=tokens[-1]) + if uncond_cache is not None: + uncond = decoder.step_logits( + uncond_cache, head_idx=head_idx, token_id=tokens[-1] + ) + logits = uncond + cfg_scale * (logits - uncond) + logits = model._mask_reserved_codec_logits(logits) + tokens.append(model._sample(logits, **sampling)) + return tokens[1:] + + +def test_cached_walk_reproduces_prefix_recompute_logits(): + model = _randomized_model() + hidden = mx.random.normal((1, 16)) + # Compare on the CPU. On the GPU the single-query cached step and the + # full-prefix pass go through different attention kernels, whose results + # differ by ~3e-3 on an M5 Max (mlx 0.32.2) even though the arithmetic is + # the same; the CPU agrees to ~1e-6. + with mx.stream(mx.cpu): + reference, _ = _reference_walk(model, hidden, 1) + cached, _ = _cached_walk(model, hidden, 1) + + assert len(cached) == len(reference) == model.num_codebooks - 1 + # Guard against a vacuous comparison: a zero head would make every logits + # tensor agree with every other one. + assert max(float(mx.abs(step).max()) for step in reference) > 0.1 + for expected, actual in zip(reference, cached): + # The walks share the same arithmetic and differ only in the order the + # accumulation happens in, so the tolerated gap is fp32 rounding, not a + # design allowance (fp64 agreement is exact). + assert mx.allclose(expected, actual, atol=1e-5).item() + assert int(mx.argmax(expected)) == int(mx.argmax(actual)) + + +def test_cached_walk_picks_the_same_codebooks_as_recompute(): + model = _randomized_model() + for seed in range(4): + mx.random.seed(seed) + hidden = mx.random.normal((1, 16)) + for first_codebook in (1, 2, 3): + assert ( + _reference_walk(model, hidden, first_codebook)[1] + == _cached_walk(model, hidden, first_codebook)[1] + ) + + +def test_depth_tokens_uses_the_cached_walk(): + model = _randomized_model() + hidden = mx.random.normal((1, 16)) + expected = _reference_walk(model, hidden, 1)[1] + tokens = model._depth_tokens( + 1, + hidden, + unconditional_hidden=None, + cfg_scale=1.0, + temperature=0.0, + top_p=1.0, + top_k=1, + ) + assert tokens == expected + + +def test_depth_tokens_cfg_matches_the_recompute_path(): + model = _randomized_model() + hidden = mx.random.normal((1, 16)) + unconditional = mx.random.normal((1, 16)) + tokens = model._depth_tokens( + 1, + hidden, + unconditional_hidden=unconditional, + cfg_scale=2.0, + temperature=0.0, + top_p=1.0, + top_k=1, + ) + decoder = model.depth_decoder + cond_cache = decoder.model.make_cache() + uncond_cache = decoder.model.make_cache() + decoder.start_frame(hidden, cond_cache) + decoder.start_frame(unconditional, uncond_cache) + expected = [0, 1] + for head_idx in range(model.num_codebooks - 1): + cond = decoder.step_logits(cond_cache, head_idx=head_idx, token_id=expected[-1]) + uncond = decoder.step_logits( + uncond_cache, head_idx=head_idx, token_id=expected[-1] + ) + logits = uncond + 2.0 * (cond - uncond) + expected.append(_next_token(model, logits)) + assert tokens == expected[1:] + + +@pytest.mark.parametrize("cfg_scale", [None, 2.0]) +def test_depth_tokens_stay_on_device_and_keep_the_rng_order(monkeypatch, cfg_scale): + """Sampled codebooks feed the next step without a host round trip. + + Reading every token back with ``.item()`` stalled the GPU 15 times per + frame. The loop must still draw from the RNG in the same order, so a seeded + stochastic run reproduces the per-token loop exactly. + """ + model = _randomized_model() + hidden = mx.random.normal((1, 16)) + unconditional = mx.random.normal((1, 16)) if cfg_scale else None + kwargs = dict( + unconditional=unconditional, + cfg_scale=cfg_scale or 1.0, + temperature=1.0, + top_p=1.0, + top_k=0, + ) + expected = [] + for seed in range(4): + mx.random.seed(seed) + expected.append(_synced_walk(model, hidden, 1, **kwargs)) + # Guard against a vacuous comparison: the seeds must pick different frames. + assert len({tuple(tokens) for tokens in expected}) > 1 + + def synced_sample(*_args, **_kwargs): + raise AssertionError("the depth loop read a token back to the host") + + monkeypatch.setattr(model, "_sample", synced_sample) + kwargs["unconditional_hidden"] = kwargs.pop("unconditional") + for seed, tokens in enumerate(expected): + mx.random.seed(seed) + assert model._depth_tokens(1, hidden, **kwargs) == tokens + + +def test_each_step_extends_the_frame_cache_once(): + model = _randomized_model() + hidden = mx.random.normal((1, 16)) + decoder = model.depth_decoder + cache = decoder.model.make_cache() + + assert [int(layer.offset) for layer in cache] == [0] * len(cache) + decoder.start_frame(hidden, cache) + assert {int(layer.offset) for layer in cache} == {1} + for head_idx in range(model.num_codebooks - 1): + decoder.step_logits(cache, head_idx=head_idx, token_id=1) + assert {int(layer.offset) for layer in cache} == {head_idx + 2} + + +def test_frames_do_not_share_cache_state(): + model = _randomized_model() + mx.random.seed(11) + first_hidden = mx.random.normal((1, 16)) + other_hidden = mx.random.normal((1, 16)) + + def frame(hidden): + return model._depth_tokens( + 1, + hidden, + unconditional_hidden=None, + cfg_scale=1.0, + temperature=0.0, + top_p=1.0, + top_k=1, + ) + + expected = frame(first_hidden) + frame(other_hidden) # interleave an unrelated frame + assert frame(first_hidden) == expected + assert frame(other_hidden) == frame(other_hidden) diff --git a/mlx_audio/tts/tests/test_breeze_tts.py b/mlx_audio/tts/tests/test_breeze_tts.py index 319d47e60..6e12d9455 100644 --- a/mlx_audio/tts/tests/test_breeze_tts.py +++ b/mlx_audio/tts/tests/test_breeze_tts.py @@ -153,16 +153,21 @@ def test_default_generic_voice_maps_to_breeze_s0(): def test_depth_cfg_applies_to_every_remaining_codebook(monkeypatch): model = Model(tiny_config()) sampled_logits = [] + frame_values = {} - def next_logits(_token_ids, hidden): - return mx.full((1, 8), hidden[0, 0]) + def start_frame(hidden, cache): + frame_values[id(cache)] = hidden[0, 0] + + def step_logits(cache, *, head_idx, token_id): + return mx.full((1, 8), frame_values[id(cache)]) def sample(logits, **_kwargs): sampled_logits.append(logits) - return 1 + return mx.array([1]) - monkeypatch.setattr(model.depth_decoder, "next_logits", next_logits) - monkeypatch.setattr(model, "_sample", sample) + monkeypatch.setattr(model.depth_decoder, "start_frame", start_frame) + monkeypatch.setattr(model.depth_decoder, "step_logits", step_logits) + monkeypatch.setattr(model, "_sample_array", sample) tokens = model._depth_tokens( 1, mx.ones((1, 16)), @@ -182,15 +187,19 @@ def test_depth_cfg_masks_reserved_tokens_at_every_step(monkeypatch): model = Model(tiny_config()) sampled_logits = [] - def next_logits(_token_ids, _hidden): + def start_frame(_hidden, _cache): + return None + + def step_logits(_cache, *, head_idx, token_id): return mx.arange(8, dtype=mx.float32)[None, :] def sample(logits, **_kwargs): sampled_logits.append(logits) - return 1 + return mx.array([1]) - monkeypatch.setattr(model.depth_decoder, "next_logits", next_logits) - monkeypatch.setattr(model, "_sample", sample) + monkeypatch.setattr(model.depth_decoder, "start_frame", start_frame) + monkeypatch.setattr(model.depth_decoder, "step_logits", step_logits) + monkeypatch.setattr(model, "_sample_array", sample) model._depth_tokens( 1, mx.ones((1, 16)),