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:
+8
-6
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user