KBaba7/llama.cpp
0
1#include "arg.h"2#include "common.h"3#include "sampling.h"4#include "log.h"5#include "llama.h"6 7#include <cstdio>8#include <string>9#include <vector>10 11struct ngram_data {12 bool active = false;13 14 llama_seq_id seq_id = -1;15 16 std::vector<int> i_batch;17 18 std::vector<llama_token> tokens;19};20 21// n-gram container22struct ngram_container {23 ngram_container(int n_vocab, int N, int G) {24 cnt.resize(n_vocab);25 head.resize(n_vocab);26 tokens.resize(n_vocab * G * (N - 1));27 }28 29 int n_total = 0;30 31 std::vector<int> cnt;32 std::vector<int> head;33 34 // [n_vocab][G][N - 1]35 // for each token of the vocab, keep a ring-buffer of capacity G of n-grams of size N - 136 std::vector<llama_token> tokens;37};38 39int main(int argc, char ** argv) {40 common_params params;41 42 if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_COMMON)) {43 return 1;44 }45 46 common_init();47 48 const int W = 15; // lookahead window49 const int N = 5; // n-gram size50 const int G = 15; // max verification n-grams51 52 const bool dump_kv_cache = params.dump_kv_cache;53 54 // init llama.cpp55 llama_backend_init();56 llama_numa_init(params.numa);57 58 // load the target model59 common_init_result llama_init = common_init_from_params(params);60 61 llama_model * model = llama_init.model.get();62 llama_context * ctx = llama_init.context.get();63 64 const llama_vocab * vocab = llama_model_get_vocab(model);65 66 // Tokenize the prompt67 std::vector<llama_token> inp;68 std::vector<llama_token> all;69 70 inp = common_tokenize(ctx, params.prompt, true, true);71 all = inp;72 73 const int max_context_size = llama_n_ctx(ctx);74 const int max_tokens_list_size = max_context_size - 4;75 76 if ((int) inp.size() > max_tokens_list_size) {77 LOG_ERR("%s: prompt too long (%d tokens, max %d)\n", __func__, (int) inp.size(), max_tokens_list_size);78 return 1;79 }80 81 LOG("\n\n");82 83 for (auto id : inp) {84 LOG("%s", common_token_to_piece(ctx, id).c_str());85 }86 87 fflush(stderr);88 89 const int n_input = inp.size();90 91 const auto t_enc_start = ggml_time_us();92 93 // eval the prompt94 llama_decode(ctx, llama_batch_get_one( inp.data(), n_input - 1));95 llama_decode(ctx, llama_batch_get_one(&inp.back(), 1));96 97 for (int s = 1; s < W + G + 1; ++s) {98 llama_kv_cache_seq_cp(ctx, 0, s, -1, -1);99 }100 101 const auto t_enc_end = ggml_time_us();102 103 int n_predict = 0;104 int n_accept = 0;105 106 int n_past = inp.size();107 108 llama_token id = 0;109 110 // used to determine end of generation111 bool has_eos = false;112 113 // for each decoded batch, we have at most W + G + 1 distinct sequences:114 // seq_id == 0 : the current input token115 // seq_id [1, W] : tokens from the past N - 1 Jacobi iterations116 // seq_id [W + 1, W + G] : verification n-grams117 llama_batch batch = llama_batch_init(params.n_ctx, 0, W + G + 1);118 119 // target model sampling context120 struct common_sampler * smpl = common_sampler_init(model, params.sampling);121 122 // verification n-grams123 std::vector<ngram_data> ngrams_cur(G);124 125 // tokens for the past N - 1 Jacobi iterations126 std::vector<llama_token> tokens_j_prev(W);127 std::vector<std::vector<llama_token>> tokens_j(N - 1);128 for (int j = 0; j < N - 1; j++) {129 tokens_j[j].resize(W);130 131 for (int i = 0; i < W; i++) {132 // there are different ways to init these tokens133 if (0) {134 // initialize randomly from the prompt tokens135 tokens_j[j][i] = all[1 + rand() % (all.size() - 1)];136 } else {137 // initialize with a sequence of increasing numbers138 tokens_j[j][i] = 100 + i;139 }140 }141 }142 143 std::vector<llama_seq_id> seq_id_look;144 145 // the input token belongs both to all sequences146 std::vector<llama_seq_id> seq_id_all(W + G + 1);147 for (int i = 0; i < W + G + 1; i++) {148 seq_id_all[i] = i;149 }150 151 // here we keep adding new n-grams as we go152 ngram_container ngrams_observed(llama_vocab_n_tokens(vocab), N, G);153 154 // debug155 struct llama_kv_cache_view kvc_view = llama_kv_cache_view_init(ctx, W + G + 1);156 157 const auto t_dec_start = ggml_time_us();158 159 // sample first token160 {161 id = common_sampler_sample(smpl, ctx, 0);162 163 common_sampler_accept(smpl, id, true);164 165 {166 const std::string token_str = common_token_to_piece(ctx, id);167 168 LOG("%s", token_str.c_str());169 fflush(stdout);170 }171 }172 173 while (true) {174 // debug175 if (dump_kv_cache) {176 llama_kv_cache_view_update(ctx, &kvc_view);177 common_kv_cache_dump_view_seqs(kvc_view, 40);178 }179 180 // build the mask from https://lmsys.org/blog/2023-11-21-lookahead-decoding/181 //182 // Example for W = 5, N = 4, G = 2:183 // (I = input, L = lookahead, V = verification)184 //185 // Batch: 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20186 // T: -2 -2 -2 -2 -1 -1 -1 -1 -1 0 0 0 0 0 0187 // Info: I L L L L L L L L L L L L L L V V V V V V188 // Pos: 0 1 2 3 4 1 2 3 4 5 2 3 4 5 6 1 2 3 1 2 3 (+ n_past)189 // Logits: 1 0 0 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 1 1 1190 // ---------------------------------------------------------------------191 // Seq: 0192 // 1 1 1193 // 2 2 2 2194 // 3 3 3 3 3195 // 4 4 4 4 4 4196 // 5 5 5 5 5 5 5197 // 6 6 6 6198 // 7 7 7 7199 // ---------------------------------------------------------------------200 // | | | | | | | | | | |201 // V V V V V | | | | | |202 // j_tokens | | | | | |203 // V V V V V V204 // id205 {206 common_batch_clear(batch);207 208 // current token - first token of the first level209 common_batch_add(batch, id, n_past, seq_id_all, true);210 211 // verification n-grams - queue this before the lookahead tokens for less KV cache fragmentation212 {213 const int g_cur = ngrams_observed.cnt[id];214 215 ngrams_cur.resize(g_cur);216 for (int g = 0; g < g_cur; g++) {217 ngrams_cur[g].active = true;218 ngrams_cur[g].tokens.resize(N);219 ngrams_cur[g].i_batch.resize(N);220 ngrams_cur[g].seq_id = W + 1 + g;221 ngrams_cur[g].i_batch[0] = 0;222 ngrams_cur[g].tokens [0] = id;223 }224 225 for (int j = 0; j < N - 1; j++) {226 for (int g = 0; g < g_cur; g++) {227 const int idx = id*(N - 1)*G + g*(N - 1);228 229 const llama_token t = ngrams_observed.tokens[idx + j];230 231 ngrams_cur[g].tokens [j + 1] = t;232 ngrams_cur[g].i_batch[j + 1] = batch.n_tokens;233 234 common_batch_add(batch, t, n_past + j + 1, { W + 1 + g }, true);235 }236 }237 }238 239 // fill the remaining W - 1 tokens for the first level240 for (int i = 1; i < W; i++) {241 seq_id_look.resize(W - i);242 for (int j = 0; j < W - i; j++) {243 seq_id_look[j] = i + j + 1;244 }245 246 common_batch_add(batch, tokens_j[0][i], n_past + i, seq_id_look, false);247 }248 249 // fill the rest of the levels250 for (int j = 1; j < N - 1; j++) {251 for (int i = 0; i < W; i++) {252 common_batch_add(batch, tokens_j[j][i], n_past + j + i, { i + 1 }, j == N - 2);253 }254 }255 }256 257 if (llama_decode(ctx, batch) != 0) {258 LOG_ERR("\n\n%s: llama_decode failed - increase KV cache size\n", __func__);259 return 1;260 }261 262 int seq_id_best = 0;263 264 for (int v = 0; v < N; ++v) {265 int i_batch = 0;266 267 // if no active ngrams are left, it means the sampled token does not pass the verification268 if (v > 0) {269 for (int g = 0; g < (int) ngrams_cur.size(); g++) {270 if (ngrams_cur[g].active) {271 i_batch = ngrams_cur[g].i_batch[v];272 seq_id_best = ngrams_cur[g].seq_id;273 274 ++n_accept;275 break;276 }277 }278 279 // no more matches -> create a new batch280 if (i_batch == 0) {281 break;282 }283 }284 285 // sample the next token286 id = common_sampler_sample(smpl, ctx, i_batch);287 288 common_sampler_accept(smpl, id, true);289 290 // print291 {292 const std::string token_str = common_token_to_piece(ctx, id);293 294 if (v == 0) {295 LOG("%s", token_str.c_str());296 } else {297 // print light cyan298 LOG("\033[0;96m%s\033[0m", token_str.c_str());299 }300 fflush(stdout);301 302 if (llama_vocab_is_eog(vocab, id)) {303 has_eos = true;304 }305 306 all.push_back(id);307 }308 309 ++n_predict;310 ++n_past;311 312 if ((params.n_predict >= 0 && n_predict > params.n_predict) || has_eos) {313 break;314 }315 316 // verify across active n-grams317 for (int g = 0; g < (int) ngrams_cur.size(); g++) {318 if (ngrams_cur[g].active) {319 if (v == N - 1) {320 ngrams_cur[g].active = false;321 } else {322 if (id != ngrams_cur[g].tokens[v + 1]) {323 ngrams_cur[g].active = false;324 }325 }326 }327 }328 329 // print known n-grams starting with token id (debug)330 if (0 && v == 0) {331 if (ngrams_observed.cnt[id] > 0) {332 LOG("\n - %d n-grams starting with '%s'\n", ngrams_observed.cnt[id], common_token_to_piece(ctx, id).c_str());333 }334 335 for (int i = 0; i < ngrams_observed.cnt[id]; i++) {336 LOG(" - ngram %2d: ", i);337 338 const int idx = id*(N - 1)*G + i*(N - 1);339 340 for (int j = 0; j < N - 1; j++) {341 const std::string token_str = common_token_to_piece(ctx, ngrams_observed.tokens[idx + j]);342 343 LOG("%s", token_str.c_str());344 }345 346 LOG("\n");347 }348 }349 350 // update lookahead tokens351 {352 for (int i = 0; i < W; i++) {353 tokens_j_prev[i] = tokens_j[0][i];354 }355 356 for (int j = 0; j < N - 2; j++) {357 tokens_j[j] = tokens_j[j + 1];358 }359 360 if (v == 0) {361 // sample from the last level362 for (int i = 0; i < W; i++) {363 tokens_j[N - 2][i] = common_sampler_sample(smpl, ctx, ngrams_cur.size()*(N-1) + W*(N - 2) + i);364 }365 } else {366 for (int i = 0; i < W; i++) {367 // there are different ways to init these tokens368 if (0) {369 // random init370 tokens_j[N - 2][i] = all[1 + rand() % (all.size() - 1)];371 } else {372 // init from the previous level373 tokens_j[N - 2][i] = tokens_j[0][i];374 }375 }376 }377 }378 379 // update observed ngrams380 if (v == 0) {381 // the first token of the n-gram is determined by the index in the container so it is not stored382 std::vector<llama_token> ngram(N - 1);383 384 // n-gram generation385 // ref: https://github.com/hao-ai-lab/LookaheadDecoding/issues/14#issuecomment-1826198518386 for (int f = 0; f < W; ++f) {387 const int ft = tokens_j_prev[f]; // first token of the n-gram388 389 for (int j = 0; j < N - 1; ++j) {390 ngram[j] = tokens_j[j][f];391 }392 393 // filter-out repeating n-grams394 {395 bool is_unique = true;396 397 for (int k = 0; k < ngrams_observed.cnt[ft]; ++k) {398 const int idx = ft*(N - 1)*G + k*(N - 1);399 400 bool is_match = true;401 for (int j = 0; j < N - 1; ++j) {402 if (ngrams_observed.tokens[idx + j] != ngram[j]) {403 is_match = false;404 break;405 }406 }407 408 if (is_match) {409 is_unique = false;410 break;411 }412 }413 414 if (!is_unique) {415 continue;416 }417 }418 419 const int head = ngrams_observed.head[ft];420 const int idx = ft*(N - 1)*G + head*(N - 1);421 422 for (int i = 0; i < N - 1; i++) {423 ngrams_observed.tokens[idx + i] = ngram[i];424 }425 426 ngrams_observed.cnt[ft] = std::min(G, ngrams_observed.cnt[ft] + 1);427 ngrams_observed.head[ft] = (head + 1) % G;428 429 ngrams_observed.n_total++;430 }431 }432 }433 434 if ((params.n_predict >= 0 && n_predict > params.n_predict) || has_eos) {435 break;436 }437 438 // KV cache management439 // if no verification token matched, we simply remove all cells from this batch -> no fragmentation440 llama_kv_cache_seq_rm(ctx, -1, n_past, -1);441 442 if (seq_id_best != 0) {443 // if a verification token matched, we keep the best sequence and remove the rest444 // this leads to some KV cache fragmentation445 llama_kv_cache_seq_keep(ctx, seq_id_best);446 llama_kv_cache_seq_cp (ctx, seq_id_best, 0, -1, -1);447 llama_kv_cache_seq_rm (ctx, seq_id_best, -1, -1);448 449 for (int s = 1; s < W + G + 1; ++s) {450 llama_kv_cache_seq_cp(ctx, 0, s, -1, -1);451 }452 }453 }454 455 auto t_dec_end = ggml_time_us();456 457 LOG("\n\n");458 459 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));460 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));461 462 LOG_INF("\n");463 LOG_INF("W = %2d\n", W);464 LOG_INF("N = %2d\n", N);465 LOG_INF("G = %2d\n", G);466 LOG_INF("\n");467 LOG_INF("n_predict = %d\n", n_predict);468 LOG_INF("n_accept = %d\n", n_accept);469 470 LOG_INF("\n");471 common_perf_print(ctx, smpl);472 473 common_sampler_free(smpl);474 475 llama_kv_cache_view_free(&kvc_view);476 477 llama_batch_free(batch);478 479 llama_backend_free();480 481 LOG("\n\n");482 483 return 0;484}485 