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_END = "" 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