chat : enable tool call in thinking for DS4 (#26269)
This commit is contained in:
+106
-100
@@ -1943,18 +1943,6 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
|||||||
adjusted_messages = deepseek_v4_sort_tool_results(inputs.messages);
|
adjusted_messages = deepseek_v4_sort_tool_results(inputs.messages);
|
||||||
}
|
}
|
||||||
|
|
||||||
data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, adjusted_messages);
|
|
||||||
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, adjusted_messages);
|
|
||||||
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
|
|
||||||
data.supports_thinking = true;
|
|
||||||
data.thinking_start_tag = "<think>";
|
|
||||||
data.thinking_end_tags = {"</think>"};
|
|
||||||
data.preserved_tokens = {
|
|
||||||
"|DSML|",
|
|
||||||
"<think>",
|
|
||||||
"</think>",
|
|
||||||
};
|
|
||||||
|
|
||||||
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
|
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
|
||||||
auto has_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object();
|
auto has_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object();
|
||||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||||
@@ -1972,6 +1960,18 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
|||||||
const std::string PARAM_END = "</" + DSML + "parameter>";
|
const std::string PARAM_END = "</" + DSML + "parameter>";
|
||||||
const std::string GEN_PROMPT = "<|Assistant|>";
|
const std::string GEN_PROMPT = "<|Assistant|>";
|
||||||
|
|
||||||
|
data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, adjusted_messages);
|
||||||
|
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, adjusted_messages);
|
||||||
|
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
|
||||||
|
data.supports_thinking = true;
|
||||||
|
data.thinking_start_tag = THINK_START;
|
||||||
|
data.thinking_end_tags = {THINK_END, FC_START};
|
||||||
|
data.preserved_tokens = {
|
||||||
|
DSML,
|
||||||
|
THINK_START,
|
||||||
|
THINK_END,
|
||||||
|
};
|
||||||
|
|
||||||
if (inputs.has_continuation()) {
|
if (inputs.has_continuation()) {
|
||||||
const auto & msg = inputs.continue_msg;
|
const auto & msg = inputs.continue_msg;
|
||||||
|
|
||||||
@@ -1983,13 +1983,101 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
|||||||
data.prompt += data.generation_prompt;
|
data.prompt += data.generation_prompt;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
bool require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||||
|
bool has_tool_calls = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
|
||||||
|
|
||||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||||
auto generation_prompt = p.literal(GEN_PROMPT);
|
auto generation_prompt = p.literal(GEN_PROMPT);
|
||||||
auto end = p.end();
|
auto end = p.end();
|
||||||
|
|
||||||
|
// build tool call section first since we might need it in reasoning
|
||||||
|
auto tool_choice = p.choice();
|
||||||
|
if (has_tool_calls) {
|
||||||
|
foreach_function(inputs.tools, [&](const json & tool) {
|
||||||
|
const auto & function = tool.at("function");
|
||||||
|
std::string name = function.at("name");
|
||||||
|
auto params = function.contains("parameters") ? function.at("parameters") : json::object();
|
||||||
|
const auto & props = params.contains("properties") ? params.at("properties") : json::object();
|
||||||
|
|
||||||
|
std::set<std::string> required;
|
||||||
|
if (params.contains("required")) {
|
||||||
|
params.at("required").get_to(required);
|
||||||
|
}
|
||||||
|
|
||||||
|
auto schema_info = common_schema_info();
|
||||||
|
schema_info.resolve_refs(params);
|
||||||
|
|
||||||
|
std::vector<common_peg_parser> required_parsers;
|
||||||
|
std::vector<common_peg_parser> optional_parsers;
|
||||||
|
for (const auto & [param_name, param_schema] : props.items()) {
|
||||||
|
bool is_required = required.find(param_name) != required.end();
|
||||||
|
bool is_string = schema_info.resolves_to_string(param_schema);
|
||||||
|
|
||||||
|
auto arg = p.tool_arg(
|
||||||
|
p.tool_arg_open(p.literal(PARAM_START + " name=\"") + p.tool_arg_name(p.literal(param_name)) +
|
||||||
|
p.literal("\" string=\"" + std::string(is_string ? "true" : "false") + "\">")) +
|
||||||
|
(is_string ?
|
||||||
|
p.tool_arg_string_value(p.until(PARAM_END)) :
|
||||||
|
p.tool_arg_json_value(p.schema(p.json(), "tool-" + name + "-arg-" + param_name + "-schema",
|
||||||
|
param_schema, false))) +
|
||||||
|
p.tool_arg_close(p.literal(PARAM_END)));
|
||||||
|
|
||||||
|
auto named_arg = p.rule("tool-" + name + "-arg-" + param_name, arg);
|
||||||
|
if (is_required) {
|
||||||
|
required_parsers.push_back(named_arg);
|
||||||
|
} else {
|
||||||
|
optional_parsers.push_back(named_arg);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
common_peg_parser args_seq = p.eps();
|
||||||
|
for (size_t i = 0; i < required_parsers.size(); i++) {
|
||||||
|
if (i > 0) {
|
||||||
|
args_seq = args_seq + p.space();
|
||||||
|
}
|
||||||
|
args_seq = args_seq + required_parsers[i];
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!optional_parsers.empty()) {
|
||||||
|
common_peg_parser any_opt = p.choice();
|
||||||
|
for (const auto & opt : optional_parsers) {
|
||||||
|
any_opt |= opt;
|
||||||
|
}
|
||||||
|
args_seq = args_seq + p.repeat(p.space() + any_opt, 0, -1);
|
||||||
|
}
|
||||||
|
|
||||||
|
common_peg_parser invoke_body = args_seq;
|
||||||
|
auto func_parser = p.tool(p.tool_open(p.literal(INVOKE_START + " name=\"") +
|
||||||
|
p.tool_name(p.literal(name)) + p.literal("\">\n")) +
|
||||||
|
invoke_body + p.space() + p.tool_close(p.literal(INVOKE_END)));
|
||||||
|
|
||||||
|
tool_choice |= p.rule("tool-" + name, func_parser);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
common_peg_parser tool_calls = p.eps();
|
||||||
|
if (inputs.parallel_tool_calls) {
|
||||||
|
tool_calls = p.trigger_rule("tool-call",
|
||||||
|
p.literal(FC_START) + p.space() + tool_choice +
|
||||||
|
p.zero_or_more(p.space() + tool_choice) + p.space() + p.literal(FC_END));
|
||||||
|
} else {
|
||||||
|
tool_calls = p.trigger_rule("tool-call",
|
||||||
|
p.literal(FC_START) + p.space() + tool_choice + p.space() + p.literal(FC_END));
|
||||||
|
}
|
||||||
|
|
||||||
auto reasoning = p.eps();
|
auto reasoning = p.eps();
|
||||||
|
auto reasoning_with_tc = p.eps();
|
||||||
|
auto obligatory_tool_calls = tool_calls;
|
||||||
|
bool allow_reasoning_with_tc = false;
|
||||||
|
|
||||||
|
if (!require_tools) {
|
||||||
|
tool_calls = p.optional(tool_calls);
|
||||||
|
}
|
||||||
|
|
||||||
if (extract_reasoning && inputs.enable_thinking) {
|
if (extract_reasoning && inputs.enable_thinking) {
|
||||||
reasoning = p.optional(THINK_START + p.reasoning(p.until(THINK_END)) + THINK_END);
|
reasoning = p.optional(THINK_START + p.reasoning(p.until(THINK_END)) + THINK_END);
|
||||||
|
reasoning_with_tc = THINK_START + p.reasoning(p.until_one_of({ FC_START, THINK_END })) + obligatory_tool_calls;
|
||||||
|
allow_reasoning_with_tc = true;
|
||||||
} else if (extract_reasoning) {
|
} else if (extract_reasoning) {
|
||||||
// Thinking disabled but reasoning extraction requested: the generation prompt
|
// Thinking disabled but reasoning extraction requested: the generation prompt
|
||||||
// contains an empty <think></think> pair (V3.2) or a bare </think> (V4) that
|
// contains an empty <think></think> pair (V3.2) or a bare </think> (V4) that
|
||||||
@@ -2007,101 +2095,19 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
|||||||
return generation_prompt + reasoning + response_format + end;
|
return generation_prompt + reasoning + response_format + end;
|
||||||
}
|
}
|
||||||
|
|
||||||
if (!has_tools || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_NONE) {
|
if (!has_tool_calls) {
|
||||||
return generation_prompt + reasoning + p.content(p.rest()) + end;
|
return generation_prompt + reasoning + p.content(p.rest()) + end;
|
||||||
}
|
}
|
||||||
|
|
||||||
auto tool_choice = p.choice();
|
auto content_before_tools = p.negate(p.literal(THINK_START)) + p.content(p.until(FC_START));
|
||||||
foreach_function(inputs.tools, [&](const json & tool) {
|
return allow_reasoning_with_tc ? generation_prompt + (reasoning_with_tc | (reasoning + content_before_tools + tool_calls)) + end :
|
||||||
const auto & function = tool.at("function");
|
generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||||
std::string name = function.at("name");
|
|
||||||
auto params = function.contains("parameters") ? function.at("parameters") : json::object();
|
|
||||||
const auto & props = params.contains("properties") ? params.at("properties") : json::object();
|
|
||||||
|
|
||||||
std::set<std::string> required;
|
|
||||||
if (params.contains("required")) {
|
|
||||||
params.at("required").get_to(required);
|
|
||||||
}
|
|
||||||
|
|
||||||
auto schema_info = common_schema_info();
|
|
||||||
schema_info.resolve_refs(params);
|
|
||||||
|
|
||||||
std::vector<common_peg_parser> required_parsers;
|
|
||||||
std::vector<common_peg_parser> optional_parsers;
|
|
||||||
for (const auto & [param_name, param_schema] : props.items()) {
|
|
||||||
bool is_required = required.find(param_name) != required.end();
|
|
||||||
bool is_string = schema_info.resolves_to_string(param_schema);
|
|
||||||
|
|
||||||
auto arg = p.tool_arg(
|
|
||||||
p.tool_arg_open(
|
|
||||||
p.literal(PARAM_START + " name=\"") +
|
|
||||||
p.tool_arg_name(p.literal(param_name)) +
|
|
||||||
p.literal("\" string=\"" + std::string(is_string ? "true" : "false") + "\">")) +
|
|
||||||
(is_string
|
|
||||||
? p.tool_arg_string_value(p.until(PARAM_END))
|
|
||||||
: p.tool_arg_json_value(p.schema(p.json(),
|
|
||||||
"tool-" + name + "-arg-" + param_name + "-schema",
|
|
||||||
param_schema, false))) +
|
|
||||||
p.tool_arg_close(p.literal(PARAM_END)));
|
|
||||||
|
|
||||||
auto named_arg = p.rule("tool-" + name + "-arg-" + param_name, arg);
|
|
||||||
if (is_required) {
|
|
||||||
required_parsers.push_back(named_arg);
|
|
||||||
} else {
|
|
||||||
optional_parsers.push_back(named_arg);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
common_peg_parser args_seq = p.eps();
|
|
||||||
for (size_t i = 0; i < required_parsers.size(); i++) {
|
|
||||||
if (i > 0) {
|
|
||||||
args_seq = args_seq + p.space();
|
|
||||||
}
|
|
||||||
args_seq = args_seq + required_parsers[i];
|
|
||||||
}
|
|
||||||
|
|
||||||
if (!optional_parsers.empty()) {
|
|
||||||
common_peg_parser any_opt = p.choice();
|
|
||||||
for (const auto & opt : optional_parsers) {
|
|
||||||
any_opt |= opt;
|
|
||||||
}
|
|
||||||
args_seq = args_seq + p.repeat(p.space() + any_opt, 0, -1);
|
|
||||||
}
|
|
||||||
|
|
||||||
common_peg_parser invoke_body = args_seq;
|
|
||||||
auto func_parser = p.tool(
|
|
||||||
p.tool_open(p.literal(INVOKE_START + " name=\"") +
|
|
||||||
p.tool_name(p.literal(name)) + p.literal("\">\n")) +
|
|
||||||
invoke_body + p.space() +
|
|
||||||
p.tool_close(p.literal(INVOKE_END)));
|
|
||||||
|
|
||||||
tool_choice |= p.rule("tool-" + name, func_parser);
|
|
||||||
});
|
|
||||||
|
|
||||||
auto require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
|
||||||
|
|
||||||
common_peg_parser tool_calls = p.eps();
|
|
||||||
if (inputs.parallel_tool_calls) {
|
|
||||||
tool_calls = p.trigger_rule("tool-call",
|
|
||||||
p.literal(FC_START) + p.space() + tool_choice +
|
|
||||||
p.zero_or_more(p.space() + tool_choice) + p.space() + p.literal(FC_END));
|
|
||||||
} else {
|
|
||||||
tool_calls = p.trigger_rule("tool-call",
|
|
||||||
p.literal(FC_START) + p.space() + tool_choice + p.space() + p.literal(FC_END));
|
|
||||||
}
|
|
||||||
|
|
||||||
if (!require_tools) {
|
|
||||||
tool_calls = p.optional(tool_calls);
|
|
||||||
}
|
|
||||||
|
|
||||||
auto content_before_tools = p.content(p.until(FC_START));
|
|
||||||
return generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
|
||||||
});
|
});
|
||||||
|
|
||||||
data.parser = parser.save();
|
data.parser = parser.save();
|
||||||
|
|
||||||
if (include_grammar) {
|
if (include_grammar) {
|
||||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
data.grammar_lazy = has_tools && !require_tools;
|
||||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||||
foreach_function(inputs.tools, [&](const json & tool) {
|
foreach_function(inputs.tools, [&](const json & tool) {
|
||||||
const auto & function = tool.at("function");
|
const auto & function = tool.at("function");
|
||||||
|
|||||||
@@ -4227,6 +4227,20 @@ static void test_template_output_peg_parsers(bool detailed_debug) {
|
|||||||
.expect_reasoning("I'm thinking")
|
.expect_reasoning("I'm thinking")
|
||||||
.expect_content("Hello, world!\nWhat's up?")
|
.expect_content("Hello, world!\nWhat's up?")
|
||||||
.run();
|
.run();
|
||||||
|
|
||||||
|
tst.test(
|
||||||
|
"Let me check the time\n\n"
|
||||||
|
"<|DSML|tool_calls>\n"
|
||||||
|
"<|DSML|invoke name=\"get_time\">\n"
|
||||||
|
"<|DSML|parameter name=\"city\" string=\"true\">Tokyo</|DSML|parameter>\n"
|
||||||
|
"</|DSML|invoke>\n"
|
||||||
|
"</|DSML|tool_calls>") // no </think> after the TC close because the grammar will immediately constrain it to end
|
||||||
|
.enable_thinking(true)
|
||||||
|
.reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK)
|
||||||
|
.tools({ get_time_tool })
|
||||||
|
.expect_reasoning("Let me check the time")
|
||||||
|
.expect_tool_calls({ { "get_time", R"({"city": "Tokyo"})", {} } })
|
||||||
|
.run();
|
||||||
}
|
}
|
||||||
|
|
||||||
// GLM-4.6 tests - format: <tool_call>function_name\n<arg_key>...</arg_key>\n<arg_value>...</arg_value>\n</tool_call>
|
// GLM-4.6 tests - format: <tool_call>function_name\n<arg_key>...</arg_key>\n<arg_value>...</arg_value>\n</tool_call>
|
||||||
|
|||||||
@@ -9,6 +9,7 @@
|
|||||||
#include "peg-parser.h"
|
#include "peg-parser.h"
|
||||||
|
|
||||||
#include <fstream>
|
#include <fstream>
|
||||||
|
#include <iterator>
|
||||||
#include <numeric>
|
#include <numeric>
|
||||||
#include <optional>
|
#include <optional>
|
||||||
#include <sstream>
|
#include <sstream>
|
||||||
@@ -398,7 +399,7 @@ int main(int argc, char ** argv) {
|
|||||||
if (std::optional<common_chat_params> spec_tmpl =
|
if (std::optional<common_chat_params> spec_tmpl =
|
||||||
common_chat_try_specialized_template(chat_template, template_source, params)) {
|
common_chat_try_specialized_template(chat_template, template_source, params)) {
|
||||||
LOG_ERR("\n");
|
LOG_ERR("\n");
|
||||||
LOG_ERR("This template uses a specialized parser, analysis results will not be available.");
|
LOG_ERR("This template uses a specialized parser, analysis results will not be available.\n");
|
||||||
parser_data = *spec_tmpl;
|
parser_data = *spec_tmpl;
|
||||||
} else {
|
} else {
|
||||||
// Render template scenarios if requested
|
// Render template scenarios if requested
|
||||||
@@ -426,7 +427,9 @@ int main(int argc, char ** argv) {
|
|||||||
// Generate Parser
|
// Generate Parser
|
||||||
parser_data = autoparser::peg_generator::generate_parser(chat_template, params, analysis);
|
parser_data = autoparser::peg_generator::generate_parser(chat_template, params, analysis);
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!std::empty(parser_data.parser)) {
|
||||||
LOG_ERR("\n=== Generated Parser ===\n");
|
LOG_ERR("\n=== Generated Parser ===\n");
|
||||||
common_peg_arena arena;
|
common_peg_arena arena;
|
||||||
arena.load(parser_data.parser);
|
arena.load(parser_data.parser);
|
||||||
|
|||||||
Reference in New Issue
Block a user