From f56d99712cf35cee9d44f1ec11d00a22a692723c Mon Sep 17 00:00:00 2001 From: Tarjei Mandt Date: Tue, 7 Apr 2026 09:00:38 +1000 Subject: [PATCH] Fix output corruption in speculative decoding (#1109) --- mlx_lm/generate.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/mlx_lm/generate.py b/mlx_lm/generate.py index 4f191f1..5ea9e8b 100644 --- a/mlx_lm/generate.py +++ b/mlx_lm/generate.py @@ -525,6 +525,12 @@ def speculative_generate_step( model_cache = prompt_cache[: len(model.layers)] draft_cache = prompt_cache[len(model.layers) :] + if not cache.can_trim_prompt_cache(model_cache): + types = {type(c).__name__ for c in model_cache if not c.is_trimmable()} + raise ValueError( + f"Speculative decoding requires a trimmable prompt cache " f"(got {types})." + ) + sampler = sampler or (lambda x: mx.argmax(x, axis=-1)) quantize_cache_fn = functools.partial( @@ -570,11 +576,12 @@ def speculative_generate_step( return _process_and_sample(None, logits.squeeze(0)) def _prefill(model, cache, y): - while y.size > prefill_step_size: - model(y[:prefill_step_size][None], cache=cache) + while y.size > 1: + n_to_process = min(prefill_step_size, y.size - 1) + model(y[:n_to_process][None], cache=cache) quantize_cache_fn(cache) mx.eval([c.state for c in cache]) - y = y[prefill_step_size:] + y = y[n_to_process:] mx.clear_cache() return y