ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
10.6 kB
"""Modilify chat-template rendering for MLX inference."""
from __future__ import annotations
import hashlib
import json
import re
from typing import Any
GEMMA_THOUGHT_CLOSE = "<channel|>"
_THINK_LINE_RE = re.compile(r"^think(?:\r?\n|$)")
_CHANNEL_BLOCK_RE = re.compile(
r"<\|channel>thought\n(.*?)\n?<channel\|>\s*(.*)",
flags=re.DOTALL,
)
_LITERAL_THINK_RE = re.compile(
r"\s*<think>(.*?)</think>\s*(.*)",
flags=re.DOTALL,
)
def _split_thought_and_content(content: Any) -> tuple[str | None, str]:
if not isinstance(content, str):
return None, ""
text = content.strip()
if not text:
return None, ""
if "�" in text:
raise ValueError("Assistant target contains a Unicode replacement character.")
if _THINK_LINE_RE.match(text):
raise ValueError("Assistant target uses ambiguous literal think without channel markers.")
channel_match = _CHANNEL_BLOCK_RE.fullmatch(text)
literal_match = _LITERAL_THINK_RE.fullmatch(text)
has_channel_token = "<|channel>" in text or GEMMA_THOUGHT_CLOSE in text
has_literal_token = "<think>" in text or "</think>" in text
if has_channel_token and channel_match is None:
raise ValueError("Malformed Gemma channel in assistant target.")
if has_literal_token and literal_match is None:
raise ValueError("Malformed `<think>` block in assistant target.")
if channel_match is not None:
thought, answer = channel_match.groups()
elif literal_match is not None:
thought, answer = literal_match.groups()
else:
return None, text
thought = thought.strip()
return (thought or None), answer.strip()
def _assistant_thought_and_content(message: dict[str, Any]) -> tuple[str | None, str]:
thought, content = _split_thought_and_content(message.get("content"))
explicit = message.get("reasoning") or message.get("reasoning_content")
if isinstance(explicit, str) and explicit.strip():
return explicit.strip(), content
return thought, content
def _deserialize_tool_call_arguments(arguments: Any) -> dict[str, Any] | None:
"""Convert OpenAI-style JSON argument strings into the mapping Gemma's template requires."""
if arguments is None or isinstance(arguments, dict):
return arguments
if not isinstance(arguments, str):
raise ValueError(
"chat_template: tool_calls[].function.arguments must be a JSON object "
f"(mapping), not a {type(arguments).__name__}."
)
text = arguments.strip()
if not text:
return {}
try:
parsed = json.loads(text)
except json.JSONDecodeError as error:
raise ValueError(
"chat_template: tool_calls[].function.arguments must be a JSON object "
"(mapping), not a string. Deserialize arguments before passing to "
f"the template: {error}"
) from error
if parsed is None or isinstance(parsed, dict):
return parsed
raise ValueError(
"chat_template: tool_calls[].function.arguments must be a JSON object "
f"(mapping), not a {type(parsed).__name__}."
)
def _stable_tool_call_id(tool_call: dict[str, Any], index: int) -> str:
"""Create a deterministic id for traces that omitted OpenAI call ids."""
payload = json.dumps(
tool_call,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
default=str,
).encode("utf-8")
return f"call_modilify_mk2_{index}_{hashlib.sha1(payload).hexdigest()[:16]}"
def _normalize_message_tool_calls(message: dict[str, Any]) -> dict[str, Any]:
tool_calls = message.get("tool_calls")
if not isinstance(tool_calls, list) or not tool_calls:
return message
updated_calls = list(tool_calls)
changed = False
for index, tool_call in enumerate(tool_calls):
if not isinstance(tool_call, dict):
continue
function = tool_call.get("function")
# Some agent traces use the compact {name, arguments} shape instead
# of OpenAI's {function: {name, arguments}} wrapper.
if not isinstance(function, dict):
name = tool_call.get("name")
if not isinstance(name, str) or not name.strip():
continue
function = {
"name": name,
"arguments": tool_call.get(
"arguments", tool_call.get("input", {})
),
}
changed = True
arguments = function.get("arguments")
parsed = (
arguments
if arguments is None or isinstance(arguments, dict)
else _deserialize_tool_call_arguments(arguments)
)
new_function = dict(function)
if parsed is not arguments:
changed = True
new_function["arguments"] = parsed
new_call = dict(tool_call)
if not isinstance(new_call.get("id"), str) or not new_call["id"]:
new_call["id"] = _stable_tool_call_id(tool_call, index)
changed = True
new_call.setdefault("type", "function")
new_call["function"] = new_function
updated_calls[index] = new_call
if not changed:
return message
updated = dict(message)
updated["tool_calls"] = updated_calls
return updated
def _normalize_assistant_message(message: dict[str, Any]) -> dict[str, Any]:
"""Lift think/channel text into ``reasoning`` and deserialize tool arguments."""
updated = _normalize_message_tool_calls(message)
if updated.get("role") != "assistant":
return updated
thought, content = _assistant_thought_and_content(updated)
content_changed = content != (updated.get("content") or "")
reasoning = updated.get("reasoning")
needs_reasoning = bool(thought) and reasoning != thought
if not content_changed and not needs_reasoning:
return updated
if updated is message:
updated = dict(message)
else:
updated = dict(updated)
if thought:
updated["reasoning"] = thought
updated["content"] = content
return updated
def normalize_chat_template_messages(messages: Any) -> Any:
"""Copy conversations into the official Gemma chat-template message schema."""
if not isinstance(messages, list) or not messages:
return messages
if isinstance(messages[0], list):
normalized_batch = None
for index, conversation in enumerate(messages):
normalized = normalize_chat_template_messages(conversation)
if normalized is conversation:
continue
if normalized_batch is None:
normalized_batch = list(messages)
normalized_batch[index] = normalized
return messages if normalized_batch is None else normalized_batch
normalized_messages = None
for index, message in enumerate(messages):
if not isinstance(message, dict):
continue
updated = _normalize_assistant_message(message)
if updated is message:
continue
if normalized_messages is None:
normalized_messages = list(messages)
normalized_messages[index] = updated
return messages if normalized_messages is None else normalized_messages
def normalize_tool_definitions(tools: Any) -> list[dict[str, Any]] | None:
"""Normalize optional tool declarations and ignore trace-only tool metadata.
The native template accepts OpenAI declarations only. ``data-new`` also
contains JSON-encoded declarations and trace metadata shaped like
``{name, arguments, tool_call_id}``; the latter are executed calls, not
declarations, and must not be passed to ``format_function_declaration``.
"""
if tools is None:
return None
pending: list[Any]
if isinstance(tools, str):
try:
parsed = json.loads(tools)
except json.JSONDecodeError:
return None
pending = parsed if isinstance(parsed, list) else [parsed]
elif isinstance(tools, dict):
pending = [tools]
elif isinstance(tools, list):
pending = list(tools)
else:
return None
normalized: list[dict[str, Any]] = []
for item in pending:
if isinstance(item, str):
try:
item = json.loads(item)
except json.JSONDecodeError:
continue
if isinstance(item, list):
pending.extend(item)
continue
if not isinstance(item, dict):
continue
function = item.get("function")
if isinstance(function, dict):
name = function.get("name")
if not isinstance(name, str) or not name.strip():
continue
declaration = dict(function)
declaration["description"] = declaration.get("description", "")
declaration["parameters"] = declaration.get("parameters") or {}
normalized.append({
"type": "function",
"function": declaration,
})
continue
# Accept the common Anthropic/tool-schema spelling when it really is
# a declaration. Execution records with only `arguments` are skipped.
name = item.get("name")
parameters = item.get("parameters", item.get("input_schema"))
if isinstance(name, str) and name.strip() and isinstance(parameters, dict):
normalized.append({
"type": "function",
"function": {
"name": name,
"description": item.get("description", ""),
"parameters": parameters,
},
})
return normalized or None
def apply_chat_template(
processor: Any,
messages: Any,
*,
think: bool,
return_tensors: str | None = None,
padding: bool | str = False,
tools: Any = None,
) -> Any:
"""Render conversations with the Modilify tokenizer template."""
template_kwargs: dict[str, Any] = {
"tokenize": True,
"add_generation_prompt": True,
"enable_thinking": think,
"return_dict": True,
}
if return_tensors is not None:
template_kwargs["return_tensors"] = return_tensors
if padding:
template_kwargs["padding"] = padding
tools = normalize_tool_definitions(tools)
if tools:
template_kwargs["tools"] = tools
messages = normalize_chat_template_messages(messages)
encoded = processor.apply_chat_template(messages, **template_kwargs)
return encoded