diff --git a/mlx_lm/generate.py b/mlx_lm/generate.py index d71de05..376f468 100644 --- a/mlx_lm/generate.py +++ b/mlx_lm/generate.py @@ -415,22 +415,24 @@ def generate_step( len(input_embeddings) if input_embeddings is not None else len(prompt) ) prompt_processed_tokens = 0 - while total_prompt_tokens - prompt_processed_tokens > prefill_step_size: + prompt_progress_callback(prompt_processed_tokens, total_prompt_tokens) + while total_prompt_tokens - prompt_processed_tokens > 1: + n_to_process = min(prefill_step_size, prompt.size - 1) _model_call( - input_tokens=prompt[:prefill_step_size][None], + input_tokens=prompt[:n_to_process][None], input_embeddings=( - input_embeddings[:prefill_step_size][None] + input_embeddings[:n_to_process][None] if input_embeddings is not None else None ), ) quantize_cache_fn(prompt_cache) mx.eval([c.state for c in prompt_cache]) + prompt_processed_tokens += n_to_process prompt_progress_callback(prompt_processed_tokens, total_prompt_tokens) - prompt_processed_tokens += prefill_step_size - prompt = prompt[prefill_step_size:] + prompt = prompt[n_to_process:] input_embeddings = ( - input_embeddings[prefill_step_size:] + input_embeddings[n_to_process:] if input_embeddings is not None else input_embeddings ) diff --git a/tests/test_generate.py b/tests/test_generate.py index 8fbd45f..b80a756 100644 --- a/tests/test_generate.py +++ b/tests/test_generate.py @@ -147,8 +147,8 @@ class TestGenerate(unittest.TestCase): self.assertEqual("TEST", response) num_embeddings = prompt_embeddings.shape[0] - self.assertEqual( - num_embeddings / prefill_step_size, num_prompt_processing_callbacks + self.assertTrue( + num_embeddings / prefill_step_size < num_prompt_processing_callbacks )