minpeter
/

Kimi-K3 / encoding_k3.py
minpeter's picture
Duplicate from moonshotai/Kimi-K3
ffd523a
Raw
History Blame Contribute Delete
22.8 kB
"""Kimi K3 XTML encoding helpers.
This module keeps chat rendering in Python.
Callers that need token IDs should consume ``EncodeSegment`` objects directly:
structural markers may be encoded as tiktoken special tokens, while user/tool
text and attribute values are encoded as ordinary text.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from typing import Any, Iterable, Optional
OPEN_TOKEN = "<|open|>"
CLOSE_TOKEN = "<|close|>"
SEP_TOKEN = "<|sep|>"
END_OF_MSG_TOKEN = "<|end_of_msg|>"
IMAGE_PLACEHOLDER = "<|kimi_image_placeholder|>"
_VALID_THINKING_EFFORTS = {"low", "high", "max"}
@dataclass(frozen=True)
class EncodeSegment:
text: str
allow_special: bool = False
class _ImagePromptState:
def __init__(self, image_prompts: Optional[list[str]] = None):
self.image_prompts = image_prompts
self.index = 0
def next_prompt(self) -> str:
if self.image_prompts is None:
return IMAGE_PLACEHOLDER
if self.index >= len(self.image_prompts):
raise ValueError("More image placeholders than image prompts.")
prompt = self.image_prompts[self.index]
self.index += 1
return prompt
def assert_consumed(self) -> None:
if self.image_prompts is None:
return
if self.index != len(self.image_prompts):
raise ValueError(
f"image prompt count {len(self.image_prompts)} != "
f"consumed placeholder count {self.index}"
)
def _segment(text: Any, *, allow_special: bool = False) -> list[EncodeSegment]:
text = str(text)
if not text:
return []
return [EncodeSegment(text, allow_special=allow_special)]
def _control(text: str) -> list[EncodeSegment]:
return _segment(text, allow_special=True)
def _text(text: Any) -> list[EncodeSegment]:
return _segment(text, allow_special=False)
def _append_text(
segments: list[EncodeSegment],
text: Any,
image_state: _ImagePromptState,
) -> None:
text = str(text)
if text == "":
return
if image_state.image_prompts is None or IMAGE_PLACEHOLDER not in text:
segments.extend(_text(text))
return
parts = text.split(IMAGE_PLACEHOLDER)
for i, part in enumerate(parts):
segments.extend(_text(part))
if i < len(parts) - 1:
segments.extend(_segment(image_state.next_prompt(),
allow_special=True))
def _escape_attr_value(value: Any) -> str:
return str(value).replace("&", "&amp;").replace('"', "&quot;")
def _attr(key: str, value: Any) -> list[EncodeSegment]:
return (
_text(f" {key}")
+ _text('="')
+ _text(_escape_attr_value(value))
+ _text('"')
)
def _open_tag(tag: str, attrs: Iterable[tuple[str, Any]] = ()) -> list[EncodeSegment]:
segments: list[EncodeSegment] = []
segments.extend(_control(OPEN_TOKEN))
segments.extend(_text(tag))
for key, value in attrs:
segments.extend(_attr(key, value))
segments.extend(_control(SEP_TOKEN))
return segments
def _close_tag(tag: str) -> list[EncodeSegment]:
segments: list[EncodeSegment] = []
segments.extend(_control(CLOSE_TOKEN))
segments.extend(_text(tag))
segments.extend(_control(SEP_TOKEN))
return segments
def _end_of_msg() -> list[EncodeSegment]:
return _control(END_OF_MSG_TOKEN)
def _json_compact(value: Any) -> str:
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
def _is_mapping(value: Any) -> bool:
return isinstance(value, dict)
def _xtml_type(value: Any) -> str:
if isinstance(value, bool):
return "boolean"
if value is None:
return "null"
if isinstance(value, (int, float)) and not isinstance(value, bool):
return "number"
if isinstance(value, str):
return "string"
if _is_mapping(value):
return "object"
return "array"
def _xtml_value(value: Any) -> str:
if isinstance(value, str):
return value
return json.dumps(value, ensure_ascii=False)
def _get_value(obj: Any, key: str, default: Any = None) -> Any:
if isinstance(obj, dict):
return obj.get(key, default)
return getattr(obj, key, default)
def extract_response_schema(response_format: Any) -> Any:
if response_format is None:
return None
json_schema = _get_value(response_format, "json_schema")
if json_schema is None:
return None
if isinstance(json_schema, dict):
return json_schema.get(
"schema",
json_schema.get("json_schema", json_schema),
)
schema = _get_value(json_schema, "schema")
if schema is not None:
return schema
schema = _get_value(json_schema, "json_schema")
if schema is not None:
return schema
return json_schema
def deep_sort_dict(obj: Any) -> Any:
if isinstance(obj, dict):
return {k: deep_sort_dict(v) for k, v in sorted(obj.items())}
if isinstance(obj, list):
return [deep_sort_dict(item) for item in obj]
return obj
def normalize_tool_arguments(arguments: Any) -> tuple[dict[str, Any], Optional[str]]:
if arguments is None:
return {}, None
if isinstance(arguments, dict):
return arguments, None
if isinstance(arguments, str):
if not arguments.strip():
return {}, None
try:
parsed = json.loads(arguments)
except json.JSONDecodeError:
return {}, arguments
if not isinstance(parsed, dict):
raise ValueError("Kimi K3 tool call arguments must be a JSON object.")
return parsed, None
raise TypeError(
"Kimi K3 tool call arguments must be a dict or a JSON object string."
)
def normalize_message(message: Any) -> Any:
if not isinstance(message, dict):
return message
normalized = dict(message)
tools = normalized.get("tools")
if tools is not None:
normalized["tools"] = deep_sort_dict(tools)
tool_calls = normalized.get("tool_calls")
if not tool_calls:
return normalized
normalized_calls = []
for tool_call in tool_calls:
if not isinstance(tool_call, dict):
normalized_calls.append(tool_call)
continue
tc = dict(tool_call)
function = tc.get("function")
if isinstance(function, dict):
fn = dict(function)
arguments, json_block = normalize_tool_arguments(fn.get("arguments"))
fn["arguments"] = arguments
if json_block is None:
fn.pop("_xtml_json_block", None)
else:
fn["_xtml_json_block"] = json_block
tc["function"] = fn
else:
arguments, json_block = normalize_tool_arguments(tc.get("arguments"))
tc["arguments"] = arguments
if json_block is None:
tc.pop("_xtml_json_block", None)
else:
tc["_xtml_json_block"] = json_block
normalized_calls.append(tc)
normalized["tool_calls"] = normalized_calls
return normalized
def normalize_conversation(conversation: Any) -> Any:
if not isinstance(conversation, list):
return conversation
def normalize_messages(messages: list[Any]) -> list[Any]:
return [normalize_message(message) for message in messages]
if conversation and isinstance(conversation[0], list):
return [normalize_messages(messages) for messages in conversation]
return normalize_messages(conversation)
def _tool_call_id_index(tool_calls: Any) -> dict:
"""Map assistant ``tool_calls[].id`` to ``(1-based position, function name)``.
The position mirrors the chat template's enumeration over ``tool_calls``
(every entry advances the position, even an id-less one). Duplicate ids keep
their first occurrence.
"""
index: dict = {}
if not isinstance(tool_calls, list):
return index
for position, tool_call in enumerate(tool_calls, start=1):
if not isinstance(tool_call, dict):
continue
call_id = tool_call.get("id")
if call_id is None:
continue
key = str(call_id)
if key in index:
continue
function = tool_call.get("function")
name = (
function.get("name") if isinstance(function, dict) else tool_call.get("name")
)
index[key] = (position, name)
return index
def normalize_xtml_tool_result_messages(messages: list[Any]) -> list[Any]:
"""Re-sort K3 XTML tool results into assistant ``tool_calls`` order.
Serving frameworks generally deliver tool results already in call order. A
direct Transformers caller, however, may pass OpenAI-style tool messages in any
order, so each run of consecutive tool messages is matched against the most
recent preceding assistant ``tool_calls`` by opaque ``tool_call_id`` ==
``tool_calls[].id`` (K3 drops the ``func:index`` format requirement) and
sorted by the matched 1-based position. The matched call is authoritative,
so each matched message's ``tool`` is set to that call's function name --
this keeps an explicit (and possibly stale) ``tool``/``name`` from drifting
out of sync with the reordered position. ``index`` is still derived from the
rendered position by the chat template. A run that cannot be fully matched is
left untouched. Re-running is idempotent.
This function is side-effect free: matched tool messages are shallow-copied
before their ``tool``/``name`` is rewritten, and every other message is
appended to the output as-is. The input list and its message objects are
never mutated.
"""
if not isinstance(messages, list):
return messages
output: list[Any] = []
current_index: dict = {}
i = 0
n = len(messages)
while i < n:
message = messages[i]
if isinstance(message, dict) and message.get("role") == "assistant":
tool_calls = message.get("tool_calls")
current_index = _tool_call_id_index(tool_calls) if tool_calls else {}
output.append(message)
i += 1
continue
if not isinstance(message, dict) or message.get("role") != "tool":
output.append(message)
i += 1
continue
run: list[tuple] = [] # (position, original_offset, message, name)
unresolved = False
offset = 0
while (
i < n and isinstance(messages[i], dict) and messages[i].get("role") == "tool"
):
tool_message = messages[i]
call_id = tool_message.get("tool_call_id", tool_message.get("id"))
matched = current_index.get(str(call_id)) if call_id is not None else None
if matched is None:
unresolved = True
run.append((None, offset, tool_message, None))
else:
position, name = matched
run.append((position, offset, tool_message, name))
offset += 1
i += 1
if unresolved:
output.extend(item[2] for item in run)
else:
run.sort(key=lambda item: (item[0], item[1]))
for _, _, tool_message, name in run:
if name is None:
output.append(tool_message)
continue
# The id-matched call is authoritative: align tool (and any
# explicit name) so the rendered XTML tool attribute cannot
# disagree with the reordered position. Copy first so the
# caller's message object is never mutated.
resolved = dict(tool_message)
resolved["tool"] = name
if "name" in resolved:
resolved["name"] = name
output.append(resolved)
return output
def is_batched_conversation(conversation: Any) -> bool:
return (
isinstance(conversation, list)
and bool(conversation)
and isinstance(conversation[0], list)
)
def _render_content_segments(
content: Any,
image_state: _ImagePromptState,
) -> list[EncodeSegment]:
segments: list[EncodeSegment] = []
if isinstance(content, str):
_append_text(segments, content, image_state)
elif content is not None:
for part in content:
if part["type"] in ["image", "image_url"]:
segments.extend(
_segment(image_state.next_prompt(), allow_special=True))
else:
_append_text(segments, part["text"], image_state)
return segments
def _internal_system_message(message_type: str, body: str) -> list[EncodeSegment]:
segments: list[EncodeSegment] = []
segments.extend(_open_tag("message", [("role", "system"), ("type", message_type)]))
segments.extend(_text(body.strip()))
segments.extend(_close_tag("message"))
segments.extend(_end_of_msg())
return segments
def _render_assistant_segments(
message: dict[str, Any],
image_state: _ImagePromptState,
thinking: bool = True,
) -> list[EncodeSegment]:
segments: list[EncodeSegment] = []
# The <think> channel is structural: in thinking mode every assistant
# message carries the open/close tags even when there is no reasoning
# content to fill in. In non-thinking mode the channel is dropped
# entirely.
if thinking:
reasoning_content = message.get("reasoning_content") or message.get(
"reasoning"
)
segments.extend(_open_tag("think"))
if reasoning_content is not None and str(reasoning_content).strip():
_append_text(segments, reasoning_content, image_state)
segments.extend(_close_tag("think"))
segments.extend(_open_tag("response"))
segments.extend(_render_content_segments(message.get("content"), image_state))
segments.extend(_close_tag("response"))
tool_calls = message.get("tool_calls")
if tool_calls:
segments.extend(_open_tag("tools"))
for index, tool_call in enumerate(tool_calls, start=1):
fn = tool_call.get("function", tool_call)
segments.extend(
_open_tag("call", [("tool", fn["name"]), ("index", index)])
)
args = fn.get("arguments", {})
json_block = fn.get("_xtml_json_block")
if json_block is not None:
segments.extend(_open_tag("json", [("type", "object")]))
_append_text(segments, json_block, image_state)
segments.extend(_close_tag("json"))
elif _is_mapping(args):
for key, value in args.items():
segments.extend(
_open_tag(
"argument",
[("key", key), ("type", _xtml_type(value))],
)
)
_append_text(segments, _xtml_value(value), image_state)
segments.extend(_close_tag("argument"))
segments.extend(_close_tag("call"))
segments.extend(_close_tag("tools"))
return segments
def _render_tool_declare(tools: Any, *, dynamic: bool = False) -> list[EncodeSegment]:
if dynamic:
body = (
"## New Tools Available\n"
"The system dynamically extends the toolset via lazy-loading.\n"
"You have access to all existing and extended tools.\n"
"Here are the specs for the extended tools.\n\n"
"```json\n"
f"{_json_compact(tools)}\n"
"```"
)
else:
body = (
"# Tools\n"
"Here are the available tools, described in JSONSchema.\n\n"
"```json\n"
f"{_json_compact(tools)}\n"
"```"
)
segments: list[EncodeSegment] = []
segments.extend(_open_tag("message", [("role", "system"), ("type", "tool-declare")]))
segments.extend(_text(body))
segments.extend(_close_tag("message"))
segments.extend(_end_of_msg())
return segments
def build_chat_segments(
messages: list[Any],
tools: Optional[list[dict]] = None,
*,
add_generation_prompt: bool = True,
thinking: bool = True,
image_prompts: Optional[list[str]] = None,
**kwargs: Any,
) -> list[EncodeSegment]:
# Re-sort tool results by tool_call_id at the lowest layer so every caller
# (processor or direct tokenizer) gets correctly ordered XTML. The helper is
# side-effect free, so the caller's message objects are left untouched.
messages = normalize_xtml_tool_result_messages(messages)
messages = normalize_conversation(messages)
tools = deep_sort_dict(tools)
kwargs = dict(kwargs)
response_format = kwargs.get("response_format")
if "response_schema" not in kwargs:
response_schema = extract_response_schema(response_format)
if response_schema is not None:
kwargs["response_schema"] = response_schema
if kwargs.get("response_schema") is not None:
kwargs["response_schema"] = deep_sort_dict(kwargs["response_schema"])
image_state = _ImagePromptState(image_prompts)
segments: list[EncodeSegment] = []
tool_calls = None
tool_index = 0
if tools:
segments.extend(_render_tool_declare(tools))
thinking_effort = kwargs.get("thinking_effort")
if thinking and thinking_effort is not None:
assert thinking_effort in _VALID_THINKING_EFFORTS, (
f"Unsupported thinking_effort={thinking_effort!r}; "
f"supported values are {sorted(_VALID_THINKING_EFFORTS)}."
)
if thinking and thinking_effort in _VALID_THINKING_EFFORTS:
segments.extend(
_internal_system_message(
"thinking-effort",
"`thinking_effort` guides on how much to think in your "
"thinking channel (not including the response channel), "
"supported values include `low`, `medium`, `high`, and `max`.\n"
f"Now the system is invoked with `thinking_effort={thinking_effort}`.",
)
)
for message_index, message in enumerate(messages):
if not isinstance(message, dict):
continue
role = message["role"]
if role == "user":
attrs = [("role", "user")]
if message.get("name"):
attrs.append(("name", message["name"]))
segments.extend(_open_tag("message", attrs))
segments.extend(_render_content_segments(message.get("content"), image_state))
segments.extend(_close_tag("message"))
segments.extend(_end_of_msg())
elif role == "system" and message.get("tools"):
segments.extend(_render_tool_declare(message["tools"], dynamic=True))
elif role == "system":
attrs = [("role", "system")]
if message.get("name"):
attrs.append(("name", message["name"]))
segments.extend(_open_tag("message", attrs))
segments.extend(_render_content_segments(message.get("content"), image_state))
segments.extend(_close_tag("message"))
segments.extend(_end_of_msg())
elif role == "tool":
tool_index += 1
tool_name = message.get("tool", message.get("name"))
if (
tool_name is None
and tool_calls is not None
and tool_index <= len(tool_calls)
):
tc = tool_calls[tool_index - 1]
fn = tc.get("function", tc)
tool_name = fn["name"]
if tool_name is None:
raise ValueError(
"Kimi K3 tool messages need a resolvable tool name: "
"carry `tool`/`name`, or match a preceding assistant "
"tool_call by order."
)
segments.extend(
_open_tag(
"message",
[("role", "tool"), ("tool", tool_name), ("index", tool_index)],
)
)
segments.extend(_render_content_segments(message.get("content"), image_state))
segments.extend(_close_tag("message"))
segments.extend(_end_of_msg())
elif role == "assistant":
tool_calls = message.get("tool_calls")
tool_index = 0
attrs = [("role", "assistant")]
if message.get("name"):
attrs.append(("name", message["name"]))
segments.extend(_open_tag("message", attrs))
segments.extend(_render_assistant_segments(message, image_state, thinking))
segments.extend(_close_tag("message"))
segments.extend(_end_of_msg())
tool_choice = kwargs.get("tool_choice")
if tool_choice == "required":
segments.extend(
_internal_system_message(
"tool-choice",
"The system is invoked with `tool_choice=required`.\n"
"You MUST call tools in the next message.",
)
)
elif tool_choice == "none":
segments.extend(
_internal_system_message(
"tool-choice",
"The system is invoked with `tool_choice=none`.\n"
"You MUST NOT call any tools in the next message.",
)
)
rf = kwargs.get("response_format")
rf_type = _get_value(rf, "type", rf) if isinstance(rf, dict) else rf
if rf_type == "json_object":
segments.extend(
_internal_system_message(
"response-format",
"The system is invoked with `response_format=json_object`.\n"
"Your response must be raw JSON data without markdown code "
"blocks (```json) or any additional formatting.",
)
)
elif rf_type == "json_schema":
schema = _json_compact(kwargs.get("response_schema"))
segments.extend(
_internal_system_message(
"response-format",
"The system is invoked with `response_format=json_schema`.\n"
"Your response must be raw JSON data without markdown code "
"blocks (```json) or any additional formatting.\n"
"The JSON data must match the following schema:\n"
f"```json\n{schema}\n```",
)
)
if add_generation_prompt:
segments.extend(_open_tag("message", [("role", "assistant")]))
segments.extend(_open_tag("think" if thinking else "response"))
image_state.assert_consumed()
return segments