Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
chat-template.hpp516 linesDownload Raw Back to common
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