Skip to content

Commit 9fe5f43

Browse files
authored
custom dsv32 chat template (#693)
* custom dsv32 chat template * use has_chat_template
1 parent 1b2d11b commit 9fe5f43

5 files changed

Lines changed: 368 additions & 35 deletions

File tree

mlx_lm/cache_prompt.py

Lines changed: 1 addition & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -41,16 +41,6 @@ def setup_arg_parser():
4141
default=None,
4242
help="End of sequence token for tokenizer",
4343
)
44-
parser.add_argument(
45-
"--ignore-chat-template",
46-
action="store_true",
47-
help="Use the raw prompt without the tokenizer's chat template.",
48-
)
49-
parser.add_argument(
50-
"--use-default-chat-template",
51-
action="store_true",
52-
help="Use the default chat template",
53-
)
5444
parser.add_argument(
5545
"--max-kv-size",
5646
type=int,
@@ -107,11 +97,7 @@ def main():
10797

10898
args.prompt = sys.stdin.read() if args.prompt == "-" else args.prompt
10999

110-
if args.use_default_chat_template:
111-
if tokenizer.chat_template is None:
112-
tokenizer.chat_template = tokenizer.default_chat_template
113-
114-
if not args.ignore_chat_template and tokenizer.chat_template is not None:
100+
if tokenizer.has_chat_template:
115101
messages = [{"role": "user", "content": args.prompt}]
116102
prompt = tokenizer.apply_chat_template(
117103
messages,
@@ -155,7 +141,6 @@ def callback(processed, total_tokens):
155141
print("Saving...")
156142
metadata = {}
157143
metadata["model"] = args.model
158-
metadata["chat_template"] = json.dumps(tokenizer.chat_template)
159144
metadata["tokenizer_config"] = json.dumps(tokenizer_config)
160145
save_prompt_cache(args.prompt_cache_file, cache, metadata)
161146

mlx_lm/generate.py

Lines changed: 1 addition & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1302,15 +1302,9 @@ def main():
13021302
if args.chat_template_config is not None:
13031303
template_kwargs = json.loads(args.chat_template_config)
13041304

1305-
if args.use_default_chat_template:
1306-
if tokenizer.chat_template is None:
1307-
tokenizer.chat_template = tokenizer.default_chat_template
1308-
elif using_cache:
1309-
tokenizer.chat_template = json.loads(metadata["chat_template"])
1310-
13111305
prompt = args.prompt.replace("\\n", "\n").replace("\\t", "\t")
13121306
prompt = sys.stdin.read() if prompt == "-" else prompt
1313-
if not args.ignore_chat_template and tokenizer.chat_template is not None:
1307+
if not args.ignore_chat_template and tokenizer.has_chat_template:
13141308
if args.system_prompt is not None:
13151309
messages = [{"role": "system", "content": args.system_prompt}]
13161310
else:

mlx_lm/tokenizer_utils.py

Lines changed: 34 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import importlib
12
import json
23
from functools import partial
34
from json import JSONDecodeError
@@ -248,7 +249,11 @@ class TokenizerWrapper:
248249
"""
249250

250251
def __init__(
251-
self, tokenizer, detokenizer_class=NaiveStreamingDetokenizer, eos_token_ids=None
252+
self,
253+
tokenizer,
254+
detokenizer_class=NaiveStreamingDetokenizer,
255+
eos_token_ids=None,
256+
chat_template=None,
252257
):
253258
self._tokenizer = tokenizer
254259
self._detokenizer_class = detokenizer_class
@@ -261,6 +266,10 @@ def __init__(
261266
self._think_end = None
262267
self._tool_call_start = None
263268
self._tool_call_end = None
269+
self._chat_template = chat_template
270+
self.has_chat_template = (
271+
tokenizer.chat_template is not None or chat_template is not None
272+
)
264273

265274
THINK_TOKENS = [("<think>", "</think>")]
266275
TOOL_CALL_TOKENS = [("<tool_call>", "</tool_call>")]
@@ -278,9 +287,15 @@ def __init__(
278287
self._tool_call_end = tool_call_end
279288
break
280289

281-
def apply_chat_template(self, *args, **kwargs):
290+
def apply_chat_template(self, *args, tokenize=True, **kwargs):
291+
if self._chat_template is not None:
292+
out = self._chat_template(*args, **kwargs)
293+
if tokenize:
294+
out = self._tokenizer.encode(out, add_special_tokens=False)
295+
return out
296+
282297
kwargs["return_dict"] = False
283-
return self._tokenizer.apply_chat_template(*args, **kwargs)
298+
return self._tokenizer.apply_chat_template(*args, tokenize=tokenize, **kwargs)
284299

285300
def add_eos_token(self, token: str):
286301
token_id = None
@@ -450,12 +465,28 @@ def load(
450465
if isinstance(eos_token_ids, int):
451466
eos_token_ids = [eos_token_ids]
452467

468+
tokenizer_config_file = model_path / "tokenizer_config.json"
469+
custom_tokenizer = None
470+
if tokenizer_config_file.exists():
471+
with open(tokenizer_config_file, "r", encoding="utf-8") as fid:
472+
try:
473+
tokenizer_config = json.load(fid)
474+
except JSONDecodeError as e:
475+
raise JSONDecodeError(
476+
"Failed to parse tokenizer_config.json", e.doc, e.pos
477+
)
478+
if tokenizer_type := tokenizer_config.get("tokenizer_type", False):
479+
custom_tokenizer = importlib.import_module(
480+
f"mlx_lm.tokenizers.{tokenizer_type}"
481+
)
482+
453483
if return_tokenizer:
454484
kwargs = tokenizer_config_extra or {}
455485
return TokenizerWrapper(
456486
AutoTokenizer.from_pretrained(model_path, **kwargs),
457487
detokenizer_class,
458488
eos_token_ids=eos_token_ids,
489+
chat_template=getattr(custom_tokenizer, "apply_chat_template", None),
459490
)
460491
else:
461492
return detokenizer_class

0 commit comments

Comments
 (0)