Tercet-R-1.0 / tiny_gdn /chatml.py
kerzgrr's picture
Upload Tercet-R-1.0 (instruct EMA @ step 4500)
48eb149 verified
Raw
History Blame Contribute Delete
4.36 kB
from __future__ import annotations
from typing import Any
from tiny_gdn.smoltalk_chat import (
NO_THINK_PREFIX,
THINK_CONTROL_PREFIX,
kwargs_from_inference_request,
materialize_smoltalk_messages,
merge_kwargs_with_system_message,
resolve_reasoning_mode,
)
from tiny_gdn.tools import prepare_inference_messages
THINK_START = "<think>"
THINK_END = "</think>"
def decode_chat_completion(tokenizer: Any, token_ids: list[int]) -> str:
return tokenizer.decode(token_ids, skip_special_tokens=False)
def append_assistant_generation_prompt(
tokenizer: Any,
token_ids: list[int],
*,
enable_thinking: bool | None = None,
) -> None:
im_start = tokenizer.token_to_id("<|im_start|>")
if im_start is None:
raise RuntimeError("Tokenizer is missing <|im_start|>")
token_ids.append(im_start)
token_ids.extend(
tokenizer.encode("assistant\n", add_special_tokens=False).ids
)
if enable_thinking is True:
token_ids.extend(
tokenizer.encode(THINK_CONTROL_PREFIX, add_special_tokens=False).ids
)
elif enable_thinking is False:
token_ids.extend(
tokenizer.encode(NO_THINK_PREFIX, add_special_tokens=False).ids
)
def encode_chatml_messages(
tokenizer: Any,
messages: list[dict[str, Any]],
*,
tools: Any | None = None,
enable_thinking: bool | None = None,
xml_tools: Any | None = None,
python_tools: Any | None = None,
) -> list[int]:
if enable_thinking is None:
prepared = prepare_inference_messages(messages, tools=tools)
assistant_thinking_control = None
else:
kwargs = kwargs_from_inference_request(
enable_thinking=enable_thinking,
xml_tools=xml_tools,
python_tools=python_tools,
tools=tools,
)
prepared = materialize_smoltalk_messages(messages, kwargs)
_, resolved = merge_kwargs_with_system_message(messages, kwargs)
assistant_thinking_control = (
resolve_reasoning_mode(
resolved.enable_thinking,
resolved.custom_instructions,
)
== "/think"
)
required_tokens = {
token: tokenizer.token_to_id(token)
for token in (
"<|begin_of_text|>",
"<|im_start|>",
"<|im_end|>",
)
}
if any(token_id is None for token_id in required_tokens.values()):
raise RuntimeError("Tokenizer is missing ChatML special tokens")
bos_id = required_tokens["<|begin_of_text|>"]
im_start_id = required_tokens["<|im_start|>"]
im_end_id = required_tokens["<|im_end|>"]
if bos_id is None or im_start_id is None or im_end_id is None:
raise RuntimeError("Tokenizer ChatML IDs could not be resolved")
token_ids = [bos_id]
newline_ids = tokenizer.encode("\n", add_special_tokens=False).ids
allowed_roles = {"system", "user", "assistant"}
for message_index, message in enumerate(prepared):
role = message["role"]
content = message["content"]
masked_prefix = message.get("masked_prefix", "")
if role not in allowed_roles:
raise ValueError(
f"Unsupported chat role at index {message_index}: {role!r}"
)
if not isinstance(masked_prefix, str):
raise ValueError(
f"Masked prefix at index {message_index} must be a string"
)
if not content.strip():
raise ValueError(
f"Chat content at index {message_index} must be non-empty"
)
token_ids.append(im_start_id)
token_ids.extend(
tokenizer.encode(f"{role}\n", add_special_tokens=False).ids
)
if masked_prefix:
token_ids.extend(
tokenizer.encode(masked_prefix, add_special_tokens=False).ids
)
token_ids.extend(
tokenizer.encode(content, add_special_tokens=False).ids
)
token_ids.append(im_end_id)
token_ids.extend(newline_ids)
if prepared[-1]["role"] != "user":
raise ValueError("The final chat message must have role 'user'")
append_assistant_generation_prompt(
tokenizer,
token_ids,
enable_thinking=assistant_thinking_control,
)
return token_ids