Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
lookahead.cpp484 linesDownload Raw Back to lookahead
1#include "arg.h"2#include "common.h"3#include "sampling.h"4#include "log.h"5#include "llama.h"6 7#include <algorithm>8#include <clocale>9#include <cstdio>10#include <string>11#include <vector>12 13struct ngram_data {14    bool active = false;15 16    llama_seq_id seq_id = -1;17 18    std::vector<int> i_batch;19 20    std::vector<llama_token> tokens;21};22 23// n-gram container24struct ngram_container {25    ngram_container(int n_vocab, int N, int G) {26        cnt.resize(n_vocab);27        head.resize(n_vocab);28        tokens.resize(n_vocab * G * (N - 1));29    }30 31    int n_total = 0;32 33    std::vector<int> cnt;34    std::vector<int> head;35 36    // [n_vocab][G][N - 1]37    // for each token of the vocab, keep a ring-buffer of capacity G of n-grams of size N - 138    std::vector<llama_token> tokens;39};40 41int main(int argc, char ** argv) {42    std::setlocale(LC_NUMERIC, "C");43 44    common_params params;45 46    common_init();47 48    if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_COMMON)) {49        return 1;50    }51 52    const int W = 15; // lookahead window53    const int N = 5;  // n-gram size54    const int G = 15; // max verification n-grams55 56    // lookahead requires W + G + 1 sequences for parallel Jacobi decoding57    params.n_parallel = W + G + 1;58 59    // unified KV cache is required for coupled sequences in batch splitting60    params.kv_unified = true;61 62    // init llama.cpp63    llama_backend_init();64    llama_numa_init(params.numa);65 66    // load the target model67    auto llama_init = common_init_from_params(params);68 69    auto * model = llama_init->model();70    auto * ctx   = llama_init->context();71 72    auto * mem = llama_get_memory(ctx);73 74    const llama_vocab * vocab = llama_model_get_vocab(model);75 76    // Tokenize the prompt77    std::vector<llama_token> inp;78    std::vector<llama_token> all;79 80    inp = common_tokenize(ctx, params.prompt, true, true);81    all = inp;82 83    const int max_context_size     = llama_n_ctx(ctx);84    const int max_tokens_list_size = max_context_size - 4;85 86    if ((int) inp.size() > max_tokens_list_size) {87        LOG_ERR("%s: prompt too long (%d tokens, max %d)\n", __func__, (int) inp.size(), max_tokens_list_size);88        return 1;89    }90 91    LOG("\n\n");92 93    for (auto id : inp) {94        LOG("%s", common_token_to_piece(ctx, id).c_str());95    }96 97    fflush(stderr);98 99    const int n_input = inp.size();100 101    const auto t_enc_start = ggml_time_us();102 103    // eval the prompt104    llama_decode(ctx, llama_batch_get_one( inp.data(), n_input - 1));105    llama_decode(ctx, llama_batch_get_one(&inp.back(),           1));106 107    for (int s = 1; s < W + G + 1; ++s) {108        llama_memory_seq_cp(mem, 0, s, -1, -1);109    }110 111    const auto t_enc_end = ggml_time_us();112 113    int n_predict = 0;114    int n_accept  = 0;115 116    int n_past = inp.size();117 118    llama_token id = 0;119 120    // used to determine end of generation121    bool has_eos = false;122 123    // for each decoded batch, we have at most W + G + 1 distinct sequences:124    // seq_id == 0           : the current input token125    // seq_id [1, W]         : tokens from the past N - 1 Jacobi iterations126    // seq_id [W + 1, W + G] : verification n-grams127    llama_batch batch = llama_batch_init(llama_n_ctx(ctx), 0, W + G + 1);128 129    // target model sampling context130    struct common_sampler * smpl = common_sampler_init(model, params.sampling);131 132    // verification n-grams133    std::vector<ngram_data> ngrams_cur(G);134 135    // tokens for the past N - 1 Jacobi iterations136    std::vector<llama_token> tokens_j_prev(W);137    std::vector<std::vector<llama_token>> tokens_j(N - 1);138    for (int j = 0; j < N - 1; j++) {139        tokens_j[j].resize(W);140 141        for (int i = 0; i < W; i++) {142            // there are different ways to init these tokens143            if (0) {144                // initialize randomly from the prompt tokens145                tokens_j[j][i] = all[1 + rand() % (all.size() - 1)];146            } else {147                // initialize with a sequence of increasing numbers148                tokens_j[j][i] = 100 + i;149            }150        }151    }152 153    std::vector<llama_seq_id> seq_id_look;154 155    // the input token belongs both to all sequences156    std::vector<llama_seq_id> seq_id_all(W + G + 1);157    for (int i = 0; i < W + G + 1; i++) {158        seq_id_all[i] = i;159    }160 161    // here we keep adding new n-grams as we go162    ngram_container ngrams_observed(llama_vocab_n_tokens(vocab), N, G);163 164    const auto t_dec_start = ggml_time_us();165 166    // sample first token167    {168        id = common_sampler_sample(smpl, ctx, 0);169 170        common_sampler_accept(smpl, id, true);171 172        {173            const std::string token_str = common_token_to_piece(ctx, id);174 175            LOG("%s", token_str.c_str());176            fflush(stdout);177        }178    }179 180    while (true) {181        // build the mask from https://lmsys.org/blog/2023-11-21-lookahead-decoding/182        //183        // Example for W = 5, N = 4, G = 2:184        // (I = input, L = lookahead, V = verification)185        //186        // Batch:  0  1  2  3  4  5  6  7  8  9 10 11 12 13 14 15 16 17 18 19 20187        // T:        -2 -2 -2 -2 -1 -1 -1 -1 -1  0  0  0  0  0  0188        // Info:   I  L  L  L  L  L  L  L  L  L  L  L  L  L  L  V  V  V  V  V  V189        // Pos:    0  1  2  3  4  1  2  3  4  5  2  3  4  5  6  1  2  3  1  2  3   (+ n_past)190        // Logits: 1  0  0  0  0  0  0  0  0  0  1  1  1  1  1  1  1  1  1  1  1191        // ---------------------------------------------------------------------192        // Seq:    0193        //         1              1              1194        //         2  2              2              2195        //         3  3  3              3              3196        //         4  4  4  4              4              4197        //         5  5  5  5  5              5              5198        //         6                                            6  6  6199        //         7                                                     7  7  7200        // ---------------------------------------------------------------------201        //                                       |  |  |  |  |  |  |  |  |  |  |202        //                                       V  V  V  V  V  |  |  |  |  |  |203        //                                         j_tokens     |  |  |  |  |  |204        //                                                      V  V  V  V  V  V205        //                                                             id206        {207            common_batch_clear(batch);208 209            // current token - first token of the first level210            common_batch_add(batch, id, n_past, seq_id_all, true);211 212            // verification n-grams - queue this before the lookahead tokens for less KV cache fragmentation213            {214                const int g_cur = ngrams_observed.cnt[id];215 216                ngrams_cur.resize(g_cur);217                for (int g = 0; g < g_cur; g++) {218                    ngrams_cur[g].active = true;219                    ngrams_cur[g].tokens.resize(N);220                    ngrams_cur[g].i_batch.resize(N);221                    ngrams_cur[g].seq_id = W + 1 + g;222                    ngrams_cur[g].i_batch[0] = 0;223                    ngrams_cur[g].tokens [0] = id;224                }225 226                for (int j = 0; j < N - 1; j++) {227                    for (int g = 0; g < g_cur; g++) {228                        const int idx = id*(N - 1)*G + g*(N - 1);229 230                        const llama_token t = ngrams_observed.tokens[idx + j];231 232                        ngrams_cur[g].tokens [j + 1] = t;233                        ngrams_cur[g].i_batch[j + 1] = batch.n_tokens;234 235                        common_batch_add(batch, t, n_past + j + 1, { W + 1 + g }, true);236                    }237                }238            }239 240            // fill the remaining W - 1 tokens for the first level241            for (int i = 1; i < W; i++) {242                seq_id_look.resize(W - i);243                for (int j = 0; j < W - i; j++) {244                    seq_id_look[j] = i + j + 1;245                }246 247                common_batch_add(batch, tokens_j[0][i], n_past + i, seq_id_look, false);248            }249 250            // fill the rest of the levels251            for (int j = 1; j < N - 1; j++) {252                for (int i = 0; i < W; i++) {253                    common_batch_add(batch, tokens_j[j][i], n_past + j + i, { i + 1 }, j == N - 2);254                }255            }256        }257 258        if (llama_decode(ctx, batch) != 0) {259            LOG_ERR("\n\n%s: llama_decode failed - increase KV cache size\n", __func__);260            return 1;261        }262 263        int seq_id_best = 0;264 265        for (int v = 0; v < N; ++v) {266            int i_batch = 0;267 268            // if no active ngrams are left, it means the sampled token does not pass the verification269            if (v > 0) {270                for (int g = 0; g < (int) ngrams_cur.size(); g++) {271                    if (ngrams_cur[g].active) {272                        i_batch = ngrams_cur[g].i_batch[v];273                        seq_id_best = ngrams_cur[g].seq_id;274 275                        ++n_accept;276                        break;277                    }278                }279 280                // no more matches -> create a new batch281                if (i_batch == 0) {282                    break;283                }284            }285 286            // sample the next token287            id = common_sampler_sample(smpl, ctx, i_batch);288 289            common_sampler_accept(smpl, id, true);290 291            // print292            {293                const std::string token_str = common_token_to_piece(ctx, id);294 295                if (v == 0) {296                    LOG("%s", token_str.c_str());297                } else {298                    // print light cyan299                    LOG("\033[0;96m%s\033[0m", token_str.c_str());300                }301                fflush(stdout);302 303                if (llama_vocab_is_eog(vocab, id)) {304                    has_eos = true;305                }306 307                all.push_back(id);308            }309 310            ++n_predict;311            ++n_past;312 313            if ((params.n_predict >= 0 && n_predict > params.n_predict) || has_eos) {314                break;315            }316 317            // verify across active n-grams318            for (int g = 0; g < (int) ngrams_cur.size(); g++) {319                if (ngrams_cur[g].active) {320                    if (v == N - 1) {321                        ngrams_cur[g].active = false;322                    } else {323                        if (id != ngrams_cur[g].tokens[v + 1]) {324                            ngrams_cur[g].active = false;325                        }326                    }327                }328            }329 330            // print known n-grams starting with token id (debug)331            if (0 && v == 0) {332                if (ngrams_observed.cnt[id] > 0) {333                    LOG("\n - %d n-grams starting with '%s'\n", ngrams_observed.cnt[id], common_token_to_piece(ctx, id).c_str());334                }335 336                for (int i = 0; i < ngrams_observed.cnt[id]; i++) {337                    LOG("   - ngram %2d: ", i);338 339                    const int idx = id*(N - 1)*G + i*(N - 1);340 341                    for (int j = 0; j < N - 1; j++) {342                        const std::string token_str = common_token_to_piece(ctx, ngrams_observed.tokens[idx + j]);343 344                        LOG("%s", token_str.c_str());345                    }346 347                    LOG("\n");348                }349            }350 351            // update lookahead tokens352            {353                for (int i = 0; i < W; i++) {354                    tokens_j_prev[i] = tokens_j[0][i];355                }356 357                for (int j = 0; j < N - 2; j++) {358                    tokens_j[j] = tokens_j[j + 1];359                }360 361                if (v == 0) {362                    // sample from the last level363                    for (int i = 0; i < W; i++) {364                        tokens_j[N - 2][i] = common_sampler_sample(smpl, ctx, ngrams_cur.size()*(N-1) + W*(N - 2) + i);365                    }366                } else {367                    for (int i = 0; i < W; i++) {368                        // there are different ways to init these tokens369                        if (0) {370                            // random init371                            tokens_j[N - 2][i] = all[1 + rand() % (all.size() - 1)];372                        } else {373                            // init from the previous level374                            tokens_j[N - 2][i] = tokens_j[0][i];375                        }376                    }377                }378            }379 380            // update observed ngrams381            if (v == 0) {382                // the first token of the n-gram is determined by the index in the container so it is not stored383                std::vector<llama_token> ngram(N - 1);384 385                // n-gram generation386                // ref: https://github.com/hao-ai-lab/LookaheadDecoding/issues/14#issuecomment-1826198518387                for (int f = 0; f < W; ++f) {388                    const int ft = tokens_j_prev[f]; // first token of the n-gram389 390                    for (int j = 0; j < N - 1; ++j) {391                        ngram[j] = tokens_j[j][f];392                    }393 394                    // filter-out repeating n-grams395                    {396                        bool is_unique = true;397 398                        for (int k = 0; k < ngrams_observed.cnt[ft]; ++k) {399                            const int idx = ft*(N - 1)*G + k*(N - 1);400 401                            bool is_match = true;402                            for (int j = 0; j < N - 1; ++j) {403                                if (ngrams_observed.tokens[idx + j] != ngram[j]) {404                                    is_match = false;405                                    break;406                                }407                            }408 409                            if (is_match) {410                                is_unique = false;411                                break;412                            }413                        }414 415                        if (!is_unique) {416                            continue;417                        }418                    }419 420                    const int head = ngrams_observed.head[ft];421                    const int idx  = ft*(N - 1)*G + head*(N - 1);422 423                    for (int i = 0; i < N - 1; i++) {424                        ngrams_observed.tokens[idx + i] = ngram[i];425                    }426 427                    ngrams_observed.cnt[ft]  = std::min(G, ngrams_observed.cnt[ft] + 1);428                    ngrams_observed.head[ft] = (head + 1) % G;429 430                    ngrams_observed.n_total++;431                }432            }433        }434 435        if ((params.n_predict >= 0 && n_predict > params.n_predict) || has_eos) {436            break;437        }438 439        // KV cache management440        // if no verification token matched, we simply remove all cells from this batch -> no fragmentation441        llama_memory_seq_rm(mem, -1, n_past, -1);442 443        if (seq_id_best != 0) {444            // if a verification token matched, we keep the best sequence and remove the rest445            // this leads to some KV cache fragmentation446            llama_memory_seq_keep(mem, seq_id_best);447            llama_memory_seq_cp  (mem, seq_id_best, 0, -1, -1);448            llama_memory_seq_rm  (mem, seq_id_best,    -1, -1);449 450            for (int s = 1; s < W + G + 1; ++s) {451                llama_memory_seq_cp(mem, 0, s, -1, -1);452            }453        }454    }455 456    auto t_dec_end = ggml_time_us();457 458    LOG("\n\n");459 460    LOG_INF("encoded %4d tokens in %8.3f seconds, speed: %8.3f t/s\n", n_input,   (t_enc_end - t_enc_start) / 1e6f, inp.size() / ((t_enc_end - t_enc_start) / 1e6f));461    LOG_INF("decoded %4d tokens in %8.3f seconds, speed: %8.3f t/s\n", n_predict, (t_dec_end - t_dec_start) / 1e6f, n_predict  / ((t_dec_end - t_dec_start) / 1e6f));462 463    LOG_INF("\n");464    LOG_INF("W = %2d\n", W);465    LOG_INF("N = %2d\n", N);466    LOG_INF("G = %2d\n", G);467    LOG_INF("\n");468    LOG_INF("n_predict = %d\n", n_predict);469    LOG_INF("n_accept  = %d\n", n_accept);470 471    LOG_INF("\n");472    common_perf_print(ctx, smpl);473 474    common_sampler_free(smpl);475 476    llama_batch_free(batch);477 478    llama_backend_free();479 480    LOG("\n\n");481 482    return 0;483}484