Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions swift/template/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -1307,6 +1307,47 @@ def _remove_history_thinking(self, inputs) -> None:
elif isinstance(content, str):
message['content'] = self._remove_thinking_content(content)

def _merge_tool_user_inputs(self, inputs: StdTemplateInputs) -> None:
# The pairwise encoder needs one query before each assistant response, so a user turn that
# follows tool results (e.g. the user interrupts after a tool call, or an agent harness injects
# a reminder) would otherwise pair a `tool` query with a `user` response and hit the
# `response_role` assertion in `_swift_encode`. Merge the tool results and the following user
# turns into a single prompt-side query, preserving each native user-turn boundary.
#
# Only applies to ChatML-style agent templates that render tool results as a user turn (those
# exposing `_get_tool_responses`). ReAct-style templates render tool results as observations
# glued to the assistant turn, so this reconstruction does not fit them and they are skipped.
# Access `self._agent_template` (raw field) instead of the `self.agent_template` property here:
# the property raises `ValueError` when no agent template was matched, which `getattr(..., None)`
# would not swallow, so it must not run before this gate on plain (non-agent) templates.
if (self.template_backend != 'swift' or not self.use_chat_template or inputs.is_multimodal
or self._agent_template is None):
return
agent_template = self.agent_template
if not hasattr(agent_template, '_get_tool_responses'):
return
messages = inputs.messages
i = 1
while i < len(messages):
if messages[i]['role'] != 'tool' or messages[i - 1]['role'] not in {'assistant', 'tool_call'}:
i += 1
continue
tool_end = i
while tool_end < len(messages) and messages[tool_end]['role'] == 'tool':
tool_end += 1
user_end = tool_end
while (user_end < len(messages) and messages[user_end]['role'] == 'user'
and isinstance(messages[user_end]['content'], str)):
user_end += 1
if user_end > tool_end:
query_prefix = ''.join(self.template_meta.prompt).split('{{QUERY}}', 1)[0]
separator = ''.join(self.template_meta.chat_sep) + query_prefix
contents = [agent_template._get_tool_responses(messages[i:tool_end])]
contents += [message['content'] for message in messages[tool_end:user_end]]
# Keep the last user message's metadata and last-round loss semantics.
messages[i:user_end] = [{**messages[user_end - 1], 'content': separator.join(contents)}]
i += 1

def _swift_prepare_inputs(self, inputs: StdTemplateInputs):
"""
Preprocesses the list of messages in the input by merging and formatting consecutive messages
Expand All @@ -1325,6 +1366,7 @@ def _swift_prepare_inputs(self, inputs: StdTemplateInputs):
Returns:
None. The input messages list is updated in-place.
"""
self._merge_tool_user_inputs(inputs)
self._preprocess_tool_call(inputs)
self._preprocess_standalone_tools(inputs)
messages = inputs.messages
Expand Down
23 changes: 23 additions & 0 deletions swift/template/templates/glm.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,29 @@ class GLMTemplateMeta(TemplateMeta):
class GLM4Template(Template):
strip_newline = True

def _swift_prepare_inputs(self, inputs: StdTemplateInputs):
super()._swift_prepare_inputs(inputs)
# The pairwise encoder needs one query before each assistant response, so a user turn that
# follows tool results would otherwise pair a `tool` query with a `user` response and hit the
# `response_role` assertion in `_swift_encode`. GLM renders each role independently: a user turn
# after tool results is just a normal `<|user|>` turn placed before the assistant transition that
# `_format_tool_responses` appends. Splice the follow-up user(s) into that tool query accordingly.
if self.template_backend != 'swift' or not self.use_chat_template or inputs.is_multimodal:
return
query_prefix = ''.join(self.template_meta.prompt).split('{{QUERY}}', 1)[0]
messages = inputs.messages
i = 1
while i < len(messages):
pre_message, message = messages[i - 1], messages[i]
tool_content = pre_message['content']
if (message['role'] == 'user' and isinstance(message['content'], str) and pre_message['role'] == 'tool'
and isinstance(tool_content, list) and tool_content and isinstance(tool_content[-1], str)
and '<|assistant|>' in tool_content[-1]):
pre_message['content'] = tool_content[:-1] + [query_prefix + message['content'], tool_content[-1]]
messages.pop(i)
continue
i += 1

def _swift_encode(self, inputs: StdTemplateInputs):
res_context_list, loss_scale_list, answer_len = super()._swift_encode(inputs)
if self.strip_newline:
Expand Down
240 changes: 240 additions & 0 deletions tests/utils/test_tool_user_followup.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,240 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
import copy
import unittest
from tokenizers import Tokenizer, decoders, models, pre_tokenizers
from transformers import PreTrainedTokenizerFast
from types import SimpleNamespace

