Skip to content

Commit 657a66c

Browse files
authored
revert return dict and wrap apply_chat_template (#691)
1 parent 595fb4b commit 657a66c

13 files changed

Lines changed: 14 additions & 26 deletions

README.md

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -71,7 +71,7 @@ prompt = "Write a story about Einstein"
7171

7272
messages = [{"role": "user", "content": prompt}]
7373
prompt = tokenizer.apply_chat_template(
74-
messages, add_generation_prompt=True, return_dict=False,
74+
messages, add_generation_prompt=True,
7575
)
7676

7777
text = generate(model, tokenizer, prompt=prompt, verbose=True)
@@ -130,7 +130,7 @@ prompt = "Write a story about Einstein"
130130

131131
messages = [{"role": "user", "content": prompt}]
132132
prompt = tokenizer.apply_chat_template(
133-
messages, add_generation_prompt=True, return_dict=False,
133+
messages, add_generation_prompt=True,
134134
)
135135

136136
for response in stream_generate(model, tokenizer, prompt, max_tokens=512):

mlx_lm/cache_prompt.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -117,7 +117,6 @@ def main():
117117
messages,
118118
add_generation_prompt=False,
119119
continue_final_message=True,
120-
return_dict=False,
121120
)
122121

123122
else:

mlx_lm/chat.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -140,7 +140,8 @@ def print_help():
140140
messages.append({"role": "system", "content": args.system_prompt})
141141
messages.append({"role": "user", "content": query})
142142
prompt = tokenizer.apply_chat_template(
143-
messages, add_generation_prompt=True, return_dict=False
143+
messages,
144+
add_generation_prompt=True,
144145
)
145146
for response in stream_generate(
146147
model,

mlx_lm/evaluate.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -63,7 +63,6 @@ def apply_chat_template(self, chat_history, add_generation_prompt=True) -> str:
6363
tokenize=False,
6464
add_generation_prompt=add_generation_prompt,
6565
continue_final_message=not add_generation_prompt,
66-
return_dict=False,
6766
**extra_kwargs,
6867
)
6968

mlx_lm/examples/batch_generate_response.py

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,6 @@
2121
tokenizer.apply_chat_template(
2222
[{"role": "user", "content": p}],
2323
add_generation_prompt=True,
24-
return_dict=False,
2524
)
2625
for p in prompts
2726
]
@@ -42,7 +41,6 @@
4241
tokenizer.apply_chat_template(
4342
[{"role": "user", "content": p}],
4443
add_generation_prompt=True,
45-
return_dict=False,
4644
)
4745
for p in prompts
4846
]

mlx_lm/examples/chat.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,8 @@
1616
prompt = "Hi my name is <Name>."
1717
messages = [{"role": "user", "content": prompt}]
1818
prompt = tokenizer.apply_chat_template(
19-
messages, add_generation_prompt=True, return_dict=False
19+
messages,
20+
add_generation_prompt=True,
2021
)
2122

2223
# Assistant response
@@ -32,7 +33,8 @@
3233
prompt = "What's my name?"
3334
messages = [{"role": "user", "content": prompt}]
3435
prompt = tokenizer.apply_chat_template(
35-
messages, add_generation_prompt=True, return_dict=False
36+
messages,
37+
add_generation_prompt=True,
3638
)
3739

3840
# Assistant response

mlx_lm/examples/generate_response.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,6 @@
1616
prompt = tokenizer.apply_chat_template(
1717
conversation=conversation,
1818
add_generation_prompt=True,
19-
return_dict=False,
2019
)
2120

2221
# Specify the maximum number of tokens

mlx_lm/examples/sharded_generate.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -64,7 +64,8 @@ def rprint(*args, **kwargs):
6464

6565
messages = [{"role": "user", "content": args.prompt}]
6666
prompt = tokenizer.apply_chat_template(
67-
messages, add_generation_prompt=True, return_dict=False
67+
messages,
68+
add_generation_prompt=True,
6869
)
6970

7071
for response in stream_generate(

mlx_lm/examples/tool_use.py

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -34,7 +34,6 @@ def multiply(a: float, b: float):
3434
messages,
3535
add_generation_prompt=True,
3636
tools=list(tools.values()),
37-
return_dict=False,
3837
)
3938

4039
prompt_cache = make_prompt_cache(model)
@@ -63,7 +62,6 @@ def multiply(a: float, b: float):
6362
prompt = tokenizer.apply_chat_template(
6463
messages,
6564
add_generation_prompt=True,
66-
return_dict=False,
6765
)
6866

6967
# Generate the final response:

mlx_lm/generate.py

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1325,7 +1325,6 @@ def main():
13251325
tokenize=False,
13261326
continue_final_message=has_prefill,
13271327
add_generation_prompt=not has_prefill,
1328-
return_dict=False,
13291328
**template_kwargs,
13301329
)
13311330

@@ -1337,7 +1336,6 @@ def main():
13371336
messages,
13381337
tokenize=False,
13391338
continue_final_message=has_prefill,
1340-
return_dict=False,
13411339
add_generation_prompt=not has_prefill,
13421340
)
13431341
prompt = prompt[test_prompt.index("<query>") :]

0 commit comments

Comments
 (0)