Felipe97/llama-cpp-compiled
01.2k
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 