from swift.template import get_template


def _make_tokenizer():
# Byte-level tokens keep this template regression independent of model downloads.
vocab = {char: i for i, char in enumerate(sorted(pre_tokenizers.ByteLevel.alphabet()))}
backend = Tokenizer(models.BPE(vocab=vocab, merges=[]))
backend.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
backend.decoder = decoders.ByteLevel()
tokenizer = PreTrainedTokenizerFast(
tokenizer_object=backend,
eos_token='<|im_end|>',
pad_token='<|endoftext|>',
additional_special_tokens=['<|im_start|>', '<think>', '</think>'])
tokenizer.model_info = SimpleNamespace(config=SimpleNamespace(), task_type='causal_lm', max_model_len=8192)
tokenizer.model_meta = SimpleNamespace(is_multimodal=False)
return tokenizer


def _make_glm_tokenizer():
# GLM renders tool results with its own control tokens, so a GLM-flavoured byte tokenizer keeps this
# regression independent of model downloads while still exercising the `<|user|>` splice.
vocab = {char: i for i, char in enumerate(sorted(pre_tokenizers.ByteLevel.alphabet()))}
backend = Tokenizer(models.BPE(vocab=vocab, merges=[]))
backend.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
backend.decoder = decoders.ByteLevel()
tokenizer = PreTrainedTokenizerFast(
tokenizer_object=backend,
eos_token='<|endoftext|>',
pad_token='<|endoftext|>',
additional_special_tokens=[
'[gMASK]', '<sop>', '<|system|>', '<|user|>', '<|assistant|>', '<|observation|>', '<think>', '</think>'
])
tokenizer.model_info = SimpleNamespace(config=SimpleNamespace(), task_type='causal_lm', max_model_len=8192)
tokenizer.model_meta = SimpleNamespace(is_multimodal=False)
return tokenizer


class TestQwenToolUserFollowup(unittest.TestCase):
"""Reference byte-match: the root fix must reproduce independent ChatML rendering for qwen3_5."""

@classmethod
def setUpClass(cls):
cls.tokenizer = _make_tokenizer()

def make_template(self, **kwargs):
template = get_template(
self.tokenizer, template_type='qwen3_5', preserve_thinking=True, add_non_thinking_prefix=False, **kwargs)
template.set_mode('train')
return template

@staticmethod
def make_data(n_tools=1, n_users=1):
call = '<think>\nplan\n</think>\n\n<tool_call>\n<function=weather>\n</function>\n</tool_call>'
answer = '<think>\nreply\n</think>\n\nfinal_answer'
messages = [{'role': 'user', 'content': 'question'}, {'role': 'assistant', 'content': call}]
messages += [{'role': 'tool', 'content': f'result_{i}'} for i in range(n_tools)]
messages += [{'role': 'user', 'content': f'followup_{i}'} for i in range(n_users)]
messages.append({'role': 'assistant', 'content': answer})
return {'messages': messages}

def test_followup_boundaries_and_labels(self):
for strategy in ('default', 'last_round'):
for n_tools, n_users in ((1, 0), (2, 0), (1, 1), (2, 1), (1, 2), (2, 2)):
with self.subTest(strategy=strategy, tools=n_tools, users=n_users):
template = self.make_template(loss_scale=strategy)
data = self.make_data(n_tools, n_users)
original = copy.deepcopy(data)
encoded = template.encode(data)
call = data['messages'][1]['content']
answer = data['messages'][-1]['content']
observations = '\n'.join(f'<tool_response>\nresult_{i}\n</tool_response>' for i in range(n_tools))
turns = [('user', 'question'), ('assistant', call), ('user', observations)]
turns += [('user', f'followup_{i}') for i in range(n_users)]
turns.append(('assistant', answer))
# Independent ChatML rendering: every user turn retains its own boundary.
text = '\n'.join(f'<|im_start|>{role}\n{content}<|im_end|>' for role, content in turns) + '\n'
self.assertEqual(self.tokenizer.decode(encoded['input_ids']), text)
self.assertEqual(data, original)
expected_labels = [-100] * len(encoded['input_ids'])
responses = [answer] if strategy == 'last_round' else [call, answer]
for response in responses:
start = text.index(response)
begin = len(self.tokenizer.encode(text[:start], add_special_tokens=False))
tokens = self.tokenizer.encode(response + '<|im_end|>\n', add_special_tokens=False)
expected_labels[begin:begin + len(tokens)] = tokens
self.assertEqual(encoded['labels'], expected_labels)

def test_no_followup_is_unchanged(self):
# The fix must be a no-op when no user turn follows the tool results.
template = self.make_template()
data = self.make_data(n_tools=2, n_users=0)
encoded = template.encode(data)
text = self.tokenizer.decode(encoded['input_ids'])
self.assertIn('<tool_response>\nresult_0\n</tool_response>', text)

