-
Notifications
You must be signed in to change notification settings - Fork 1.1k
Expand file tree
/
Copy pathbatch_generate_response.py
More file actions
51 lines (43 loc) · 1.17 KB
/
Copy pathbatch_generate_response.py
File metadata and controls
51 lines (43 loc) · 1.17 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
# Copyright © 2025 Apple Inc.
from mlx_lm import batch_generate, load
# Specify the checkpoint
checkpoint = "mlx-community/Llama-3.2-3B-Instruct-4bit"
# Load the corresponding model and tokenizer
model, tokenizer = load(path_or_hf_repo=checkpoint)
# A batch of prompts
prompts = [
"Write a story about Einstein.",
"Why is the sky blue?",
"What time is it?",
"How tall is Mt Everest?",
]
# Apply the chat template and encode to tokens
prompts = [
tokenizer.apply_chat_template(
[{"role": "user", "content": p}],
add_generation_prompt=True,
)
for p in prompts
]
# Set `verbose=True` to see generation statistics
result = batch_generate(
model, tokenizer, prompts, verbose=False, return_prompt_caches=True, max_tokens=2048
)
print(result.texts[-1])
prompts = [
"Could you summarize that?",
"And what about the sea?",
"Try again?",
"And Mt Olympus?",
]
prompts = [
tokenizer.apply_chat_template(
[{"role": "user", "content": p}],
add_generation_prompt=True,
)
for p in prompts
]
result = batch_generate(
model, tokenizer, prompts, verbose=False, prompt_caches=result.caches
)
print(result.texts[-1])