Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
main.cpp925 linesDownload Raw Back to main
1#include "arg.h"2#include "common.h"3#include "console.h"4#include "log.h"5#include "sampling.h"6#include "llama.h"7#include "chat-template.hpp"8 9#include <cstdio>10#include <cstring>11#include <ctime>12#include <fstream>13#include <iostream>14#include <sstream>15#include <string>16#include <vector>17 18#if defined (__unix__) || (defined (__APPLE__) && defined (__MACH__))19#include <signal.h>20#include <unistd.h>21#elif defined (_WIN32)22#define WIN32_LEAN_AND_MEAN23#ifndef NOMINMAX24#define NOMINMAX25#endif26#include <windows.h>27#include <signal.h>28#endif29 30#if defined(_MSC_VER)31#pragma warning(disable: 4244 4267) // possible loss of data32#endif33 34static const char * DEFAULT_SYSTEM_MESSAGE = "You are a helpful assistant";35 36static llama_context           ** g_ctx;37static llama_model             ** g_model;38static common_sampler          ** g_smpl;39static common_params            * g_params;40static std::vector<llama_token> * g_input_tokens;41static std::ostringstream       * g_output_ss;42static std::vector<llama_token> * g_output_tokens;43static bool is_interacting  = false;44static bool need_insert_eot = false;45 46static void print_usage(int argc, char ** argv) {47    (void) argc;48 49    LOG("\nexample usage:\n");50    LOG("\n  text generation:     %s -m your_model.gguf -p \"I believe the meaning of life is\" -n 128\n", argv[0]);51    LOG("\n  chat (conversation): %s -m your_model.gguf -p \"You are a helpful assistant\" -cnv\n", argv[0]);52    LOG("\n");53}54 55static bool file_exists(const std::string & path) {56    std::ifstream f(path.c_str());57    return f.good();58}59 60static bool file_is_empty(const std::string & path) {61    std::ifstream f;62    f.exceptions(std::ifstream::failbit | std::ifstream::badbit);63    f.open(path.c_str(), std::ios::in | std::ios::binary | std::ios::ate);64    return f.tellg() == 0;65}66 67#if defined (__unix__) || (defined (__APPLE__) && defined (__MACH__)) || defined (_WIN32)68static void sigint_handler(int signo) {69    if (signo == SIGINT) {70        if (!is_interacting && g_params->interactive) {71            is_interacting  = true;72            need_insert_eot = true;73        } else {74            console::cleanup();75            LOG("\n");76            common_perf_print(*g_ctx, *g_smpl);77 78            // make sure all logs are flushed79            LOG("Interrupted by user\n");80            common_log_pause(common_log_main());81 82            _exit(130);83        }84    }85}86#endif87 88int main(int argc, char ** argv) {89    common_params params;90    g_params = &params;91    if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_MAIN, print_usage)) {92        return 1;93    }94 95    common_init();96 97    auto & sparams = params.sampling;98 99    // save choice to use color for later100    // (note for later: this is a slightly awkward choice)101    console::init(params.simple_io, params.use_color);102    atexit([]() { console::cleanup(); });103 104    if (params.logits_all) {105        LOG_ERR("************\n");106        LOG_ERR("%s: please use the 'perplexity' tool for perplexity calculations\n", __func__);107        LOG_ERR("************\n\n");108 109        return 0;110    }111 112    if (params.embedding) {113        LOG_ERR("************\n");114        LOG_ERR("%s: please use the 'embedding' tool for embedding calculations\n", __func__);115        LOG_ERR("************\n\n");116 117        return 0;118    }119 120    if (params.n_ctx != 0 && params.n_ctx < 8) {121        LOG_WRN("%s: warning: minimum context size is 8, using minimum size.\n", __func__);122        params.n_ctx = 8;123    }124 125    if (params.rope_freq_base != 0.0) {126        LOG_WRN("%s: warning: changing RoPE frequency base to %g.\n", __func__, params.rope_freq_base);127    }128 129    if (params.rope_freq_scale != 0.0) {130        LOG_WRN("%s: warning: scaling RoPE frequency by %g.\n", __func__, params.rope_freq_scale);131    }132 133    LOG_INF("%s: llama backend init\n", __func__);134 135    llama_backend_init();136    llama_numa_init(params.numa);137 138    llama_model * model = nullptr;139    llama_context * ctx = nullptr;140    common_sampler * smpl = nullptr;141 142    g_model = &model;143    g_ctx = &ctx;144    g_smpl = &smpl;145 146    std::vector<common_chat_msg> chat_msgs;147 148    // load the model and apply lora adapter, if any149    LOG_INF("%s: load the model and apply lora adapter, if any\n", __func__);150    common_init_result llama_init = common_init_from_params(params);151 152    model = llama_init.model.get();153    ctx = llama_init.context.get();154 155    if (model == NULL) {156        LOG_ERR("%s: error: unable to load model\n", __func__);157        return 1;158    }159 160    const llama_vocab * vocab = llama_model_get_vocab(model);161    auto chat_templates = common_chat_templates_from_model(model, params.chat_template);162 163    LOG_INF("%s: llama threadpool init, n_threads = %d\n", __func__, (int) params.cpuparams.n_threads);164 165    auto * reg = ggml_backend_dev_backend_reg(ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU));166    auto * ggml_threadpool_new_fn = (decltype(ggml_threadpool_new) *) ggml_backend_reg_get_proc_address(reg, "ggml_threadpool_new");167    auto * ggml_threadpool_free_fn = (decltype(ggml_threadpool_free) *) ggml_backend_reg_get_proc_address(reg, "ggml_threadpool_free");168 169    struct ggml_threadpool_params tpp_batch =170            ggml_threadpool_params_from_cpu_params(params.cpuparams_batch);171    struct ggml_threadpool_params tpp =172            ggml_threadpool_params_from_cpu_params(params.cpuparams);173 174    set_process_priority(params.cpuparams.priority);175 176    struct ggml_threadpool * threadpool_batch = NULL;177    if (!ggml_threadpool_params_match(&tpp, &tpp_batch)) {178        threadpool_batch = ggml_threadpool_new_fn(&tpp_batch);179        if (!threadpool_batch) {180            LOG_ERR("%s: batch threadpool create failed : n_threads %d\n", __func__, tpp_batch.n_threads);181            return 1;182        }183 184        // Start the non-batch threadpool in the paused state185        tpp.paused = true;186    }187 188    struct ggml_threadpool * threadpool = ggml_threadpool_new_fn(&tpp);189    if (!threadpool) {190        LOG_ERR("%s: threadpool create failed : n_threads %d\n", __func__, tpp.n_threads);191        return 1;192    }193 194    llama_attach_threadpool(ctx, threadpool, threadpool_batch);195 196    const int n_ctx_train = llama_model_n_ctx_train(model);197    const int n_ctx = llama_n_ctx(ctx);198 199    if (n_ctx > n_ctx_train) {200        LOG_WRN("%s: model was trained on only %d context tokens (%d specified)\n", __func__, n_ctx_train, n_ctx);201    }202 203    // auto enable conversation mode if chat template is available204    const bool has_chat_template = chat_templates.has_explicit_template && chat_templates.template_default;205    if (params.conversation_mode == COMMON_CONVERSATION_MODE_AUTO) {206        if (has_chat_template) {207            LOG_INF("%s: chat template is available, enabling conversation mode (disable it with -no-cnv)\n", __func__);208            params.conversation_mode = COMMON_CONVERSATION_MODE_ENABLED;209        } else {210            params.conversation_mode = COMMON_CONVERSATION_MODE_DISABLED;211        }212    }213 214    // in case user force-activate conversation mode (via -cnv) without proper chat template, we show a warning215    if (params.conversation_mode && !has_chat_template) {216        LOG_WRN("%s: chat template is not available or is not supported. This may cause the model to output suboptimal responses\n", __func__);217    }218 219    // print chat template example in conversation mode220    if (params.conversation_mode) {221        if (params.enable_chat_template) {222            LOG_INF("%s: chat template example:\n%s\n", __func__, common_chat_format_example(*chat_templates.template_default, params.use_jinja).c_str());223        } else {224            LOG_INF("%s: in-suffix/prefix is specified, chat template will be disabled\n", __func__);225        }226    }227 228    // print system information229    {230        LOG_INF("\n");231        LOG_INF("%s\n", common_params_get_system_info(params).c_str());232        LOG_INF("\n");233    }234 235    std::string path_session = params.path_prompt_cache;236    std::vector<llama_token> session_tokens;237 238    if (!path_session.empty()) {239        LOG_INF("%s: attempting to load saved session from '%s'\n", __func__, path_session.c_str());240        if (!file_exists(path_session)) {241            LOG_INF("%s: session file does not exist, will create.\n", __func__);242        } else if (file_is_empty(path_session)) {243            LOG_INF("%s: The session file is empty. A new session will be initialized.\n", __func__);244        } else {245            // The file exists and is not empty246            session_tokens.resize(n_ctx);247            size_t n_token_count_out = 0;248            if (!llama_state_load_file(ctx, path_session.c_str(), session_tokens.data(), session_tokens.capacity(), &n_token_count_out)) {249                LOG_ERR("%s: failed to load session file '%s'\n", __func__, path_session.c_str());250                return 1;251            }252            session_tokens.resize(n_token_count_out);253            LOG_INF("%s: loaded a session with prompt size of %d tokens\n", __func__, (int)session_tokens.size());254        }255    }256 257    const bool add_bos = llama_vocab_get_add_bos(vocab) && !params.use_jinja;258    if (!llama_model_has_encoder(model)) {259        GGML_ASSERT(!llama_vocab_get_add_eos(vocab));260    }261 262    LOG_DBG("n_ctx: %d, add_bos: %d\n", n_ctx, add_bos);263 264    std::vector<llama_token> embd_inp;265 266    auto chat_add_and_format = [&chat_msgs, &chat_templates](const std::string & role, const std::string & content) {267        common_chat_msg new_msg{role, content, {}};268        auto formatted = common_chat_format_single(*chat_templates.template_default, chat_msgs, new_msg, role == "user", g_params->use_jinja);269        chat_msgs.push_back({role, content, {}});270        LOG_DBG("formatted: '%s'\n", formatted.c_str());271        return formatted;272    };273 274    {275        auto prompt = (params.conversation_mode && params.enable_chat_template)276            // format the system prompt in conversation mode (fallback to default if empty)277            ? chat_add_and_format("system", params.prompt.empty() ? DEFAULT_SYSTEM_MESSAGE : params.prompt)278            // otherwise use the prompt as is279            : params.prompt;280        if (params.interactive_first || !params.prompt.empty() || session_tokens.empty()) {281            LOG_DBG("tokenize the prompt\n");282            embd_inp = common_tokenize(ctx, prompt, true, true);283        } else {284            LOG_DBG("use session tokens\n");285            embd_inp = session_tokens;286        }287 288        LOG_DBG("prompt: \"%s\"\n", prompt.c_str());289        LOG_DBG("tokens: %s\n", string_from(ctx, embd_inp).c_str());290    }291 292    // Should not run without any tokens293    if (embd_inp.empty()) {294        if (add_bos) {295            embd_inp.push_back(llama_vocab_bos(vocab));296            LOG_WRN("embd_inp was considered empty and bos was added: %s\n", string_from(ctx, embd_inp).c_str());297        } else {298            LOG_ERR("input is empty\n");299            return -1;300        }301    }302 303    // Tokenize negative prompt304    if ((int) embd_inp.size() > n_ctx - 4) {305        LOG_ERR("%s: prompt is too long (%d tokens, max %d)\n", __func__, (int) embd_inp.size(), n_ctx - 4);306        return 1;307    }308 309    // debug message about similarity of saved session, if applicable310    size_t n_matching_session_tokens = 0;311    if (!session_tokens.empty()) {312        for (llama_token id : session_tokens) {313            if (n_matching_session_tokens >= embd_inp.size() || id != embd_inp[n_matching_session_tokens]) {314                break;315            }316            n_matching_session_tokens++;317        }318        if (params.prompt.empty() && n_matching_session_tokens == embd_inp.size()) {319            LOG_INF("%s: using full prompt from session file\n", __func__);320        } else if (n_matching_session_tokens >= embd_inp.size()) {321            LOG_INF("%s: session file has exact match for prompt!\n", __func__);322        } else if (n_matching_session_tokens < (embd_inp.size() / 2)) {323            LOG_WRN("%s: session file has low similarity to prompt (%zu / %zu tokens); will mostly be reevaluated\n",324                    __func__, n_matching_session_tokens, embd_inp.size());325        } else {326            LOG_INF("%s: session file matches %zu / %zu tokens of prompt\n",327                    __func__, n_matching_session_tokens, embd_inp.size());328        }329 330        // remove any "future" tokens that we might have inherited from the previous session331        llama_kv_cache_seq_rm(ctx, -1, n_matching_session_tokens, -1);332    }333 334    LOG_DBG("recalculate the cached logits (check): embd_inp.size() %zu, n_matching_session_tokens %zu, embd_inp.size() %zu, session_tokens.size() %zu\n",335         embd_inp.size(), n_matching_session_tokens, embd_inp.size(), session_tokens.size());336 337    // if we will use the cache for the full prompt without reaching the end of the cache, force338    // reevaluation of the last token to recalculate the cached logits339    if (!embd_inp.empty() && n_matching_session_tokens == embd_inp.size() && session_tokens.size() > embd_inp.size()) {340        LOG_DBG("recalculate the cached logits (do): session_tokens.resize( %zu )\n", embd_inp.size() - 1);341 342        session_tokens.resize(embd_inp.size() - 1);343    }344 345    // number of tokens to keep when resetting context346    if (params.n_keep < 0 || params.n_keep > (int) embd_inp.size()) {347        params.n_keep = (int)embd_inp.size();348    } else {349        params.n_keep += add_bos; // always keep the BOS token350    }351 352    if (params.conversation_mode) {353        params.interactive_first = true;354    }355 356    // enable interactive mode if interactive start is specified357    if (params.interactive_first) {358        params.interactive = true;359    }360 361    if (params.verbose_prompt) {362        LOG_INF("%s: prompt: '%s'\n", __func__, params.prompt.c_str());363        LOG_INF("%s: number of tokens in prompt = %zu\n", __func__, embd_inp.size());364        for (int i = 0; i < (int) embd_inp.size(); i++) {365            LOG_INF("%6d -> '%s'\n", embd_inp[i], common_token_to_piece(ctx, embd_inp[i]).c_str());366        }367 368        if (params.n_keep > add_bos) {369            LOG_INF("%s: static prompt based on n_keep: '", __func__);370            for (int i = 0; i < params.n_keep; i++) {371                LOG_CNT("%s", common_token_to_piece(ctx, embd_inp[i]).c_str());372            }373            LOG_CNT("'\n");374        }375        LOG_INF("\n");376    }377 378    // ctrl+C handling379    {380#if defined (__unix__) || (defined (__APPLE__) && defined (__MACH__))381        struct sigaction sigint_action;382        sigint_action.sa_handler = sigint_handler;383        sigemptyset (&sigint_action.sa_mask);384        sigint_action.sa_flags = 0;385        sigaction(SIGINT, &sigint_action, NULL);386#elif defined (_WIN32)387        auto console_ctrl_handler = +[](DWORD ctrl_type) -> BOOL {388            return (ctrl_type == CTRL_C_EVENT) ? (sigint_handler(SIGINT), true) : false;389        };390        SetConsoleCtrlHandler(reinterpret_cast<PHANDLER_ROUTINE>(console_ctrl_handler), true);391#endif392    }393 394    if (params.interactive) {395        LOG_INF("%s: interactive mode on.\n", __func__);396 397        if (!params.antiprompt.empty()) {398            for (const auto & antiprompt : params.antiprompt) {399                LOG_INF("Reverse prompt: '%s'\n", antiprompt.c_str());400                if (params.verbose_prompt) {401                    auto tmp = common_tokenize(ctx, antiprompt, false, true);402                    for (int i = 0; i < (int) tmp.size(); i++) {403                        LOG_INF("%6d -> '%s'\n", tmp[i], common_token_to_piece(ctx, tmp[i]).c_str());404                    }405                }406            }407        }408 409        if (params.input_prefix_bos) {410            LOG_INF("Input prefix with BOS\n");411        }412 413        if (!params.input_prefix.empty()) {414            LOG_INF("Input prefix: '%s'\n", params.input_prefix.c_str());415            if (params.verbose_prompt) {416                auto tmp = common_tokenize(ctx, params.input_prefix, true, true);417                for (int i = 0; i < (int) tmp.size(); i++) {418                    LOG_INF("%6d -> '%s'\n", tmp[i], common_token_to_piece(ctx, tmp[i]).c_str());419                }420            }421        }422 423        if (!params.input_suffix.empty()) {424            LOG_INF("Input suffix: '%s'\n", params.input_suffix.c_str());425            if (params.verbose_prompt) {426                auto tmp = common_tokenize(ctx, params.input_suffix, false, true);427                for (int i = 0; i < (int) tmp.size(); i++) {428                    LOG_INF("%6d -> '%s'\n", tmp[i], common_token_to_piece(ctx, tmp[i]).c_str());429                }430            }431        }432    }433 434    smpl = common_sampler_init(model, sparams);435    if (!smpl) {436        LOG_ERR("%s: failed to initialize sampling subsystem\n", __func__);437        return 1;438    }439 440    LOG_INF("sampler seed: %u\n",     common_sampler_get_seed(smpl));441    LOG_INF("sampler params: \n%s\n", sparams.print().c_str());442    LOG_INF("sampler chain: %s\n",    common_sampler_print(smpl).c_str());443 444    LOG_INF("generate: n_ctx = %d, n_batch = %d, n_predict = %d, n_keep = %d\n", n_ctx, params.n_batch, params.n_predict, params.n_keep);445 446    // group-attention state447    // number of grouped KV tokens so far (used only if params.grp_attn_n > 1)448    int ga_i = 0;449 450    const int ga_n = params.grp_attn_n;451    const int ga_w = params.grp_attn_w;452 453    if (ga_n != 1) {454        GGML_ASSERT(ga_n > 0                    && "grp_attn_n must be positive");                     // NOLINT455        GGML_ASSERT(ga_w % ga_n == 0            && "grp_attn_w must be a multiple of grp_attn_n");     // NOLINT456      //GGML_ASSERT(n_ctx_train % ga_w == 0     && "n_ctx_train must be a multiple of grp_attn_w");    // NOLINT457      //GGML_ASSERT(n_ctx >= n_ctx_train * ga_n && "n_ctx must be at least n_ctx_train * grp_attn_n"); // NOLINT458        LOG_INF("self-extend: n_ctx_train = %d, grp_attn_n = %d, grp_attn_w = %d\n", n_ctx_train, ga_n, ga_w);459    }460    LOG_INF("\n");461 462    if (params.interactive) {463        const char * control_message;464        if (params.multiline_input) {465            control_message = " - To return control to the AI, end your input with '\\'.\n"466                              " - To return control without starting a new line, end your input with '/'.\n";467        } else {468            control_message = " - Press Return to return control to the AI.\n"469                              " - To return control without starting a new line, end your input with '/'.\n"470                              " - If you want to submit another line, end your input with '\\'.\n";471        }472        LOG_INF("== Running in interactive mode. ==\n");473#if defined (__unix__) || (defined (__APPLE__) && defined (__MACH__)) || defined (_WIN32)474        LOG_INF(       " - Press Ctrl+C to interject at any time.\n");475#endif476        LOG_INF(       "%s", control_message);477        if (params.conversation_mode && params.enable_chat_template && params.prompt.empty()) {478            LOG_INF(   " - Using default system message. To change it, set a different value via -p PROMPT or -f FILE argument.\n");479        }480        LOG_INF("\n");481 482        is_interacting = params.interactive_first;483    }484 485    bool is_antiprompt        = false;486    bool input_echo           = true;487    bool display              = true;488    bool need_to_save_session = !path_session.empty() && n_matching_session_tokens < embd_inp.size();489 490    int n_past             = 0;491    int n_remain           = params.n_predict;492    int n_consumed         = 0;493    int n_session_consumed = 0;494 495    std::vector<int>   input_tokens;  g_input_tokens  = &input_tokens;496    std::vector<int>   output_tokens; g_output_tokens = &output_tokens;497    std::ostringstream output_ss;     g_output_ss     = &output_ss;498    std::ostringstream assistant_ss; // for storing current assistant message, used in conversation mode499 500    // the first thing we will do is to output the prompt, so set color accordingly501    console::set_display(console::prompt);502    display = params.display_prompt;503 504    std::vector<llama_token> embd;505 506    // single-token antiprompts507    std::vector<llama_token> antiprompt_token;508 509    for (const std::string & antiprompt : params.antiprompt) {510        auto ids = ::common_tokenize(ctx, antiprompt, false, true);511        if (ids.size() == 1) {512            antiprompt_token.push_back(ids[0]);513        }514    }515 516    if (llama_model_has_encoder(model)) {517        int enc_input_size = embd_inp.size();518        llama_token * enc_input_buf = embd_inp.data();519 520        if (llama_encode(ctx, llama_batch_get_one(enc_input_buf, enc_input_size))) {521            LOG_ERR("%s : failed to eval\n", __func__);522            return 1;523        }524 525        llama_token decoder_start_token_id = llama_model_decoder_start_token(model);526        if (decoder_start_token_id == LLAMA_TOKEN_NULL) {527            decoder_start_token_id = llama_vocab_bos(vocab);528        }529 530        embd_inp.clear();531        embd_inp.push_back(decoder_start_token_id);532    }533 534    while ((n_remain != 0 && !is_antiprompt) || params.interactive) {535        // predict536        if (!embd.empty()) {537            // Note: (n_ctx - 4) here is to match the logic for commandline prompt handling via538            // --prompt or --file which uses the same value.539            int max_embd_size = n_ctx - 4;540 541            // Ensure the input doesn't exceed the context size by truncating embd if necessary.542            if ((int) embd.size() > max_embd_size) {543                const int skipped_tokens = (int) embd.size() - max_embd_size;544                embd.resize(max_embd_size);545 546                console::set_display(console::error);547                LOG_WRN("<<input too long: skipped %d token%s>>", skipped_tokens, skipped_tokens != 1 ? "s" : "");548                console::set_display(console::reset);549            }550 551            if (ga_n == 1) {552                // infinite text generation via context shifting553                // if we run out of context:554                // - take the n_keep first tokens from the original prompt (via n_past)555                // - take half of the last (n_ctx - n_keep) tokens and recompute the logits in batches556 557                if (n_past + (int) embd.size() >= n_ctx) {558                    if (!params.ctx_shift){559                        LOG_DBG("\n\n%s: context full and context shift is disabled => stopping\n", __func__);560                        break;561                    }562 563                    if (params.n_predict == -2) {564                        LOG_DBG("\n\n%s: context full and n_predict == -%d => stopping\n", __func__, params.n_predict);565                        break;566                    }567 568                    const int n_left    = n_past - params.n_keep;569                    const int n_discard = n_left/2;570 571                    LOG_DBG("context full, swapping: n_past = %d, n_left = %d, n_ctx = %d, n_keep = %d, n_discard = %d\n",572                            n_past, n_left, n_ctx, params.n_keep, n_discard);573 574                    llama_kv_cache_seq_rm (ctx, 0, params.n_keep            , params.n_keep + n_discard);575                    llama_kv_cache_seq_add(ctx, 0, params.n_keep + n_discard, n_past, -n_discard);576 577                    n_past -= n_discard;578 579                    LOG_DBG("after swap: n_past = %d\n", n_past);580 581                    LOG_DBG("embd: %s\n", string_from(ctx, embd).c_str());582 583                    LOG_DBG("clear session path\n");584                    path_session.clear();585                }586            } else {587                // context extension via Self-Extend588                while (n_past >= ga_i + ga_w) {589                    const int ib = (ga_n*ga_i)/ga_w;590                    const int bd = (ga_w/ga_n)*(ga_n - 1);591                    const int dd = (ga_w/ga_n) - ib*bd - ga_w;592 593                    LOG_DBG("\n");594                    LOG_DBG("shift: [%6d, %6d] + %6d -> [%6d, %6d]\n", ga_i, n_past, ib*bd, ga_i + ib*bd, n_past + ib*bd);595                    LOG_DBG("div:   [%6d, %6d] / %6d -> [%6d, %6d]\n", ga_i + ib*bd, ga_i + ib*bd + ga_w, ga_n, (ga_i + ib*bd)/ga_n, (ga_i + ib*bd + ga_w)/ga_n);596                    LOG_DBG("shift: [%6d, %6d] + %6d -> [%6d, %6d]\n", ga_i + ib*bd + ga_w, n_past + ib*bd, dd, ga_i + ib*bd + ga_w + dd, n_past + ib*bd + dd);597 598                    llama_kv_cache_seq_add(ctx, 0, ga_i,                n_past,              ib*bd);599                    llama_kv_cache_seq_div(ctx, 0, ga_i + ib*bd,        ga_i + ib*bd + ga_w, ga_n);600                    llama_kv_cache_seq_add(ctx, 0, ga_i + ib*bd + ga_w, n_past + ib*bd,      dd);601 602                    n_past -= bd;603 604                    ga_i += ga_w/ga_n;605 606                    LOG_DBG("\nn_past_old = %d, n_past = %d, ga_i = %d\n\n", n_past + bd, n_past, ga_i);607                }608            }609 610            // try to reuse a matching prefix from the loaded session instead of re-eval (via n_past)611            if (n_session_consumed < (int) session_tokens.size()) {612                size_t i = 0;613                for ( ; i < embd.size(); i++) {614                    if (embd[i] != session_tokens[n_session_consumed]) {615                        session_tokens.resize(n_session_consumed);616                        break;617                    }618 619                    n_past++;620                    n_session_consumed++;621 622                    if (n_session_consumed >= (int) session_tokens.size()) {623                        ++i;624                        break;625                    }626                }627                if (i > 0) {628                    embd.erase(embd.begin(), embd.begin() + i);629                }630            }631 632            for (int i = 0; i < (int) embd.size(); i += params.n_batch) {633                int n_eval = (int) embd.size() - i;634                if (n_eval > params.n_batch) {635                    n_eval = params.n_batch;636                }637 638                LOG_DBG("eval: %s\n", string_from(ctx, embd).c_str());639 640                if (llama_decode(ctx, llama_batch_get_one(&embd[i], n_eval))) {641                    LOG_ERR("%s : failed to eval\n", __func__);642                    return 1;643                }644 645                n_past += n_eval;646 647                LOG_DBG("n_past = %d\n", n_past);648                // Display total tokens alongside total time649                if (params.n_print > 0 && n_past % params.n_print == 0) {650                    LOG_DBG("\n\033[31mTokens consumed so far = %d / %d \033[0m\n", n_past, n_ctx);651                }652            }653 654            if (!embd.empty() && !path_session.empty()) {655                session_tokens.insert(session_tokens.end(), embd.begin(), embd.end());656                n_session_consumed = session_tokens.size();657            }658        }659 660        embd.clear();661 662        if ((int) embd_inp.size() <= n_consumed && !is_interacting) {663            // optionally save the session on first sample (for faster prompt loading next time)664            if (!path_session.empty() && need_to_save_session && !params.prompt_cache_ro) {665                need_to_save_session = false;666                llama_state_save_file(ctx, path_session.c_str(), session_tokens.data(), session_tokens.size());667 668                LOG_DBG("saved session to %s\n", path_session.c_str());669            }670 671            const llama_token id = common_sampler_sample(smpl, ctx, -1);672 673            common_sampler_accept(smpl, id, /* accept_grammar= */ true);674 675            // LOG_DBG("last: %s\n", string_from(ctx, smpl->prev.to_vector()).c_str());676 677            embd.push_back(id);678 679            // echo this to console680            input_echo = true;681 682            // decrement remaining sampling budget683            --n_remain;684 685            LOG_DBG("n_remain: %d\n", n_remain);686        } else {687            // some user input remains from prompt or interaction, forward it to processing688            LOG_DBG("embd_inp.size(): %d, n_consumed: %d\n", (int) embd_inp.size(), n_consumed);689            while ((int) embd_inp.size() > n_consumed) {690                embd.push_back(embd_inp[n_consumed]);691 692                // push the prompt in the sampling context in order to apply repetition penalties later693                // for the prompt, we don't apply grammar rules694                common_sampler_accept(smpl, embd_inp[n_consumed], /* accept_grammar= */ false);695 696                ++n_consumed;697                if ((int) embd.size() >= params.n_batch) {698                    break;699                }700            }701        }702 703        // display text704        if (input_echo && display) {705            for (auto id : embd) {706                const std::string token_str = common_token_to_piece(ctx, id, params.special);707 708                // Console/Stream Output709                LOG("%s", token_str.c_str());710 711                // Record Displayed Tokens To Log712                // Note: Generated tokens are created one by one hence this check713                if (embd.size() > 1) {714                    // Incoming Requested Tokens715                    input_tokens.push_back(id);716                } else {717                    // Outgoing Generated Tokens718                    output_tokens.push_back(id);719                    output_ss << token_str;720                }721            }722        }723 724        // reset color to default if there is no pending user input725        if (input_echo && (int) embd_inp.size() == n_consumed) {726            console::set_display(console::reset);727            display = true;728        }729 730        // if not currently processing queued inputs;731        if ((int) embd_inp.size() <= n_consumed) {732            // check for reverse prompt in the last n_prev tokens733            if (!params.antiprompt.empty()) {734                const int n_prev = 32;735                const std::string last_output = common_sampler_prev_str(smpl, ctx, n_prev);736 737                is_antiprompt = false;738                // Check if each of the reverse prompts appears at the end of the output.739                // If we're not running interactively, the reverse prompt might be tokenized with some following characters740                // so we'll compensate for that by widening the search window a bit.741                for (std::string & antiprompt : params.antiprompt) {742                    size_t extra_padding = params.interactive ? 0 : 2;743                    size_t search_start_pos = last_output.length() > static_cast<size_t>(antiprompt.length() + extra_padding)744                        ? last_output.length() - static_cast<size_t>(antiprompt.length() + extra_padding)745                        : 0;746 747                    if (last_output.find(antiprompt, search_start_pos) != std::string::npos) {748                        if (params.interactive) {749                            is_interacting = true;750                        }751                        is_antiprompt = true;752                        break;753                    }754                }755 756                // check for reverse prompt using special tokens757                llama_token last_token = common_sampler_last(smpl);758                if (std::find(antiprompt_token.begin(), antiprompt_token.end(), last_token) != antiprompt_token.end()) {759                    if (params.interactive) {760                        is_interacting = true;761                    }762                    is_antiprompt = true;763                }764 765                if (is_antiprompt) {766                    LOG_DBG("found antiprompt: %s\n", last_output.c_str());767                }768            }769 770            // deal with end of generation tokens in interactive mode771            if (llama_vocab_is_eog(vocab, common_sampler_last(smpl))) {772                LOG_DBG("found an EOG token\n");773 774                if (params.interactive) {775                    if (!params.antiprompt.empty()) {776                        // tokenize and inject first reverse prompt777                        const auto first_antiprompt = common_tokenize(ctx, params.antiprompt.front(), false, true);778                        embd_inp.insert(embd_inp.end(), first_antiprompt.begin(), first_antiprompt.end());779                        is_antiprompt = true;780                    }781 782                    if (params.enable_chat_template) {783                        chat_add_and_format("assistant", assistant_ss.str());784                    }785                    is_interacting = true;786                    LOG("\n");787                }788            }789 790            // if current token is not EOG, we add it to current assistant message791            if (params.conversation_mode) {792                const auto id = common_sampler_last(smpl);793                assistant_ss << common_token_to_piece(ctx, id, false);794            }795 796            if (n_past > 0 && is_interacting) {797                LOG_DBG("waiting for user input\n");798 799                if (params.conversation_mode) {800                    LOG("\n> ");801                }802 803                if (params.input_prefix_bos) {804                    LOG_DBG("adding input prefix BOS token\n");805                    embd_inp.push_back(llama_vocab_bos(vocab));806                }807 808                std::string buffer;809                if (!params.input_prefix.empty() && !params.conversation_mode) {810                    LOG_DBG("appending input prefix: '%s'\n", params.input_prefix.c_str());811                    LOG("%s", params.input_prefix.c_str());812                }813 814                // color user input only815                console::set_display(console::user_input);816                display = params.display_prompt;817 818                std::string line;819                bool another_line = true;820                do {821                    another_line = console::readline(line, params.multiline_input);822                    buffer += line;823                } while (another_line);824 825                // done taking input, reset color826                console::set_display(console::reset);827                display = true;828 829                // Add tokens to embd only if the input buffer is non-empty830                // Entering a empty line lets the user pass control back831                if (buffer.length() > 1) {832                    // append input suffix if any833                    if (!params.input_suffix.empty() && !params.conversation_mode) {834                        LOG_DBG("appending input suffix: '%s'\n", params.input_suffix.c_str());835                        LOG("%s", params.input_suffix.c_str());836                    }837 838                    LOG_DBG("buffer: '%s'\n", buffer.c_str());839 840                    const size_t original_size = embd_inp.size();841 842                    if (params.escape) {843                        string_process_escapes(buffer);844                    }845 846                    bool format_chat = params.conversation_mode && params.enable_chat_template;847                    std::string user_inp = format_chat848                        ? chat_add_and_format("user", std::move(buffer))849                        : std::move(buffer);850                    // TODO: one inconvenient of current chat template implementation is that we can't distinguish between user input and special tokens (prefix/postfix)851                    const auto line_pfx = common_tokenize(ctx, params.input_prefix, false, true);852                    const auto line_inp = common_tokenize(ctx, user_inp,            false, format_chat);853                    const auto line_sfx = common_tokenize(ctx, params.input_suffix, false, true);854 855                    LOG_DBG("input tokens: %s\n", string_from(ctx, line_inp).c_str());856 857                    // if user stop generation mid-way, we must add EOT to finish model's last response858                    if (need_insert_eot && format_chat) {859                        llama_token eot = llama_vocab_eot(vocab);860                        embd_inp.push_back(eot == LLAMA_TOKEN_NULL ? llama_vocab_eos(vocab) : eot);861                        need_insert_eot = false;862                    }863 864                    embd_inp.insert(embd_inp.end(), line_pfx.begin(), line_pfx.end());865                    embd_inp.insert(embd_inp.end(), line_inp.begin(), line_inp.end());866                    embd_inp.insert(embd_inp.end(), line_sfx.begin(), line_sfx.end());867 868                    for (size_t i = original_size; i < embd_inp.size(); ++i) {869                        const llama_token token = embd_inp[i];870                        output_tokens.push_back(token);871                        output_ss << common_token_to_piece(ctx, token);872                    }873 874                    // reset assistant message875                    assistant_ss.str("");876 877                    n_remain -= line_inp.size();878                    LOG_DBG("n_remain: %d\n", n_remain);879                } else {880                    LOG_DBG("empty line, passing control back\n");881                }882 883                input_echo = false; // do not echo this again884            }885 886            if (n_past > 0) {887                if (is_interacting) {888                    common_sampler_reset(smpl);889                }890                is_interacting = false;891            }892        }893 894        // end of generation895        if (!embd.empty() && llama_vocab_is_eog(vocab, embd.back()) && !(params.interactive)) {896            LOG(" [end of text]\n");897            break;898        }899 900        // In interactive mode, respect the maximum number of tokens and drop back to user input when reached.901        // We skip this logic when n_predict == -1 (infinite) or -2 (stop at context size).902        if (params.interactive && n_remain <= 0 && params.n_predict >= 0) {903            n_remain = params.n_predict;904            is_interacting = true;905        }906    }907 908    if (!path_session.empty() && params.prompt_cache_all && !params.prompt_cache_ro) {909        LOG("\n%s: saving final output to session file '%s'\n", __func__, path_session.c_str());910        llama_state_save_file(ctx, path_session.c_str(), session_tokens.data(), session_tokens.size());911    }912 913    LOG("\n\n");914    common_perf_print(ctx, smpl);915 916    common_sampler_free(smpl);917 918    llama_backend_free();919 920    ggml_threadpool_free_fn(threadpool);921    ggml_threadpool_free_fn(threadpool_batch);922 923    return 0;924}925