File size: 2,750 Bytes
be18299 | 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 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 | from transformers import LlamaTokenizerFast
class Llmjp4Tokenizer(LlamaTokenizerFast):
_HARMONY_TOKENS: set[str] = {
"<|start|>",
"<|message|>",
"<|channel|>",
"<|constrain|>",
"<|end|>",
"<|return|>",
"<|call|>",
}
# NOTE(odashi):
# Response schemas are not recognized automatically.
# We need to define them manually.
# https://github.com/huggingface/trl/issues/4609
_RESPONSE_SCHEMA = {
"type": "object",
"properties": {
"role": {"const": "assistant"},
"content": {"type": "string", "x-regex": r"<\|channel\|>final<\|message\|>(.*?)(?:<\|end\|>|<\|return\|>|$)"},
"thinking": {"type": "string", "x-regex": r"<\|channel\|>analysis<\|message\|>(.*?)<\|end\|>"},
"tool_calls": {
"x-regex-iterator": r"<\|channel\|>commentary (to=functions\..*?<\|message\|>.*?)(?:<\|call\|>|$)",
"type": "array",
"items": {
"type": "object",
"properties": {
"type": {"const": "function"},
"function": {
"type": "object",
"properties": {
"name": {"type": "string", "x-regex": r"^to=functions\.(\w+)"},
"arguments": {
"type": "object",
"x-regex": r"<\|message\|>(.*)",
"x-parser": "json",
"additionalProperties": {"type": "any"},
},
},
},
},
},
},
},
}
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.response_schema = self._RESPONSE_SCHEMA
self._harmony_token_ids = {
self.convert_tokens_to_ids(token)
for token in self._HARMONY_TOKENS
}
def _decode(self, token_ids: int | list[int], *args, **kwargs):
if isinstance(token_ids, int):
token_ids = [token_ids]
result: list[str] = []
prev_pos = 0
# NOTE(odashi):
# Ensure that text tokens are decoded without preceding Harmony tokens
# to avoid incorrect addition of whitespaces.
for pos, token_id in enumerate(token_ids, start=1):
if token_id in self._harmony_token_ids or pos == len(token_ids):
result.append(super()._decode(token_ids[prev_pos:pos], *args, **kwargs))
prev_pos = pos
return "".join(result)
|