Studio: split parallel tool calls for Llama 3.x chat templates (#7426)
* split parallel tool calls for single-call-only chat templates * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
This commit is contained in:
parent
56fb522746
commit
9e2b47d2b5
2 changed files with 220 additions and 6 deletions
|
|
@ -326,6 +326,58 @@ def _normalize_tool_call_arguments(messages: list) -> list:
|
|||
return out if mutated else messages
|
||||
|
||||
|
||||
def _take_tool_result(pending: list, call_id) -> Optional[dict]:
|
||||
if call_id:
|
||||
for i, result in enumerate(pending):
|
||||
if result.get("tool_call_id") == call_id:
|
||||
return pending.pop(i)
|
||||
for i, result in enumerate(pending):
|
||||
if not result.get("tool_call_id"):
|
||||
return pending.pop(i)
|
||||
return None
|
||||
|
||||
|
||||
def _split_parallel_tool_calls(messages: list) -> list:
|
||||
"""Llama 3.x templates render one call per message, so split parallel calls
|
||||
into consecutive single-call messages, each followed by its own result."""
|
||||
if not any(isinstance(m, dict) and len(m.get("tool_calls") or ()) > 1 for m in messages):
|
||||
return messages
|
||||
|
||||
out: list = []
|
||||
i = 0
|
||||
total = len(messages)
|
||||
while i < total:
|
||||
msg = messages[i]
|
||||
calls = msg.get("tool_calls") if isinstance(msg, dict) else None
|
||||
if not calls or len(calls) <= 1:
|
||||
out.append(msg)
|
||||
i += 1
|
||||
continue
|
||||
|
||||
# Tool results right after this message answer its calls.
|
||||
j = i + 1
|
||||
pending: list = []
|
||||
while (
|
||||
j < total
|
||||
and isinstance(messages[j], dict)
|
||||
and messages[j].get("role") in ("tool", "ipython")
|
||||
):
|
||||
pending.append(messages[j])
|
||||
j += 1
|
||||
|
||||
for idx, call in enumerate(calls):
|
||||
piece = {**msg, "tool_calls": [call]}
|
||||
if idx:
|
||||
piece["content"] = ""
|
||||
out.append(piece)
|
||||
result = _take_tool_result(pending, call.get("id") if isinstance(call, dict) else None)
|
||||
if result is not None:
|
||||
out.append(result)
|
||||
out.extend(pending)
|
||||
i = j
|
||||
return out
|
||||
|
||||
|
||||
def apply_chat_template_for_generation(
|
||||
tokenizer,
|
||||
messages: list,
|
||||
|
|
@ -378,13 +430,21 @@ def apply_chat_template_for_generation(
|
|||
try:
|
||||
return _render(messages)
|
||||
except Exception:
|
||||
# Strict tool templates reject the JSON-string ``arguments`` form via
|
||||
# TypeError or a broad Jinja raise_exception, so retry with dicts coerced.
|
||||
# Original messages render first, so working templates stay byte-identical.
|
||||
# Retry with repairs applied cumulatively. Originals render first, so
|
||||
# working templates stay byte-identical.
|
||||
candidates: list = []
|
||||
normalized = _normalize_tool_call_arguments(messages)
|
||||
if normalized is messages:
|
||||
raise
|
||||
return _render(normalized)
|
||||
if normalized is not messages:
|
||||
candidates.append(normalized)
|
||||
split = _split_parallel_tool_calls(normalized)
|
||||
if split is not normalized:
|
||||
candidates.append(split)
|
||||
for candidate in candidates:
|
||||
try:
|
||||
return _render(candidate)
|
||||
except Exception:
|
||||
continue
|
||||
raise
|
||||
|
||||
|
||||
def render_native_template(
|
||||
|
|
|
|||
|
|
@ -6,10 +6,14 @@ from the OpenAI JSON-string form to a dict before rendering. Strict tool
|
|||
templates (e.g. mlx-community Qwen3.5 checkpoints) iterate arguments.items() and
|
||||
raise "Can only get item pairs from a mapping." on the string form when a prior
|
||||
tool call is re-rendered on the next turn (MLX + transformers paths).
|
||||
|
||||
It must likewise split parallel tool calls for templates that render only one
|
||||
call per message (Llama 3.x).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
|
@ -21,6 +25,7 @@ if str(_BACKEND) not in sys.path:
|
|||
|
||||
from core.inference.chat_template_helpers import ( # noqa: E402
|
||||
_normalize_tool_call_arguments,
|
||||
_split_parallel_tool_calls,
|
||||
apply_chat_template_for_generation,
|
||||
)
|
||||
|
||||
|
|
@ -155,3 +160,152 @@ def test_unrelated_template_error_still_propagates_with_dict_args():
|
|||
|
||||
with pytest.raises(ValueError, match = "broken"):
|
||||
apply_chat_template_for_generation(_AlwaysRaises(), _conv({"query": "x"}))
|
||||
|
||||
|
||||
def _parallel_conv(
|
||||
*,
|
||||
ids = ("c1", "c2"),
|
||||
results_have_ids = True,
|
||||
content = "sure",
|
||||
):
|
||||
a, b = ids
|
||||
return [
|
||||
{"role": "user", "content": "search then render"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"tool_calls": [
|
||||
{
|
||||
"type": "function",
|
||||
"id": a,
|
||||
"function": {"name": "web_search", "arguments": {"query": "x"}},
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"id": b,
|
||||
"function": {"name": "render_html", "arguments": {"html": "<canvas>"}},
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"name": "web_search",
|
||||
**({"tool_call_id": a} if results_have_ids else {}),
|
||||
"content": "no text",
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"name": "render_html",
|
||||
**({"tool_call_id": b} if results_have_ids else {}),
|
||||
"content": "ok",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class _SingleToolCallTokenizer:
|
||||
"""Mimics the Llama 3.x template: rejects >1 call per message."""
|
||||
|
||||
def apply_chat_template(
|
||||
self,
|
||||
messages,
|
||||
*,
|
||||
tokenize = False,
|
||||
add_generation_prompt = True,
|
||||
**kw,
|
||||
):
|
||||
for msg in messages:
|
||||
if len(msg.get("tool_calls") or ()) > 1:
|
||||
raise ValueError("This model only supports single tool-calls at once!")
|
||||
return "RENDERED"
|
||||
|
||||
|
||||
def test_parallel_calls_split_into_sequential_single_call_turns():
|
||||
out = _split_parallel_tool_calls(_parallel_conv())
|
||||
assert [(m["role"], m.get("name")) for m in out] == [
|
||||
("user", None),
|
||||
("assistant", None),
|
||||
("tool", "web_search"),
|
||||
("assistant", None),
|
||||
("tool", "render_html"),
|
||||
]
|
||||
assert [len(m["tool_calls"]) for m in out if m.get("tool_calls")] == [1, 1]
|
||||
assert out[1]["tool_calls"][0]["function"]["name"] == "web_search"
|
||||
assert out[3]["tool_calls"][0]["function"]["name"] == "render_html"
|
||||
|
||||
|
||||
def test_split_pairs_results_by_tool_call_id_not_position():
|
||||
conv = _parallel_conv()
|
||||
conv[2], conv[3] = conv[3], conv[2] # results arrive out of order
|
||||
out = _split_parallel_tool_calls(conv)
|
||||
assert out[1]["tool_calls"][0]["id"] == "c1" and out[2]["tool_call_id"] == "c1"
|
||||
assert out[3]["tool_calls"][0]["id"] == "c2" and out[4]["tool_call_id"] == "c2"
|
||||
|
||||
|
||||
def test_split_falls_back_to_order_when_results_have_no_ids():
|
||||
out = _split_parallel_tool_calls(_parallel_conv(results_have_ids = False))
|
||||
assert [m["role"] for m in out] == ["user", "assistant", "tool", "assistant", "tool"]
|
||||
assert out[2]["name"] == "web_search" and out[4]["name"] == "render_html"
|
||||
|
||||
|
||||
def test_split_keeps_content_on_first_piece_only():
|
||||
out = _split_parallel_tool_calls(_parallel_conv(content = "sure"))
|
||||
assert out[1]["content"] == "sure"
|
||||
assert out[3]["content"] == ""
|
||||
|
||||
|
||||
def test_split_keeps_unmatched_results_after_the_split():
|
||||
conv = _parallel_conv()
|
||||
del conv[3] # second call never returned a result
|
||||
out = _split_parallel_tool_calls(conv)
|
||||
assert [m["role"] for m in out] == ["user", "assistant", "tool", "assistant"]
|
||||
|
||||
|
||||
def test_split_leaves_later_turns_intact():
|
||||
conv = _parallel_conv() + [
|
||||
{"role": "assistant", "content": "done"},
|
||||
{"role": "user", "content": "thanks"},
|
||||
]
|
||||
out = _split_parallel_tool_calls(conv)
|
||||
assert [m["role"] for m in out[-2:]] == ["assistant", "user"]
|
||||
assert out[-2]["content"] == "done"
|
||||
|
||||
|
||||
def test_single_call_and_plain_conversations_pass_through_unchanged():
|
||||
conv = _conv({"query": "x"})
|
||||
assert _split_parallel_tool_calls(conv) is conv
|
||||
plain = [{"role": "user", "content": "hi"}]
|
||||
assert _split_parallel_tool_calls(plain) is plain
|
||||
|
||||
|
||||
def test_render_succeeds_on_single_call_template_with_parallel_calls():
|
||||
# Regression: two calls in one turn used to break every later render.
|
||||
result = apply_chat_template_for_generation(_SingleToolCallTokenizer(), _parallel_conv())
|
||||
assert result == "RENDERED"
|
||||
|
||||
|
||||
def test_string_arguments_and_parallel_calls_are_repaired_together():
|
||||
conv = _parallel_conv()
|
||||
for call in conv[1]["tool_calls"]:
|
||||
call["function"]["arguments"] = json.dumps(call["function"]["arguments"])
|
||||
|
||||
class _StrictAndSingleCall(_SingleToolCallTokenizer):
|
||||
def apply_chat_template(self, messages, **kw):
|
||||
for msg in messages:
|
||||
for call in msg.get("tool_calls", []) or []:
|
||||
if isinstance(call.get("function", {}).get("arguments"), str):
|
||||
raise TypeError("Can only get item pairs from a mapping.")
|
||||
return super().apply_chat_template(messages, **kw)
|
||||
|
||||
assert apply_chat_template_for_generation(_StrictAndSingleCall(), conv) == "RENDERED"
|
||||
|
||||
|
||||
def test_lenient_template_never_sees_a_split_conversation():
|
||||
seen = {}
|
||||
|
||||
class _Lenient:
|
||||
def apply_chat_template(self, messages, **kw):
|
||||
seen["n"] = len(messages)
|
||||
return "RENDERED"
|
||||
|
||||
apply_chat_template_for_generation(_Lenient(), _parallel_conv())
|
||||
assert seen["n"] == 4 # unsplit
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue