KBaba7/llama.cpp
0
1/*2 Copyright 2024 Google LLC3 4 Use of this source code is governed by an MIT-style5 license that can be found in the LICENSE file or at6 https://opensource.org/licenses/MIT.7*/8// SPDX-License-Identifier: MIT9#pragma once10 11#include "minja.hpp"12#include <json.hpp>13#include <string>14#include <vector>15 16using json = nlohmann::ordered_json;17 18namespace minja {19 20struct chat_template_caps {21 bool supports_tools = false;22 bool supports_tool_calls = false;23 bool supports_tool_responses = false;24 bool supports_system_role = false;25 bool supports_parallel_tool_calls = false;26 bool supports_tool_call_id = false;27 // meta-llama/Llama-3.1-8B-Instruct expects arguments to be an object.28 // Most other templates (and OpenAI's API) expect the arguments object to be stringified.29 bool requires_object_arguments = false;30 // CohereForAI/c4ai-command-r-plus simple variant31 bool requires_non_null_content = false;32 // MiniMaxAI/MiniMax-Text-01 special33 bool requires_typed_content = false;34};35 36struct chat_template_inputs {37 nlohmann::ordered_json messages;38 nlohmann::ordered_json tools;39 bool add_generation_prompt = true;40 nlohmann::ordered_json extra_context;41 std::chrono::system_clock::time_point now = std::chrono::system_clock::now();42};43 44struct chat_template_options {45 bool apply_polyfills = true;46 bool use_bos_token = true;47 bool use_eos_token = true;48 bool define_strftime_now = true;49 50 bool polyfill_tools = true;51 bool polyfill_tool_call_examples = true;52 bool polyfill_tool_calls = true;53 bool polyfill_tool_responses = true;54 bool polyfill_system_role = true;55 bool polyfill_object_arguments = true;56 bool polyfill_typed_content = true;57};58 59class chat_template {60 61 private:62 chat_template_caps caps_;63 std::string source_;64 std::string bos_token_;65 std::string eos_token_;66 std::shared_ptr<minja::TemplateNode> template_root_;67 std::string tool_call_example_;68 69 std::string try_raw_render(70 const nlohmann::ordered_json & messages,71 const nlohmann::ordered_json & tools,72 bool add_generation_prompt,73 const nlohmann::ordered_json & extra_context = nlohmann::ordered_json()) const74 {75 try {76 chat_template_inputs inputs;77 inputs.messages = messages;78 inputs.tools = tools;79 inputs.add_generation_prompt = add_generation_prompt;80 inputs.extra_context = extra_context;81 // Use fixed date for tests82 inputs.now = std::chrono::system_clock::from_time_t(0);83 84 chat_template_options opts;85 opts.apply_polyfills = false;86 87 auto prompt = apply(inputs, opts);88 // fprintf(stderr, "try_raw_render: %s\n", prompt.c_str());89 return prompt;90 } catch (const std::exception & e) {91 // fprintf(stderr, "try_raw_render error: %s\n", e.what());92 return "";93 }94 }95 96 public:97 98 chat_template(const std::string & source, const std::string & bos_token, const std::string & eos_token)99 : source_(source), bos_token_(bos_token), eos_token_(eos_token)100 {101 template_root_ = minja::Parser::parse(source_, {102 /* .trim_blocks = */ true,103 /* .lstrip_blocks = */ true,104 /* .keep_trailing_newline = */ false,105 });106 107 auto contains = [](const std::string & haystack, const std::string & needle) {108 return haystack.find(needle) != std::string::npos;109 };110 111 const std::string user_needle = "<User Needle>";112 const std::string sys_needle = "<System Needle>";113 const json dummy_str_user_msg = {{"role", "user"}, {"content", user_needle}};114 const json dummy_typed_user_msg = {{"role", "user"}, {"content", json::array({{{"type", "text"}, {"text", user_needle}}})}};115 116 caps_.requires_typed_content =117 !contains(try_raw_render(json::array({dummy_str_user_msg}), {}, false), user_needle)118 && contains(try_raw_render(json::array({dummy_typed_user_msg}), {}, false), user_needle);119 120 const auto dummy_user_msg = caps_.requires_typed_content121 ? dummy_typed_user_msg122 : dummy_str_user_msg;123 const json needle_system_msg = {124 {"role", "system"},125 {"content", caps_.requires_typed_content ? json::array({{{"type", "text"}, {"text", sys_needle}}}) : json(sys_needle)},126 };127 128 caps_.supports_system_role = contains(try_raw_render({needle_system_msg, dummy_user_msg,}, {}, false), sys_needle);129 130 auto out = try_raw_render(json::array({131 dummy_user_msg132 }), json::array({133 {134 {"name", "some_tool"},135 {"type", "function"},136 {"function", {137 {"name", "some_tool"},138 {"description", "Some tool."},139 {"parameters", {140 {"type", "object"},141 {"properties", {142 {"arg", {143 {"type", "string"},144 {"description", "Some argument."},145 }},146 }},147 {"required", json::array({ "arg" })},148 }},149 }},150 },151 }), false);152 caps_.supports_tools = contains(out, "some_tool");153 154 auto make_tool_calls_msg = [&](const json & tool_calls) {155 return json {156 {"role", "assistant"},157 {"content", nullptr},158 {"tool_calls", tool_calls},159 };160 };161 auto make_tool_call = [](const std::string & tool_name, const json & arguments) {162 return json {163 {"id", "call_1___"},164 {"type", "function"},165 {"function", {166 {"arguments", arguments},167 {"name", tool_name},168 }},169 };170 };171 const json dummy_args_obj {{"argument_needle", "print('Hello, World!')"}};172 173 // Note: the arguments are rendered in both cases, but may be double-escaped, which we don't want.174 out = try_raw_render(json::array({175 dummy_user_msg,176 make_tool_calls_msg(json::array({make_tool_call("ipython", dummy_args_obj.dump())})),177 }), {}, false);178 auto tool_call_renders_str_arguments = contains(out, "\"argument_needle\":") || contains(out, "'argument_needle':");179 out = try_raw_render(json::array({180 dummy_user_msg,181 make_tool_calls_msg(json::array({make_tool_call("ipython", dummy_args_obj)})),182 }), {}, false);183 auto tool_call_renders_obj_arguments = contains(out, "\"argument_needle\":") || contains(out, "'argument_needle':");184 185 caps_.supports_tool_calls = tool_call_renders_str_arguments || tool_call_renders_obj_arguments;186 caps_.requires_object_arguments = !tool_call_renders_str_arguments && tool_call_renders_obj_arguments;187 auto out_empty = try_raw_render(json::array({dummy_user_msg, {{"role", "assistant"}, {"content", ""}}}), {}, false);188 auto out_null = try_raw_render(json::array({dummy_user_msg, {{"role", "assistant"}, {"content", nullptr}}}), {}, false);189 caps_.requires_non_null_content = contains(out_empty, user_needle) && !contains(out_null, user_needle);190 191 if (caps_.supports_tool_calls) {192 auto dummy_args = caps_.requires_object_arguments ? dummy_args_obj : json(dummy_args_obj.dump());193 auto tc1 = make_tool_call("test_tool1", dummy_args);194 auto tc2 = make_tool_call("test_tool2", dummy_args);195 auto out = try_raw_render(json::array({196 dummy_user_msg,197 make_tool_calls_msg(json::array({tc1, tc2})),198 }), {}, false);199 caps_.supports_parallel_tool_calls = contains(out, "test_tool1") && contains(out, "test_tool2");200 201 out = try_raw_render(json::array({202 dummy_user_msg,203 make_tool_calls_msg(json::array({tc1})),204 {205 {"role", "tool"},206 {"name", "test_tool1"},207 {"content", "Some response!"},208 {"tool_call_id", "call_911_"},209 }210 }), {}, false);211 caps_.supports_tool_responses = contains(out, "Some response!");212 caps_.supports_tool_call_id = contains(out, "call_911_");213 }214 215 try {216 if (!caps_.supports_tools) {217 const json user_msg {218 {"role", "user"},219 {"content", "Hey"},220 };221 const json args {222 {"arg1", "some_value"},223 };224 const json tool_call_msg {225 {"role", "assistant"},226 {"content", nullptr},227 {"tool_calls", json::array({228 {229 // TODO: detect if requires numerical id or fixed length == 6 like Nemo230 {"id", "call_1___"},231 {"type", "function"},232 {"function", {233 {"name", "tool_name"},234 {"arguments", (caps_.requires_object_arguments ? args : json(minja::Value(args).dump(-1, /* to_json= */ true)))},235 }},236 },237 })},238 };239 std::string prefix, full;240 {241 chat_template_inputs inputs;242 inputs.messages = json::array({user_msg});243 inputs.add_generation_prompt = true;244 prefix = apply(inputs);245 }246 {247 chat_template_inputs inputs;248 inputs.messages = json::array({user_msg, tool_call_msg});249 inputs.add_generation_prompt = false;250 full = apply(inputs);251 }252 253 if (full.find(prefix) != 0) {254 if (prefix.rfind(eos_token_) == prefix.size() - eos_token_.size()) {255 prefix = prefix.substr(0, prefix.size() - eos_token_.size());256 }257 }258 if (full.find(prefix) != 0) {259 fprintf(stderr, "Failed to infer a tool call example (possible template bug)\n");260 }261 tool_call_example_ = full.substr(prefix.size());262 }263 } catch (const std::exception & e) {264 fprintf(stderr, "Failed to generate tool call example: %s\n", e.what());265 }266 }267 268 const std::string & source() const { return source_; }269 const std::string & bos_token() const { return bos_token_; }270 const std::string & eos_token() const { return eos_token_; }271 const chat_template_caps & original_caps() const { return caps_; }272 273 // Deprecated, please use the form with chat_template_inputs and chat_template_options274 std::string apply(275 const nlohmann::ordered_json & messages,276 const nlohmann::ordered_json & tools,277 bool add_generation_prompt,278 const nlohmann::ordered_json & extra_context = nlohmann::ordered_json(),279 bool apply_polyfills = true)280 {281 fprintf(stderr, "[%s] Deprecated!\n", __func__);282 chat_template_inputs inputs;283 inputs.messages = messages;284 inputs.tools = tools;285 inputs.add_generation_prompt = add_generation_prompt;286 inputs.extra_context = extra_context;287 inputs.now = std::chrono::system_clock::now();288 289 chat_template_options opts;290 opts.apply_polyfills = apply_polyfills;291 292 return apply(inputs, opts);293 }294 295 std::string apply(296 const chat_template_inputs & inputs,297 const chat_template_options & opts = chat_template_options()) const298 {299 json actual_messages;300 301 auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();302 auto has_tool_calls = false;303 auto has_tool_responses = false;304 auto has_string_content = false;305 for (const auto & message : inputs.messages) {306 if (message.contains("tool_calls") && !message["tool_calls"].is_null()) {307 has_tool_calls = true;308 }309 if (message.contains("role") && message["role"] == "tool") {310 has_tool_responses = true;311 }312 if (message.contains("content") && message["content"].is_string()) {313 has_string_content = true;314 }315 }316 317 auto polyfill_system_role = opts.polyfill_system_role && !caps_.supports_system_role;318 auto polyfill_tools = opts.polyfill_tools && has_tools && !caps_.supports_tools;319 auto polyfill_tool_call_example = polyfill_tools && opts.polyfill_tool_call_examples;320 auto polyfill_tool_calls = opts.polyfill_tool_calls && has_tool_calls && !caps_.supports_tool_calls;321 auto polyfill_tool_responses = opts.polyfill_tool_responses && has_tool_responses && !caps_.supports_tool_responses;322 auto polyfill_object_arguments = opts.polyfill_object_arguments && has_tool_calls && caps_.requires_object_arguments;323 auto polyfill_typed_content = opts.polyfill_typed_content && has_string_content && caps_.requires_typed_content;324 325 auto needs_polyfills = opts.apply_polyfills && (false326 || polyfill_system_role327 || polyfill_tools328 || polyfill_tool_calls329 || polyfill_tool_responses330 || polyfill_object_arguments331 || polyfill_typed_content332 );333 334 if (needs_polyfills) {335 actual_messages = json::array();336 337 auto add_message = [&](const json & msg) {338 if (polyfill_typed_content && msg.contains("content") && !msg.at("content").is_null() && msg.at("content").is_string()) {339 actual_messages.push_back({340 {"role", msg.at("role")},341 {"content", {{342 {"type", "text"},343 {"text", msg.at("content")},344 }}},345 });346 } else {347 actual_messages.push_back(msg);348 }349 };350 351 std::string pending_system;352 auto flush_sys = [&]() {353 if (!pending_system.empty()) {354 add_message({355 {"role", "user"},356 {"content", pending_system},357 });358 pending_system.clear();359 }360 };361 362 json adjusted_messages;363 if (polyfill_tools) {364 adjusted_messages = add_system(inputs.messages,365 "You can call any of the following tools to satisfy the user's requests: " + minja::Value(inputs.tools).dump(2, /* to_json= */ true) +366 (!polyfill_tool_call_example || tool_call_example_.empty() ? "" : "\n\nExample tool call syntax:\n\n" + tool_call_example_));367 } else {368 adjusted_messages = inputs.messages;369 }370 371 for (const auto & message_ : adjusted_messages) {372 auto message = message_;373 if (!message.contains("role") || !message.contains("content")) {374 throw std::runtime_error("message must have 'role' and 'content' fields: " + message.dump());375 }376 std::string role = message.at("role");377 378 if (message.contains("tool_calls")) {379 if (polyfill_object_arguments || polyfill_tool_calls) {380 for (auto & tool_call : message.at("tool_calls")) {381 if (tool_call["type"] == "function") {382 auto & function = tool_call.at("function");383 auto & arguments = function.at("arguments");384 if (arguments.is_string()) {385 try {386 arguments = json::parse(arguments.get<std::string>());387 } catch (const std::exception & ecvt) {388 fprintf(stderr, "Failed to parse arguments: %s\n", ecvt.what());389 }390 }391 }392 }393 }394 if (polyfill_tool_calls) {395 auto content = message.at("content");396 auto tool_calls = json::array();397 for (const auto & tool_call : message.at("tool_calls")) {398 if (tool_call.at("type") != "function") {399 continue;400 }401 const auto & function = tool_call.at("function");402 auto tc = json {403 {"name", function.at("name")},404 {"arguments", function.at("arguments")},405 };406 if (tool_call.contains("id")) {407 tc["id"] = tool_call["id"];408 }409 tool_calls.push_back(tc);410 }411 auto obj = json {412 {"tool_calls", tool_calls},413 };414 if (!content.is_null() && content != "") {415 obj["content"] = content;416 }417 message["content"] = obj.dump(2);418 message.erase("tool_calls");419 }420 }421 if (polyfill_tool_responses && role == "tool") {422 message["role"] = "user";423 auto obj = json {424 {"tool_response", {425 {"content", message.at("content")},426 }},427 };428 if (message.contains("name")) {429 obj["tool_response"]["name"] = message.at("name");430 }431 if (message.contains("tool_call_id")) {432 obj["tool_response"]["tool_call_id"] = message.at("tool_call_id");433 }434 message["content"] = obj.dump(2);435 message.erase("name");436 }437 438 if (!message["content"].is_null() && polyfill_system_role) {439 std::string content = message.at("content");440 if (role == "system") {441 if (!pending_system.empty()) pending_system += "\n";442 pending_system += content;443 continue;444 } else {445 if (role == "user") {446 if (!pending_system.empty()) {447 message["content"] = pending_system + (content.empty() ? "" : "\n" + content);448 pending_system.clear();449 }450 } else {451 flush_sys();452 }453 }454 }455 add_message(message);456 }457 flush_sys();458 } else {459 actual_messages = inputs.messages;460 }461 462 auto context = minja::Context::make(json({463 {"messages", actual_messages},464 {"add_generation_prompt", inputs.add_generation_prompt},465 }));466 context->set("bos_token", opts.use_bos_token ? bos_token_ : "");467 context->set("eos_token", opts.use_eos_token ? eos_token_ : "");468 if (opts.define_strftime_now) {469 auto now = inputs.now;470 context->set("strftime_now", Value::callable([now](const std::shared_ptr<minja::Context> &, minja::ArgumentsValue & args) {471 args.expectArgs("strftime_now", {1, 1}, {0, 0});472 auto format = args.args[0].get<std::string>();473 474 auto time = std::chrono::system_clock::to_time_t(now);475 auto local_time = *std::localtime(&time);476 std::ostringstream ss;477 ss << std::put_time(&local_time, format.c_str());478 return ss.str();479 }));480 }481 if (!inputs.tools.is_null()) {482 context->set("tools", minja::Value(inputs.tools));483 }484 if (!inputs.extra_context.is_null()) {485 for (auto & kv : inputs.extra_context.items()) {486 context->set(kv.key(), minja::Value(kv.value()));487 }488 }489 490 auto ret = template_root_->render(context);491 // fprintf(stderr, "actual_messages: %s\n", actual_messages.dump(2).c_str());492 // fprintf(stderr, "apply: %s\n\n", ret.c_str());493 return ret;494 }495 496 static nlohmann::ordered_json add_system(const nlohmann::ordered_json & messages, const std::string & system_prompt) {497 json messages_with_system = messages;498 499 if (messages_with_system.size() > 0 && messages_with_system[0].at("role") == "system") {500 std::string existing_system = messages_with_system.at(0).at("content");501 messages_with_system[0] = json {502 {"role", "system"},503 {"content", existing_system + "\n\n" + system_prompt},504 };505 } else {506 messages_with_system.insert(messages_with_system.begin(), json {507 {"role", "system"},508 {"content", system_prompt},509 });510 }511 return messages_with_system;512 }513};514 515} // namespace minja516 