* Allow prompt and input_embeddings * Prefer Optional to "| None" * Require prompt for generate_step * Formatting
157 lines
4.9 KiB
Python
157 lines
4.9 KiB
Python
# Copyright © 2024 Apple Inc.
|
|
|
|
import unittest
|
|
from typing import List
|
|
|
|
from mlx_lm.generate import (
|
|
GenerationResponse,
|
|
generate,
|
|
stream_generate,
|
|
)
|
|
from mlx_lm.sample_utils import make_logits_processors, make_sampler
|
|
from mlx_lm.utils import load
|
|
|
|
|
|
class TestGenerate(unittest.TestCase):
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.HF_MODEL_PATH = "mlx-community/Qwen1.5-0.5B-Chat-4bit"
|
|
cls.model, cls.tokenizer = load(cls.HF_MODEL_PATH)
|
|
|
|
def test_generate(self):
|
|
# Simple test that generation runs
|
|
text = generate(
|
|
self.model, self.tokenizer, "hello", max_tokens=5, verbose=False
|
|
)
|
|
|
|
def test_generate_with_logit_bias(self):
|
|
logit_bias = {0: 2000.0, 1: -20.0}
|
|
text = generate(
|
|
self.model,
|
|
self.tokenizer,
|
|
"hello",
|
|
max_tokens=5,
|
|
logits_processors=make_logits_processors(logit_bias),
|
|
verbose=False,
|
|
)
|
|
self.assertEqual(text, "!!!!!")
|
|
|
|
def test_generate_with_processor(self):
|
|
init_toks = self.tokenizer.encode("hello")
|
|
|
|
all_toks = None
|
|
|
|
def logits_processor(toks, logits):
|
|
nonlocal all_toks
|
|
all_toks = toks
|
|
return logits
|
|
|
|
generate(
|
|
self.model,
|
|
self.tokenizer,
|
|
"hello",
|
|
max_tokens=5,
|
|
verbose=False,
|
|
logits_processors=[logits_processor],
|
|
)
|
|
self.assertEqual(len(all_toks), len(init_toks) + 5)
|
|
|
|
def test_stream_generate_speculative(self):
|
|
# Use same model as draft model, this is not a speed test
|
|
draft_model, _ = load(self.HF_MODEL_PATH)
|
|
|
|
results: List[GenerationResponse] = []
|
|
drafted: List[bool] = []
|
|
|
|
# make a determinate sampler
|
|
sampler = make_sampler(temp=0.0)
|
|
messages = [{"role": "user", "content": "hello"}]
|
|
prompt = self.tokenizer.apply_chat_template(
|
|
messages, add_generation_prompt=True
|
|
)
|
|
|
|
for generation_result in stream_generate(
|
|
model=self.model,
|
|
tokenizer=self.tokenizer,
|
|
prompt=prompt,
|
|
max_tokens=5,
|
|
draft_model=draft_model,
|
|
num_draft_tokens=2,
|
|
sampler=sampler,
|
|
):
|
|
drafted.append(generation_result.from_draft)
|
|
results.append(generation_result)
|
|
|
|
self.assertEqual(len(results), 6)
|
|
drafted.pop()
|
|
# since num_draft_tokens is 2 and draft model is the same, the
|
|
# first 2 generations should be drafts, the third should come
|
|
# from the target model, and last two should be drafts
|
|
self.assertEqual(drafted, [True, True, False, True, True])
|
|
|
|
def test_stream_generate_input_embeddings(self):
|
|
sampler = make_sampler(temp=0.0) # determinate sampler
|
|
|
|
# get prompt embeddings
|
|
messages = [{"role": "user", "content": "Say 'TEST' and nothing else"}]
|
|
prompt = self.tokenizer.apply_chat_template(
|
|
messages, add_generation_prompt=True
|
|
)
|
|
prompt_embeddings = self.model.model.embed_tokens(prompt)
|
|
|
|
response = ""
|
|
for generation_result in stream_generate(
|
|
model=self.model,
|
|
tokenizer=self.tokenizer,
|
|
prompt=prompt,
|
|
max_tokens=5,
|
|
sampler=sampler,
|
|
input_embeddings=prompt_embeddings,
|
|
):
|
|
response += generation_result.text
|
|
|
|
self.assertEqual("TEST", response)
|
|
|
|
def test_stream_generate_input_embeddings_prefill(self):
|
|
sampler = make_sampler(temp=0.0) # determinate sampler
|
|
|
|
# get prompt embeddings
|
|
messages = [{"role": "user", "content": "Say 'TEST' and nothing else"}]
|
|
prompt = self.tokenizer.apply_chat_template(
|
|
messages, add_generation_prompt=True
|
|
)
|
|
prompt_embeddings = self.model.model.embed_tokens(prompt)
|
|
|
|
# setup prompt progress callback to track batched prefill
|
|
num_prompt_processing_callbacks = 0
|
|
|
|
def progress_callback(processed: int, total: int) -> None:
|
|
nonlocal num_prompt_processing_callbacks
|
|
num_prompt_processing_callbacks += 1
|
|
|
|
# generate
|
|
prefill_step_size = 5
|
|
response = ""
|
|
for generation_result in stream_generate(
|
|
model=self.model,
|
|
tokenizer=self.tokenizer,
|
|
prompt=prompt,
|
|
max_tokens=5,
|
|
sampler=sampler,
|
|
input_embeddings=prompt_embeddings,
|
|
prefill_step_size=prefill_step_size,
|
|
prompt_progress_callback=progress_callback,
|
|
):
|
|
response += generation_result.text
|
|
|
|
self.assertEqual("TEST", response)
|
|
num_embeddings = prompt_embeddings.shape[0]
|
|
self.assertEqual(
|
|
num_embeddings / prefill_step_size, num_prompt_processing_callbacks
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|