only apply lm_head to the last token (#406)

* only apply lm_head to the last token

* peel off last token instead and use lazy eval

* fix
This commit is contained in:
Awni Hannun
2025-08-28 12:31:12 -07:00
committed by GitHub
parent da1309f5a7
commit 24fefe3d05
2 changed files with 10 additions and 8 deletions
+8 -6
View File
@@ -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
)
+2 -2
View File
@@ -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
)