File size: 10,611 Bytes
e4f7326
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
"""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