# Copyright © 2024 Apple Inc.
import abc
import copy
import functools
import importlib
import inspect
import json
from functools import partial
from json import JSONDecodeError
from typing import Any, Dict, List, Optional
from transformers import AutoTokenizer, PreTrainedTokenizerFast
from transformers.tokenization_utils_base import PreTrainedTokenizerBase
class StreamingDetokenizer(abc.ABC):
"""The streaming detokenizer interface so that we can detokenize one token at a time.
Example usage is as follows:
detokenizer = ...
# Reset the tokenizer state
detokenizer.reset()
for token in generate(...):
detokenizer.add_token(token.item())
# Contains the whole text so far. Some tokens may not be included
# since it contains whole words usually.
detokenizer.text
# Contains the printable segment (usually a word) since the last
# time it was accessed
detokenizer.last_segment
# Contains all the tokens added so far
detokenizer.tokens
# Make sure that we detokenize any remaining tokens
detokenizer.finalize()
# Now detokenizer.text should match tokenizer.decode(detokenizer.tokens)
"""
# Set by reset(); text is a property on some subclasses.
text: str
tokens: List[int]
offset: int
@abc.abstractmethod
def reset(self):
"""Drop all streaming state, keeping data derived from the tokenizer."""
@abc.abstractmethod
def add_token(self, token):
"""Consume one token id."""
@abc.abstractmethod
def finalize(self):
"""Flush any text held back waiting for more tokens."""
@property
def last_segment(self):
"""Return the last segment of readable text since last time this property was accessed."""
text = self.text
segment = text[self.offset :]
self.offset = len(text)
return segment
class NaiveStreamingDetokenizer(StreamingDetokenizer):
"""NaiveStreamingDetokenizer relies on the underlying tokenizer
implementation and should work with every tokenizer.
Its complexity is O(T^2) where T is the longest line since it will
repeatedly detokenize the same tokens until a new line is generated.
"""
def __init__(self, tokenizer):
super().__init__()
self._tokenizer = tokenizer
self._tokenizer.decode([0])
probe = tokenizer.encode("a ,b", add_special_tokens=False)
self._clean_spaces = " ," not in tokenizer.decode(probe)
self.reset()
def reset(self):
self.offset = 0
self.tokens = []
self._text = ""
self._current_tokens = []
self._current_text = ""
def add_token(self, token):
self._current_tokens.append(token)
self.tokens.append(token)
def finalize(self):
self._text += self._tokenizer.decode(self._current_tokens)
self._current_tokens = []
self._current_text = ""
@property
def text(self):
if self._current_tokens:
self._current_text = self._tokenizer.decode(self._current_tokens)
if self._current_text.endswith("\ufffd"):
# An incomplete character can decode to several replacements.
self._current_text = self._current_text.rstrip("\ufffd")
elif (
self._clean_spaces
and len(self._current_text) > 0
and self._current_text[-1] == " "
):
self._current_text = self._current_text[:-1]
if self._current_text and self._current_text[-1] == "\n":
self._text += self._current_text
self._current_tokens.clear()
self._current_text = ""
return self._text + self._current_text
class SPMStreamingDetokenizer(StreamingDetokenizer):
"""A streaming detokenizer for SPM models.
It adds tokens to the text if the next token starts with the special SPM
underscore which results in linear complexity.
"""
_sep = "\u2581".encode("utf-8")
def __init__(self, tokenizer, trim_space=True):
super().__init__()
self.trim_space = trim_space
ids = list(range(len(tokenizer)))
tokens = tokenizer.convert_ids_to_tokens(ids)
self.tokenmap = [
# Byte tokens carry their value in hex.
bytes([int(t[3:5], 16)]) if t.startswith("<0x") else t.encode("utf-8")
for t in tokens
]
self.reset()
def reset(self):
self.offset = 0
self._unflushed = b""
self.text = ""
self.tokens = []
def _try_flush(self, force=False):
text = self._unflushed.replace(self._sep, b" ").decode("utf-8", "replace")
if not force and text.endswith("\ufffd"):
return
if not self.text and self.trim_space and text and text[0] == " ":
text = text[1:]
self.text += text
self._unflushed = b""
def add_token(self, token):
self.tokens.append(token)
v = self.tokenmap[token]
self._unflushed += v
self._try_flush()
def finalize(self):
self._try_flush(force=True)
self._unflushed = b""
@functools.lru_cache(maxsize=1)
def _byte_decoder():
"""See https://github.com/openai/gpt-2/blob/master/src/encoder.py for the rationale."""
char_to_bytes = {}
limits = [
0,
ord("!"),
ord("~") + 1,
ord("¡"),
ord("¬") + 1,
ord("®"),
ord("ÿ") + 1,
]
n = 0
for i, (start, stop) in enumerate(zip(limits, limits[1:])):
if i % 2 == 0:
for b in range(start, stop):
char_to_bytes[chr(2**8 + n)] = b
n += 1
else:
for b in range(start, stop):
char_to_bytes[chr(b)] = b
return char_to_bytes
class BPEStreamingDetokenizer(StreamingDetokenizer):
"""A streaming detokenizer for OpenAI style BPE models.
It adds tokens to the text if the next token starts with a space similar to
the SPM detokenizer.
"""
def __init__(self, tokenizer):
super().__init__()
ids = list(range(len(tokenizer)))
self.tokenmap = tokenizer.convert_ids_to_tokens(ids)
self.reset()
def reset(self):
self.offset = 0
self._unflushed = ""
self.text = ""
self.tokens = []
def _decode_bytes(self, seq):
byte_decoder = _byte_decoder()
barr = bytearray()
for c in seq:
res = byte_decoder.get(c, False)
if res:
barr.append(res)
else:
barr.extend(bytes(c, "utf-8"))
return barr.decode("utf-8", "replace")
def _maybe_trim_space(self, current_text):
if len(current_text) == 0:
return current_text
elif current_text[0] != " ":
return current_text
elif not self.text:
return current_text[1:]
return current_text
def add_token(self, token):
self.tokens.append(token)
# Undocumented fallback from #418, likely for a padded model vocab.
# TODO(michalk8): check whether this is still needed.
v = self.tokenmap[token] if token < len(self.tokenmap) else "!"
self._unflushed += v
text = self._decode_bytes(self._unflushed)
# For multi-byte utf-8 wait until they are complete
# For single spaces wait until the next token to clean it if needed
if not text.endswith("\ufffd") and not (
len(v) == 1 and _byte_decoder().get(v[0]) == 32
):
self.text += self._maybe_trim_space(text)
self._unflushed = ""
def finalize(self):
byte_decoder = _byte_decoder()
current_text = bytearray(byte_decoder[c] for c in self._unflushed).decode(
"utf-8",
"replace",
)
self.text += self._maybe_trim_space(current_text)
self._unflushed = ""
def _infer_thinking(tokenizer):
vocab = tokenizer.get_vocab()
THINK_TOKENS = [
("", ""),
("", ""),
("<|think:start|>", "<|think:end|>"),
]
# Single token thinking modes
for think_start, think_end in THINK_TOKENS:
if think_start in vocab and think_end in vocab:
return (
think_start,
think_end,
(vocab[think_start],),
(vocab[think_end],),
)
# Multi token thinking modes
if "<|channel>" in vocab and "" in vocab:
think_start = "<|channel>thought"
think_end = ""
return (
think_start,
think_end,
tuple(tokenizer.encode(think_start, add_special_tokens=False)),
tuple(tokenizer.encode(think_end, add_special_tokens=False)),
)
if _is_xtml_vocab(vocab):
think_start = "<|open|>think<|sep|>"
think_end = "<|close|>think<|sep|>"
return (
think_start,
think_end,
tuple(tokenizer.encode(think_start, add_special_tokens=False)),
tuple(tokenizer.encode(think_end, add_special_tokens=False)),
)
return (None, None, None, None)
def _is_xtml_vocab(vocab):
return all(
t in vocab for t in ("<|open|>", "<|close|>", "<|sep|>", "<|end_of_msg|>")
)
def _infer_structural_markers(tokenizer):
if _is_xtml_vocab(tokenizer.get_vocab()):
return (
"<|open|>response<|sep|>",
"<|close|>response<|sep|>",
"<|close|>message<|sep|>",
)
return ()
def _infer_thinking_kwarg(tokenizer):
custom_renderer = (
getattr(type(tokenizer), "apply_chat_template", None)
is not PreTrainedTokenizerBase.apply_chat_template
)
if custom_renderer:
try:
params = inspect.signature(type(tokenizer).apply_chat_template).parameters
if "thinking" in params and "enable_thinking" not in params:
return "thinking", custom_renderer
except (ValueError, TypeError):
pass
return "enable_thinking", custom_renderer
class TokenizerWrapper:
"""A wrapper that combines an HF tokenizer and a detokenizer.
Accessing any attribute other than the ``detokenizer`` is forwarded to the
huggingface tokenizer.
"""
def __init__(
self,
tokenizer,
detokenizer_class=NaiveStreamingDetokenizer,
eos_token_ids=None,
chat_template=None,
tool_call_start=None,
tool_call_end=None,
tool_parser=None,
):
self._tokenizer = tokenizer
# Built once, since building the token map is expensive.
self._detokenizer = detokenizer_class(tokenizer)
self._eos_token_ids = set(eos_token_ids or [])
if tokenizer.eos_token_id is not None:
self._eos_token_ids.add(tokenizer.eos_token_id)
(
self._think_start,
self._think_end,
self._think_start_tokens,
self._think_end_tokens,
) = _infer_thinking(tokenizer)
self._structural_markers = _infer_structural_markers(tokenizer)
self._chat_template = chat_template
self._thinking_kwarg, has_custom_renderer = _infer_thinking_kwarg(tokenizer)
self.has_chat_template = (
tokenizer.chat_template is not None
or chat_template is not None
or has_custom_renderer
)
self._tool_parser = tool_parser
self._tool_call_start = tool_call_start
self._tool_call_end = tool_call_end
self._tool_call_start_tokens = None
self._tool_call_end_tokens = None
if tool_call_start is not None:
self._tool_call_start_tokens = tuple(
tokenizer.encode(tool_call_start, add_special_tokens=False)
)
self._tool_call_end_tokens = tuple(
tokenizer.encode(tool_call_end, add_special_tokens=False)
)
def apply_chat_template(self, *args, tokenize=True, **kwargs):
if self._thinking_kwarg != "enable_thinking" and "enable_thinking" in kwargs:
kwargs[self._thinking_kwarg] = kwargs.pop("enable_thinking")
if self._thinking_kwarg not in kwargs:
kwargs[self._thinking_kwarg] = self.has_thinking
if self._chat_template is not None:
out = self._chat_template(*args, **kwargs)
if tokenize:
out = self._tokenizer.encode(out, add_special_tokens=False)
return out
kwargs["return_dict"] = False
return self._tokenizer.apply_chat_template(*args, tokenize=tokenize, **kwargs)
def add_eos_token(self, token: str):
token_id = None
try:
token_id = int(token)
except ValueError:
token_id = self._tokenizer.convert_tokens_to_ids(token)
if token_id is None:
raise ValueError(f"'{token}' is not a token for this tokenizer")
self._eos_token_ids.add(token_id)
@staticmethod
def _find(tokens, sequence, start=None, end=None, reverse=False):
start = max(start or 0, 0)
end = end or len(tokens)
outer_loop = (
range(end - len(sequence), start - 1, -1)
if reverse
else range(start, end - len(sequence) + 1)
)
for i in outer_loop:
if tokens[i] == sequence[0]:
if all(tokens[i + j] == sequence[j] for j in range(1, len(sequence))):
return i
return -1
def find_think_start(self, tokens, start=None, end=None):
return self._find(tokens, self._think_start_tokens, start=start, end=end)
def rfind_think_start(self, tokens, start=None, end=None):
return self._find(
tokens, self._think_start_tokens, start=start, end=end, reverse=True
)
def find_think_end(self, tokens, start=None, end=None):
return self._find(tokens, self._think_end_tokens, start=start, end=end)
def rfind_think_end(self, tokens, start=None, end=None):
return self._find(
tokens, self._think_end_tokens, start=start, end=end, reverse=True
)
@property
def has_thinking(self):
return self._think_start is not None
@property
def think_start(self):
return self._think_start
@property
def think_start_id(self):
if self._think_start_tokens is None:
return None
if len(self._think_start_tokens) > 1:
raise ValueError("The start thinking sequence is more than 1 token")
return self._think_start_tokens[0]
@property
def think_start_tokens(self):
return self._think_start_tokens
@property
def think_end(self):
return self._think_end
@property
def think_end_id(self):
if self._think_end_tokens is None:
return None
if len(self._think_end_tokens) > 1:
raise ValueError("The end thinking sequence is more than 1 token")
return self._think_end_tokens[0]
@property
def think_end_tokens(self):
return self._think_end_tokens
@property
def has_tool_calling(self):
return self._tool_call_start is not None
@property
def tool_call_start(self):
return self._tool_call_start
@property
def tool_call_start_tokens(self):
return self._tool_call_start_tokens
@property
def tool_call_end(self):
return self._tool_call_end
@property
def tool_call_end_tokens(self):
return self._tool_call_end_tokens
@property
def structural_markers(self):
return self._structural_markers
@property
def tool_parser(self):
return self._tool_parser
@property
def detokenizer(self):
"""
Get a stateful streaming detokenizer.
"""
# A copy per caller, since requests are detokenized concurrently.
detokenizer = copy.copy(self._detokenizer)
detokenizer.reset()
return detokenizer
@property
def eos_token_ids(self):
return self._eos_token_ids
@eos_token_ids.setter
def eos_token_ids(self, value):
self._eos_token_ids = set(value) if value is not None else set()
def __len__(self):
# Special methods bypass __getattr__, so proxy this one explicitly.
return len(self._tokenizer)
def __getattr__(self, attr):
# Names this class defines are not delegated, so a property that
# raises reports its own error.
if attr.startswith("_") or hasattr(type(self), attr):
raise AttributeError(
f"{type(self).__name__!r} object has no attribute {attr!r}"
)
return getattr(self._tokenizer, attr)
def __setattr__(self, attr, value):
# Defer to the class so properties keep their setters.
if attr.startswith("_") or hasattr(type(self), attr):
super().__setattr__(attr, value)
else:
setattr(self._tokenizer, attr, value)
class NewlineTokenizer(PreTrainedTokenizerFast):
"""A tokenizer that replaces newlines with and with new line."""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def _preprocess_text(self, text):
return text.replace("\n", "")
def _postprocess_text(self, text):
return text.replace("", "\n")
def encode(self, text, **kwargs):
return super().encode(self._preprocess_text(text), **kwargs)
def encode_batch(self, texts, **kwargs):
return super().encode_batch([self._preprocess_text(t) for t in texts], **kwargs)
def decode(self, *args, **kwargs):
return self._postprocess_text(super().decode(*args, **kwargs))
def batch_decode(self, *args, **kwargs):
decoded = super().batch_decode(*args, **kwargs)
return [self._postprocess_text(d) for d in decoded]
AutoTokenizer.register(NewlineTokenizer, fast_tokenizer_class=NewlineTokenizer)
def _match(a, b):
if type(a) != type(b):
return False
if isinstance(a, dict):
return len(a) == len(b) and all(k in b and _match(a[k], b[k]) for k in a)
if isinstance(a, list):
return len(a) == len(b) and all(_match(ai, bi) for ai, bi in zip(a, b))
return a == b
def _is_spm_decoder(decoder):
_target_description = {
"type": "Sequence",
"decoders": [
{"type": "Replace", "pattern": {"String": "▁"}, "content": " "},
{"type": "ByteFallback"},
{"type": "Fuse"},
{"type": "Strip", "content": " ", "start": 1, "stop": 0},
],
}
return _match(_target_description, decoder)
def _is_spm_decoder_no_space(decoder):
_target_description = {
"type": "Sequence",
"decoders": [
{"type": "Replace", "pattern": {"String": "▁"}, "content": " "},
{"type": "ByteFallback"},
{"type": "Fuse"},
],
}
return _match(_target_description, decoder)
def _is_bpe_decoder(decoder):
return isinstance(decoder, dict) and decoder.get("type", None) == "ByteLevel"
def _infer_tool_parser(tokenizer):
"""Attempt to auto-infer a tool parser from the chat template or vocab."""
chat_template = tokenizer.chat_template
if isinstance(chat_template, str):
if "" in chat_template:
return "minimax_m2"
elif "<|tool_call>" in chat_template and "" in chat_template:
return "gemma4"
elif "" in chat_template:
return "function_gemma"
elif "" in chat_template:
return "longcat"
elif "" in chat_template:
return "glm47"
elif (
"<|tool_list_start|>" in chat_template
or "<|tool_call_start|>" in chat_template
):
return "pythonic"
elif (
"\\n\n" in chat_template:
return "kimi_k2"
elif "[TOOL_CALLS]" in chat_template:
return "mistral"
elif "" in chat_template and "tool_call.name" in chat_template:
return "json_tools"
# No template match, so fall back to the vocab. LFM2.5 conversions for
# example have the tool-call tokens but do not name them in the template.
vocab = tokenizer.get_vocab()
if _is_xtml_vocab(vocab):
return "kimi_k3"
elif "<|tool_call_start|>" in vocab and "<|tool_call_end|>" in vocab:
return "pythonic"
return None
def load(
model_path,
tokenizer_config_extra: Optional[Dict[str, Any]] = None,
eos_token_ids=None,
) -> TokenizerWrapper:
"""Load a huggingface tokenizer and try to infer the type of streaming
detokenizer to use.
Note, to use a fast streaming tokenizer, pass a local file path rather than
a Hugging Face repo ID.
"""
detokenizer_class = NaiveStreamingDetokenizer
tokenizer_file = model_path / "tokenizer.json"
if tokenizer_file.exists():
with open(tokenizer_file, "r", encoding="utf-8") as fid:
try:
tokenizer_content = json.load(fid)
except JSONDecodeError as e:
raise JSONDecodeError(
"Failed to parse tokenizer.json", e.doc, e.pos
) from e
if "decoder" in tokenizer_content:
if _is_spm_decoder(tokenizer_content["decoder"]):
detokenizer_class = SPMStreamingDetokenizer
elif _is_spm_decoder_no_space(tokenizer_content["decoder"]):
detokenizer_class = partial(SPMStreamingDetokenizer, trim_space=False)
elif _is_bpe_decoder(tokenizer_content["decoder"]):
detokenizer_class = BPEStreamingDetokenizer
if isinstance(eos_token_ids, int):
eos_token_ids = [eos_token_ids]
chat_template = None
tokenizer = AutoTokenizer.from_pretrained(
model_path, **(tokenizer_config_extra or {})
)
tokenizer_config = tokenizer.init_kwargs
if chat_template_type := tokenizer_config.get("chat_template_type", False):
chat_template = importlib.import_module(
f"mlx_lm.chat_templates.{chat_template_type}"
).apply_chat_template
tool_parser_type = tokenizer_config.get("tool_parser_type")
if tool_parser_type is None:
tool_parser_type = _infer_tool_parser(tokenizer)
if tool_parser_type is not None:
tool_module = importlib.import_module(f"mlx_lm.tool_parsers.{tool_parser_type}")
tool_parser = tool_module.parse_tool_call
tool_call_start = tool_module.tool_call_start
tool_call_end = tool_module.tool_call_end
tokenizer_config["tool_parser_type"] = tool_parser_type
else:
tool_parser = None
tool_call_start = None
tool_call_end = None
return TokenizerWrapper(
tokenizer,
detokenizer_class,
eos_token_ids=eos_token_ids,
chat_template=chat_template,
tool_parser=tool_parser,
tool_call_start=tool_call_start,
tool_call_end=tool_call_end,
)
def no_bos_or_eos(sequence: List, bos: int, eos: int) -> List:
removed_bos = sequence if sequence[0] != bos else sequence[1:]
return removed_bos[:-1] if removed_bos[-1] == eos else removed_bos