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:
Nilay 2026-07-28 01:05:34 +05:30 committed by GitHub
commit 9e2b47d2b5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 220 additions and 6 deletions

View file

@ -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(

View file

@ -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