"""Modilify chat-template rendering for MLX inference.""" from __future__ import annotations import hashlib import json import re from typing import Any GEMMA_THOUGHT_CLOSE = "" _THINK_LINE_RE = re.compile(r"^think(?:\r?\n|$)") _CHANNEL_BLOCK_RE = re.compile( r"<\|channel>thought\n(.*?)\n?\s*(.*)", flags=re.DOTALL, ) _LITERAL_THINK_RE = re.compile( r"\s*(.*?)\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 = "" in text or "" 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 `` 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