| 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 |
|
|