def test_openai_and_swift_tool_calls_match(self):
# The follow-up path must treat OpenAI-format tool_calls and native tool_call/tool_response alike.
call = {'name': 'weather', 'arguments': {'city': 'Beijing'}}
data = self.make_data()
data['messages'][1] = {
'role': 'assistant',
'content': '',
'tool_calls': [{
'id': 'call_1',
'type': 'function',
'function': call
}]
}
data['messages'][2]['tool_call_id'] = 'call_1'
native = copy.deepcopy(data)
native['messages'][1] = {'role': 'tool_call', 'content': call}
native['messages'][2]['role'] = 'tool_response'
template = self.make_template()
self.assertEqual(template.encode(data), template.encode(native))

def test_response_loss_weight_is_preserved(self):
# Splicing the follow-up into the query must not disturb the assistant response loss weights.
template = self.make_template(is_binary_loss_scale=False)
data = self.make_data()
data['messages'][-1]['loss_scale'] = 0.4
encoded = template.encode(data)
text = self.tokenizer.decode(encoded['input_ids'])
start = text.index(data['messages'][-1]['content'])
begin = len(self.tokenizer.encode(text[:start], add_special_tokens=False))
length = len(self.tokenizer.encode(data['messages'][-1]['content'], add_special_tokens=False))
self.assertEqual(encoded['loss_scale'][begin:begin + length], [0.4] * length)
for label, weight in zip(encoded['labels'], encoded['loss_scale']):
if label == -100:
self.assertEqual(weight, 0.)


class TestHermesToolUserFollowup(unittest.TestCase):
"""The root fix generalizes: a non-qwen ChatML template stops crashing and keeps the user boundary."""

@classmethod
def setUpClass(cls):
cls.tokenizer = _make_tokenizer()

def make_template(self):
template = get_template(self.tokenizer, template_type='qwen2_5', agent_template='hermes')
template.set_mode('train')
return template

@staticmethod
def make_data():
messages = [
{
'role': 'user',
'content': 'question'
},
{
'role': 'assistant',
'content': '<tool_call>\n{"name": "weather", "arguments": {}}\n</tool_call>'
},
{
'role': 'tool',
'content': 'result_0'
},
{
'role': 'user',
'content': 'followup_0'
},
{
'role': 'assistant',
'content': 'final_answer'
},
]
return {'messages': messages}

def test_does_not_crash_and_keeps_boundary(self):
template = self.make_template()
encoded = template.encode(self.make_data())
text = self.tokenizer.decode(encoded['input_ids'])
self.assertIn('<tool_response>\nresult_0\n</tool_response>', text)
self.assertIn('followup_0', text)
# The follow-up user keeps its own turn boundary rather than being glued to the tool result.
self.assertIn('</tool_response><|im_end|>\n<|im_start|>user\nfollowup_0', text)


class TestGLMToolUserFollowup(unittest.TestCase):
"""GLM renders tool results independently; a follow-up user must become a normal `<|user|>` turn."""

@classmethod
def setUpClass(cls):
cls.tokenizer = _make_glm_tokenizer()

def make_template(self):
template = get_template(self.tokenizer, template_type='glm4_5')
template.set_mode('train')
return template

@staticmethod
def make_data(n_users=1):
call = '<tool_call>weather\n<arg_key>city</arg_key>\n<arg_value>BJ</arg_value>\n</tool_call>'
messages = [
{
'role': 'user',
'content': 'question'
},
{
'role': 'assistant',
'content': call
},
{
'role': 'tool',
'content': 'result_0'
},
]
messages += [{'role': 'user', 'content': f'followup_{i}'} for i in range(n_users)]
messages.append({'role': 'assistant', 'content': 'final_answer'})
return {'messages': messages}

def test_followup_spliced_before_assistant(self):
template = self.make_template()
encoded = template.encode(self.make_data(n_users=1))
text = self.tokenizer.decode(encoded['input_ids'])
# Matches the official jinja: `</tool_response><|user|>\n{followup}<|assistant|>`.
self.assertIn('</tool_response><|user|>\nfollowup_0<|assistant|>', text)
# The follow-up belongs to the query side and must not be supervised.
supervised = self.tokenizer.decode([t for t in encoded['labels'] if t != -100])
self.assertNotIn('followup_0', supervised)

def test_no_followup_is_unchanged(self):
template = self.make_template()
encoded = template.encode(self.make_data(n_users=0))
text = self.tokenizer.decode(encoded['input_ids'])
self.assertIn('</tool_response><|assistant|>', text)
self.assertNotIn('<|user|>\nfollowup', text)


if __name__ == '__main__':
unittest.main()
Loading