KBaba7/llama.cpp
0
1#include "llama-context.h"2 3#include "llama-impl.h"4#include "llama-mmap.h"5 6#include <cassert>7#include <cmath>8#include <cstring>9#include <stdexcept>10 11void llama_set_k_shift(struct llama_context & lctx) {12 const int64_t kv_size = lctx.kv_self.size;13 14 assert(ggml_backend_buffer_is_host(lctx.inp_K_shift->buffer));15 16 int32_t * data = (int32_t *) lctx.inp_K_shift->data;17 18 for (int i = 0; i < kv_size; ++i) {19 data[i] = lctx.kv_self.cells[i].delta;20 }21}22 23void llama_set_s_copy(struct llama_context & lctx) {24 const int64_t kv_size = lctx.kv_self.size;25 26 assert(ggml_backend_buffer_is_host(lctx.inp_s_copy->buffer));27 28 int32_t * data = (int32_t *) lctx.inp_s_copy->data;29 30 for (int i = 0; i < kv_size; ++i) {31 data[i] = lctx.kv_self.cells[i].src;32 }33}34 35// llama input36 37static int32_t llama_relative_position_bucket(llama_pos x, llama_pos y, uint64_t n_buckets, bool bidirectional) {38 // TODO move to hparams if a T5 variant appears that uses a different value39 const int64_t max_distance = 128;40 41 if (bidirectional) {42 n_buckets >>= 1;43 }44 45 const int64_t max_exact = n_buckets >> 1;46 47 int32_t relative_position = x - y;48 int32_t relative_bucket = 0;49 if (bidirectional) {50 relative_bucket += (relative_position > 0) * n_buckets;51 relative_position = abs(relative_position);52 } else {53 relative_position = -std::min<int32_t>(relative_position, 0);54 }55 int32_t relative_position_if_large = floorf(max_exact + logf(1.0 * relative_position / max_exact) * (n_buckets - max_exact) / log(1.0 * max_distance / max_exact));56 relative_position_if_large = std::min<int32_t>(relative_position_if_large, n_buckets - 1);57 relative_bucket += (relative_position < max_exact ? relative_position : relative_position_if_large);58 return relative_bucket;59}60 61void llama_set_inputs(llama_context & lctx, const llama_ubatch & ubatch) {62 //63 // set input data64 //65 66 const auto & hparams = lctx.model.hparams;67 const auto & cparams = lctx.cparams;68 const auto & kv_self = lctx.kv_self;69 70 if (ubatch.token) {71 const int64_t n_tokens = ubatch.n_tokens;72 73 ggml_backend_tensor_set(lctx.inp_tokens, ubatch.token, 0, n_tokens*ggml_element_size(lctx.inp_tokens));74 }75 76 if (ubatch.embd) {77 const int64_t n_embd = hparams.n_embd;78 const int64_t n_tokens = ubatch.n_tokens;79 80 ggml_backend_tensor_set(lctx.inp_embd, ubatch.embd, 0, n_tokens*n_embd*ggml_element_size(lctx.inp_embd));81 }82 83 if (ubatch.pos && lctx.inp_pos) {84 const int64_t n_tokens = ubatch.n_tokens;85 auto n_pos = lctx.n_pos_per_token;86 ggml_backend_tensor_set(lctx.inp_pos, ubatch.pos, 0, n_tokens*n_pos*ggml_element_size(lctx.inp_pos));87 }88 89 if (hparams.causal_attn || cparams.pooling_type == LLAMA_POOLING_TYPE_NONE) {90 //GGML_ASSERT(lctx.inp_out_ids && "every model that can must skip unused outputs");91 92 if (!lctx.inp_out_ids) {93 LLAMA_LOG_WARN("%s: 'lctx.inp_out_ids' is not created\n", __func__);94 } else {95 const int64_t n_tokens = ubatch.n_tokens;96 97 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_out_ids->buffer));98 int32_t * data = (int32_t *) lctx.inp_out_ids->data;99 100 if (lctx.n_outputs == n_tokens) {101 for (int i = 0; i < n_tokens; ++i) {102 data[i] = i;103 }104 } else if (ubatch.output) {105 int32_t n_outputs = 0;106 for (int i = 0; i < n_tokens; ++i) {107 if (ubatch.output[i]) {108 data[n_outputs++] = i;109 }110 }111 // the graph needs to have been passed the correct number of outputs112 GGML_ASSERT(lctx.n_outputs == n_outputs);113 } else if (lctx.n_outputs == 1) {114 // only keep last output115 data[0] = n_tokens - 1;116 } else {117 GGML_ASSERT(lctx.n_outputs == 0);118 }119 }120 }121 122 GGML_ASSERT(123 // (!a || b) is a logical implication (a -> b)124 // !hparams.causal_attn -> !cparams.causal_attn125 (hparams.causal_attn || !cparams.causal_attn) &&126 "causal attention is not supported by this model"127 );128 129 if (lctx.inp_KQ_mask || lctx.inp_KQ_mask_swa) {130 // NOTE: hparams.causal_attn indicates the model is capable of generation and uses the kv cache.131 if (cparams.causal_attn && !lctx.is_encoding) {132 const int64_t n_kv = kv_self.n;133 const int64_t n_tokens = ubatch.n_tokens;134 const int64_t n_seq_tokens = ubatch.n_seq_tokens;135 const int64_t n_seqs = ubatch.n_seqs;136 137 138 float * data = nullptr;139 float * data_swa = nullptr;140 141 if (lctx.inp_KQ_mask) {142 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_KQ_mask->buffer));143 data = (float *) lctx.inp_KQ_mask->data;144 }145 146 if (lctx.inp_KQ_mask_swa) {147 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_KQ_mask_swa->buffer));148 data_swa = (float *) lctx.inp_KQ_mask_swa->data;149 }150 151 // For causal attention, use only the previous KV cells152 // of the correct sequence for each token of the ubatch.153 // It's assumed that if a token in the batch has multiple sequences, they are equivalent.154 for (int h = 0; h < 1; ++h) {155 for (int s = 0; s < n_seqs; ++s) {156 const llama_seq_id seq_id = ubatch.seq_id[s][0];157 158 for (int j = 0; j < n_seq_tokens; ++j) {159 const llama_pos pos = ubatch.pos[s*n_seq_tokens + j];160 161 for (int i = 0; i < n_kv; ++i) {162 float f;163 if (!kv_self.cells[i].has_seq_id(seq_id) || kv_self.cells[i].pos > pos) {164 f = -INFINITY;165 } else {166 if (hparams.use_alibi) {167 f = -std::abs(kv_self.cells[i].pos - pos);168 } else {169 f = 0.0f;170 }171 }172 173 if (data) {174 data[h*(n_kv*n_tokens) + s*(n_kv*n_seq_tokens) + j*n_kv + i] = f;175 }176 177 // may need to cut off old tokens for sliding window178 if (data_swa) {179 if (pos - kv_self.cells[i].pos >= (int32_t)hparams.n_swa) {180 f = -INFINITY;181 }182 data_swa[h*(n_kv*n_tokens) + s*(n_kv*n_seq_tokens) + j*n_kv + i] = f;183 }184 }185 }186 }187 188 if (data) {189 for (int i = n_tokens; i < GGML_PAD(n_tokens, GGML_KQ_MASK_PAD); ++i) {190 for (int j = 0; j < n_kv; ++j) {191 data[h*(n_kv*n_tokens) + i*n_kv + j] = -INFINITY;192 }193 }194 }195 196 if (data_swa) {197 for (int i = n_tokens; i < GGML_PAD(n_tokens, GGML_KQ_MASK_PAD); ++i) {198 for (int j = 0; j < n_kv; ++j) {199 data_swa[h*(n_kv*n_tokens) + i*n_kv + j] = -INFINITY;200 }201 }202 }203 }204 } else {205 const int64_t n_tokens = ubatch.n_tokens;206 const int64_t n_seq_tokens = ubatch.n_seq_tokens;207 const int64_t n_seqs = ubatch.n_seqs;208 // when using kv cache, the mask needs to match the kv cache size209 const int64_t n_stride = hparams.causal_attn && !lctx.is_encoding ? kv_self.n : n_tokens;210 211 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_KQ_mask->buffer));212 213 float * data = (float *) lctx.inp_KQ_mask->data;214 215 for (int h = 0; h < 1; ++h) {216 for (int s1 = 0; s1 < n_seqs; ++s1) {217 const llama_seq_id seq_id = ubatch.seq_id[s1][0];218 219 for (int j = 0; j < n_seq_tokens; ++j) {220 const int32_t tj = s1*n_seq_tokens + j;221 222 for (int s0 = 0; s0 < n_seqs; ++s0) {223 for (int i = 0; i < n_seq_tokens; ++i) {224 const int32_t ti = s0*n_seq_tokens + i;225 float f = -INFINITY;226 227 for (int s = 0; s < ubatch.n_seq_id[s0]; ++s) {228 if (ubatch.seq_id[s0][s] == seq_id) {229 if (hparams.use_alibi) {230 f = -std::abs(ubatch.pos[ti] - ubatch.pos[tj]);231 } else {232 f = 0.0f;233 }234 break;235 }236 }237 238 data[h*(n_tokens*n_tokens) + tj*n_stride + ti] = f;239 }240 }241 242 for (int i = n_tokens; i < n_stride; ++i) {243 data[h*(n_tokens*n_tokens) + tj*n_stride + i] = -INFINITY;244 }245 }246 }247 }248 }249 }250 251 if (cparams.embeddings && cparams.pooling_type == LLAMA_POOLING_TYPE_MEAN) {252 const int64_t n_tokens = ubatch.n_tokens;253 const int64_t n_seq_tokens = ubatch.n_seq_tokens;254 const int64_t n_seqs = ubatch.n_seqs;255 256 GGML_ASSERT(lctx.inp_mean);257 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_mean->buffer));258 259 float * data = (float *) lctx.inp_mean->data;260 memset(lctx.inp_mean->data, 0, n_tokens * n_tokens * ggml_element_size(lctx.inp_mean));261 262 std::vector<uint64_t> sum(n_tokens, 0);263 264 for (int s = 0; s < n_seqs; ++s) {265 const llama_seq_id seq_id = ubatch.seq_id[s][0];266 267 // TODO: adapt limits to n_seqs when ubatch.equal_seqs is true268 GGML_ASSERT(seq_id < n_tokens && "seq_id cannot be larger than n_tokens with pooling_type == MEAN");269 270 sum[seq_id] += ubatch.n_seq_tokens;271 }272 273 std::vector<float> div(n_tokens, 0.0f);274 for (int i = 0; i < n_tokens; ++i) {275 const uint64_t s = sum[i];276 if (s > 0) {277 div[i] = 1.0f/float(s);278 }279 }280 281 for (int s = 0; s < n_seqs; ++s) {282 const llama_seq_id seq_id = ubatch.seq_id[s][0];283 284 for (int i = 0; i < n_seq_tokens; ++i) {285 data[seq_id*n_tokens + s*n_seq_tokens + i] = div[seq_id];286 }287 }288 }289 290 if (cparams.embeddings && (291 cparams.pooling_type == LLAMA_POOLING_TYPE_CLS ||292 cparams.pooling_type == LLAMA_POOLING_TYPE_RANK)) {293 const int64_t n_tokens = ubatch.n_tokens;294 const int64_t n_seq_tokens = ubatch.n_seq_tokens;295 const int64_t n_seqs = ubatch.n_seqs;296 297 GGML_ASSERT(lctx.inp_cls);298 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_cls->buffer));299 300 uint32_t * data = (uint32_t *) lctx.inp_cls->data;301 memset(lctx.inp_cls->data, 0, n_tokens * ggml_element_size(lctx.inp_cls));302 303 for (int s = 0; s < n_seqs; ++s) {304 const llama_seq_id seq_id = ubatch.seq_id[s][0];305 306 // TODO: adapt limits to n_seqs when ubatch.equal_seqs is true307 GGML_ASSERT(seq_id < n_tokens && "seq_id cannot be larger than n_tokens with pooling_type == CLS or RANK");308 309 for (int i = 0; i < n_seq_tokens; ++i) {310 const llama_pos pos = ubatch.pos[s*n_seq_tokens + i];311 312 if (pos == 0) {313 data[seq_id] = s*n_seq_tokens + i;314 }315 }316 }317 }318 319 if (cparams.embeddings && cparams.pooling_type == LLAMA_POOLING_TYPE_LAST) {320 const int64_t n_tokens = ubatch.n_tokens;321 const int64_t n_seq_tokens = ubatch.n_seq_tokens;322 const int64_t n_seqs = ubatch.n_seqs;323 324 GGML_ASSERT(lctx.inp_cls);325 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_cls->buffer));326 327 uint32_t * data = (uint32_t *) lctx.inp_cls->data;328 memset(lctx.inp_cls->data, 0, n_tokens * ggml_element_size(lctx.inp_cls));329 330 std::vector<int> last_pos(n_tokens, -1);331 std::vector<int> last_row(n_tokens, -1);332 333 for (int s = 0; s < n_seqs; ++s) {334 const llama_seq_id seq_id = ubatch.seq_id[s][0];335 336 // TODO: adapt limits to n_seqs when ubatch.equal_seqs is true337 GGML_ASSERT(seq_id < n_tokens && "seq_id cannot be larger than n_tokens with pooling_type == LAST");338 339 for (int i = 0; i < n_seq_tokens; ++i) {340 const llama_pos pos = ubatch.pos[s*n_seq_tokens + i];341 342 if (pos >= last_pos[seq_id]) {343 last_pos[seq_id] = pos;344 last_row[seq_id] = s*n_seq_tokens + i;345 }346 }347 }348 349 for (int i = 0; i < n_tokens; ++i) {350 if (last_row[i] >= 0) {351 data[i] = last_row[i];352 }353 }354 }355 356 if (kv_self.recurrent) {357 const int64_t n_kv = kv_self.n;358 359 if (lctx.inp_s_mask) {360 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_s_mask->buffer));361 float * data = (float *) lctx.inp_s_mask->data;362 363 // clear unused states364 for (int i = 0; i < n_kv; ++i) {365 const uint32_t cell_id = i + kv_self.head;366 llama_kv_cell & kv_cell = lctx.kv_self.cells[cell_id];367 368 data[i] = (float) (kv_cell.src >= 0);369 370 // only clear once371 if (kv_cell.src < 0) {372 kv_cell.src = cell_id;373 }374 }375 }376 377 if (lctx.inp_s_copy) {378 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_s_copy->buffer));379 int32_t * data = (int32_t *) lctx.inp_s_copy->data;380 381 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n382 for (uint32_t i = 0; i < n_kv; ++i) {383 const uint32_t cell_id = i + kv_self.head;384 llama_kv_cell & kv_cell = lctx.kv_self.cells[cell_id];385 386 // prevent out-of-bound sources387 if (kv_cell.src < 0 || (uint32_t) kv_cell.src >= kv_self.size) {388 kv_cell.src = cell_id;389 }390 391 data[i] = kv_cell.src;392 393 // ensure copy only happens once394 if (kv_cell.src != (int32_t) cell_id) {395 kv_cell.src = cell_id;396 }397 }398 }399 }400 401 if (lctx.inp_pos_bucket) {402 const int64_t n_tokens = ubatch.n_tokens;403 404 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_pos_bucket->buffer));405 GGML_ASSERT(!ubatch.equal_seqs); // TODO: use ubatch.n_seqs instead of failing406 407 int32_t * data = (int32_t *) lctx.inp_pos_bucket->data;408 409 if (!lctx.is_encoding) {410 const int64_t n_kv = kv_self.n;411 for (int h = 0; h < 1; ++h) {412 for (int j = 0; j < n_tokens; ++j) {413 for (int i = 0; i < n_kv; ++i) {414 data[h*(n_kv*n_tokens) + j*n_kv + i] = llama_relative_position_bucket(lctx.kv_self.cells[i].pos, ubatch.pos[j], hparams.n_rel_attn_bkts, lctx.is_encoding);415 }416 }417 }418 } else {419 for (int h = 0; h < 1; ++h) {420 for (int j = 0; j < n_tokens; ++j) {421 for (int i = 0; i < n_tokens; ++i) {422 data[h*(n_tokens*n_tokens) + j*n_tokens + i] = llama_relative_position_bucket(ubatch.pos[i], ubatch.pos[j], hparams.n_rel_attn_bkts, lctx.is_encoding);423 }424 }425 }426 }427 }428 429 if (!lctx.is_encoding && lctx.inp_embd_enc) {430 assert(lctx.inp_embd_enc->type == GGML_TYPE_F32);431 assert((size_t) ggml_nelements(lctx.inp_embd_enc) == lctx.embd_enc.size());432 433 ggml_backend_tensor_set(lctx.inp_embd_enc, lctx.embd_enc.data(), 0, ggml_nbytes(lctx.inp_embd_enc));434 }435 436 if (!lctx.is_encoding && lctx.inp_KQ_mask_cross) {437 const int64_t n_output_enc = lctx.embd_enc.size() / hparams.n_embd;438 const int64_t n_tokens = ubatch.n_tokens;439 440 GGML_ASSERT(ggml_backend_buffer_is_host(lctx.inp_KQ_mask_cross->buffer));441 GGML_ASSERT(!ubatch.equal_seqs); // TODO: use ubatch.n_seqs instead of failing442 443 float * data = (float *) lctx.inp_KQ_mask_cross->data;444 445 for (int h = 0; h < 1; ++h) {446 for (int j = 0; j < n_tokens; ++j) {447 for (int i = 0; i < n_output_enc; ++i) {448 float f = -INFINITY;449 for (int s = 0; s < ubatch.n_seq_id[j]; ++s) {450 const llama_seq_id seq_id = ubatch.seq_id[j][s];451 if (lctx.seq_ids_enc[i].find(seq_id) != lctx.seq_ids_enc[i].end()) {452 f = 0.0f;453 }454 }455 data[h*(n_output_enc*n_tokens) + j*n_output_enc + i] = f;456 }457 }458 459 for (int i = n_tokens; i < GGML_PAD(n_tokens, GGML_KQ_MASK_PAD); ++i) {460 for (int j = 0; j < n_output_enc; ++j) {461 data[h*(n_output_enc*n_tokens) + i*n_output_enc + j] = -INFINITY;462 }463 }464 }465 }466}467 468// llama output469 470size_t llama_output_reserve(struct llama_context & lctx, size_t n_outputs) {471 const auto & cparams = lctx.cparams;472 const auto & hparams = lctx.model.hparams;473 const auto & vocab = lctx.model.vocab;474 475 const size_t n_outputs_max = std::max(n_outputs, (size_t) cparams.n_seq_max);476 477 const auto n_batch = cparams.n_batch;478 const auto n_vocab = vocab.n_tokens();479 const auto n_embd = hparams.n_embd;480 481 // TODO: use a per-batch flag for logits presence instead482 const bool has_logits = !cparams.embeddings;483 const bool has_embd = cparams.embeddings && (cparams.pooling_type == LLAMA_POOLING_TYPE_NONE);484 485 const size_t logits_size = has_logits ? n_vocab*n_outputs_max : 0;486 const size_t embd_size = has_embd ? n_embd*n_outputs_max : 0;487 488 if (lctx.output_ids.empty()) {489 // init, never resized afterwards490 lctx.output_ids.resize(n_batch);491 }492 493 const size_t prev_size = lctx.buf_output ? ggml_backend_buffer_get_size(lctx.buf_output.get()) : 0;494 const size_t new_size = (logits_size + embd_size) * sizeof(float);495 496 // alloc only when more than the current capacity is required497 // TODO: also consider shrinking the buffer498 if (!lctx.buf_output || prev_size < new_size) {499 if (lctx.buf_output) {500#ifndef NDEBUG501 // This doesn't happen often, but may be annoying in some cases (like the HellaSwag benchmark)502 LLAMA_LOG_INFO("%s: reallocating output buffer from size %.02f MiB to %.02f MiB\n", __func__, prev_size / 1024.0 / 1024.0, new_size / 1024.0 / 1024.0);503#endif504 lctx.buf_output = nullptr;505 lctx.logits = nullptr;506 lctx.embd = nullptr;507 }508 509 auto * buft = ggml_backend_cpu_buffer_type();510 // try to use the host buffer of the device where the output tensor is allocated for faster transfer to system memory511 auto * output_dev = lctx.model.dev_output();512 auto * output_dev_host_buft = output_dev ? ggml_backend_dev_host_buffer_type(output_dev) : nullptr;513 if (output_dev_host_buft) {514 buft = output_dev_host_buft;515 }516 lctx.buf_output.reset(ggml_backend_buft_alloc_buffer(buft, new_size));517 if (lctx.buf_output == nullptr) {518 LLAMA_LOG_ERROR("%s: failed to allocate output buffer of size %.2f MiB\n", __func__, new_size / (1024.0 * 1024.0));519 return 0;520 }521 }522 523 float * output_base = (float *) ggml_backend_buffer_get_base(lctx.buf_output.get());524 525 lctx.logits = has_logits ? output_base : nullptr;526 lctx.embd = has_embd ? output_base + logits_size : nullptr;527 528 lctx.output_size = n_outputs_max;529 lctx.logits_size = logits_size;530 lctx.embd_size = embd_size;531 532 // set all ids as invalid (negative)533 std::fill(lctx.output_ids.begin(), lctx.output_ids.end(), -1);534 535 ggml_backend_buffer_clear(lctx.buf_output.get(), 0);536 537 lctx.n_outputs = 0;538 539 return n_outputs_max;540}541 542void llama_output_reorder(struct llama_context & ctx) {543 std::vector<size_t> & out_ids = ctx.sbatch.out_ids;544 if (!out_ids.empty()) {545 const uint32_t n_vocab = ctx.model.vocab.n_tokens();546 const uint32_t n_embd = ctx.model.hparams.n_embd;547 548 const int32_t n_outputs = ctx.n_outputs;549 GGML_ASSERT((size_t) n_outputs == out_ids.size());550 551 // TODO: is there something more efficient which also minimizes swaps?552 // selection sort, to minimize swaps (from https://en.wikipedia.org/wiki/Selection_sort)553 for (int32_t i = 0; i < n_outputs - 1; ++i) {554 int32_t j_min = i;555 for (int32_t j = i + 1; j < n_outputs; ++j) {556 if (out_ids[j] < out_ids[j_min]) {557 j_min = j;558 }559 }560 if (j_min == i) { continue; }561 std::swap(out_ids[i], out_ids[j_min]);562 if (ctx.logits_size > 0) {563 for (uint32_t k = 0; k < n_vocab; k++) {564 std::swap(ctx.logits[i*n_vocab + k], ctx.logits[j_min*n_vocab + k]);565 }566 }567 if (ctx.embd_size > 0) {568 for (uint32_t k = 0; k < n_embd; k++) {569 std::swap(ctx.embd[i*n_embd + k], ctx.embd[j_min*n_embd + k]);570 }571 }572 }573 std::fill(ctx.output_ids.begin(), ctx.output_ids.end(), -1);574 for (int32_t i = 0; i < n_outputs; ++i) {575 ctx.output_ids[out_ids[i]] = i;576 }577 out_ids.clear();578 }579}580 581//582// interface implementation583//584 585void llama_free(struct llama_context * ctx) {586 delete ctx;587}588 589uint32_t llama_n_ctx(const struct llama_context * ctx) {590 return ctx->cparams.n_ctx;591}592 593uint32_t llama_n_batch(const struct llama_context * ctx) {594 return ctx->cparams.n_batch;595}596 597uint32_t llama_n_ubatch(const struct llama_context * ctx) {598 return ctx->cparams.n_ubatch;599}600 601uint32_t llama_n_seq_max(const struct llama_context * ctx) {602 return ctx->kv_self.size;603}604 605const struct llama_model * llama_get_model(const struct llama_context * ctx) {606 return &ctx->model;607}608 609enum llama_pooling_type llama_pooling_type(const struct llama_context * ctx) {610 return ctx->cparams.pooling_type;611}612 613void llama_attach_threadpool(614 struct llama_context * ctx,615 ggml_threadpool_t threadpool,616 ggml_threadpool_t threadpool_batch) {617 ctx->threadpool = threadpool;618 ctx->threadpool_batch = threadpool_batch ? threadpool_batch : threadpool;619}620 621void llama_detach_threadpool(struct llama_context * ctx) {622 ctx->threadpool = nullptr;623 ctx->threadpool_batch = nullptr;624}625 626void llama_set_n_threads(struct llama_context * ctx, int32_t n_threads, int32_t n_threads_batch) {627 ctx->cparams.n_threads = n_threads;628 ctx->cparams.n_threads_batch = n_threads_batch;629}630 631int32_t llama_n_threads(struct llama_context * ctx) {632 return ctx->cparams.n_threads;633}634 635int32_t llama_n_threads_batch(struct llama_context * ctx) {636 return ctx->cparams.n_threads_batch;637}638 639void llama_set_abort_callback(struct llama_context * ctx, bool (*abort_callback)(void * data), void * abort_callback_data) {640 ctx->abort_callback = abort_callback;641 ctx->abort_callback_data = abort_callback_data;642 643 for (auto & backend : ctx->backends) {644 auto * reg = ggml_backend_dev_backend_reg(ggml_backend_get_device(backend.get()));645 auto * set_abort_callback_fn = (ggml_backend_set_abort_callback_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_set_abort_callback");646 if (set_abort_callback_fn) {647 set_abort_callback_fn(backend.get(), ctx->abort_callback, ctx->abort_callback_data);648 }649 }650}651 652void llama_set_embeddings(struct llama_context * ctx, bool embeddings) {653 ctx->cparams.embeddings = embeddings;654}655 656void llama_set_causal_attn(struct llama_context * ctx, bool causal_attn) {657 ctx->cparams.causal_attn = causal_attn;658}659 660void llama_synchronize(struct llama_context * ctx) {661 ggml_backend_sched_synchronize(ctx->sched.get());662 663 // FIXME: if multiple single tokens are evaluated without a synchronization,664 // the stats will be added to the prompt evaluation stats665 // this should only happen when using batch size 1 to evaluate a batch666 667 // add the evaluation to the stats668 if (ctx->n_queued_tokens == 1) {669 if (!ctx->cparams.no_perf) {670 ctx->t_eval_us += ggml_time_us() - ctx->t_compute_start_us;671 }672 ctx->n_eval++;673 } else if (ctx->n_queued_tokens > 1) {674 if (!ctx->cparams.no_perf) {675 ctx->t_p_eval_us += ggml_time_us() - ctx->t_compute_start_us;676 }677 ctx->n_p_eval += ctx->n_queued_tokens;678 }679 680 // get a more accurate load time, upon first eval681 if (ctx->n_queued_tokens > 0 && !ctx->has_evaluated_once) {682 ctx->t_load_us = ggml_time_us() - ctx->t_start_us;683 ctx->has_evaluated_once = true;684 }685 686 ctx->n_queued_tokens = 0;687 ctx->t_compute_start_us = 0;688}689 690float * llama_get_logits(struct llama_context * ctx) {691 llama_synchronize(ctx);692 693 // reorder logits for backward compatibility694 // TODO: maybe deprecate this695 llama_output_reorder(*ctx);696 697 return ctx->logits;698}699 700float * llama_get_logits_ith(struct llama_context * ctx, int32_t i) {701 int32_t j = -1;702 703 llama_synchronize(ctx);704 705 try {706 if (ctx->logits == nullptr) {707 throw std::runtime_error("no logits");708 }709 710 if (i < 0) {711 j = ctx->n_outputs + i;712 if (j < 0) {713 throw std::runtime_error(format("negative index out of range [0, %d)", ctx->n_outputs));714 }715 } else if ((size_t) i >= ctx->output_ids.size()) {716 throw std::runtime_error(format("out of range [0, %zu)", ctx->output_ids.size()));717 } else {718 j = ctx->output_ids[i];719 }720 721 if (j < 0) {722 throw std::runtime_error(format("batch.logits[%d] != true", i));723 }724 if (j >= ctx->n_outputs) {725 // This should not happen726 throw std::runtime_error(format("corrupt output buffer (j=%d, n_outputs=%d)", j, ctx->n_outputs));727 }728 729 return ctx->logits + j*ctx->model.vocab.n_tokens();730 } catch (const std::exception & err) {731 LLAMA_LOG_ERROR("%s: invalid logits id %d, reason: %s\n", __func__, i, err.what());732#ifndef NDEBUG733 GGML_ABORT("fatal error");734#else735 return nullptr;736#endif737 }738}739 740float * llama_get_embeddings(struct llama_context * ctx) {741 llama_synchronize(ctx);742 743 // reorder embeddings for backward compatibility744 // TODO: maybe deprecate this745 llama_output_reorder(*ctx);746 747 return ctx->embd;748}749 750float * llama_get_embeddings_ith(struct llama_context * ctx, int32_t i) {751 int32_t j = -1;752 753 llama_synchronize(ctx);754 755 try {756 if (ctx->embd == nullptr) {757 throw std::runtime_error("no embeddings");758 }759 760 if (i < 0) {761 j = ctx->n_outputs + i;762 if (j < 0) {763 throw std::runtime_error(format("negative index out of range [0, %d)", ctx->n_outputs));764 }765 } else if ((size_t) i >= ctx->output_ids.size()) {766 throw std::runtime_error(format("out of range [0, %zu)", ctx->output_ids.size()));767 } else {768 j = ctx->output_ids[i];769 }770 771 if (j < 0) {772 throw std::runtime_error(format("batch.logits[%d] != true", i));773 }774 if (j >= ctx->n_outputs) {775 // This should not happen776 throw std::runtime_error(format("corrupt output buffer (j=%d, n_outputs=%d)", j, ctx->n_outputs));777 }778 779 return ctx->embd + j*ctx->model.hparams.n_embd;780 } catch (const std::exception & err) {781 LLAMA_LOG_ERROR("%s: invalid embeddings id %d, reason: %s\n", __func__, i, err.what());782#ifndef NDEBUG783 GGML_ABORT("fatal error");784#else785 return nullptr;786#endif787 }788}789 790float * llama_get_embeddings_seq(struct llama_context * ctx, llama_seq_id seq_id) {791 llama_synchronize(ctx);792 793 auto it = ctx->embd_seq.find(seq_id);794 if (it == ctx->embd_seq.end()) {795 return nullptr;796 }797 798 return it->second.data();799}800 801// llama state API802 803// deprecated804size_t llama_get_state_size(struct llama_context * ctx) {805 return llama_state_get_size(ctx);806}807 808// deprecated809size_t llama_copy_state_data(struct llama_context * ctx, uint8_t * dst) {810 return llama_state_get_data(ctx, dst, -1);811}812 813// deprecated814size_t llama_set_state_data(struct llama_context * ctx, const uint8_t * src) {815 return llama_state_set_data(ctx, src, -1);816}817 818// deprecated819bool llama_load_session_file(struct llama_context * ctx, const char * path_session, llama_token * tokens_out, size_t n_token_capacity, size_t * n_token_count_out) {820 return llama_state_load_file(ctx, path_session, tokens_out, n_token_capacity, n_token_count_out);821}822 823// deprecated824bool llama_save_session_file(struct llama_context * ctx, const char * path_session, const llama_token * tokens, size_t n_token_count) {825 return llama_state_save_file(ctx, path_session, tokens, n_token_count);826}827 828// TODO: replace all non-fatal assertions with returned errors or exceptions829struct llama_data_write {830 virtual void write(const void * src, size_t size) = 0;831 virtual void write_tensor_data(const struct ggml_tensor * tensor, size_t offset, size_t size) = 0;832 virtual size_t get_size_written() = 0;833 virtual ~llama_data_write() = default;834 835 void write_string(const std::string & str) {836 uint32_t str_size = str.size();837 838 write(&str_size, sizeof(str_size));839 write(str.data(), str_size);840 }841 842 void write_model_info(const struct llama_context * ctx) {843 const std::string arch_str = llm_arch_name(ctx->model.arch);844 write_string(arch_str);845 // TODO: add more model-specific info which should prevent loading the session file if not identical846 }847 848 //void write_rng(const std::mt19937 & rng) {849 // std::ostringstream rng_ss;850 // rng_ss << rng;851 852 // const std::string & rng_str = rng_ss.str();853 854 // write_string(rng_str);855 //}856 857 void write_output_ids(struct llama_context * ctx) {858 llama_output_reorder(*ctx);859 860 const uint32_t n_outputs = ctx->n_outputs;861 862 std::vector<int32_t> output_pos;863 864 const size_t n_batch = ctx->cparams.n_batch;865 const auto & output_ids = ctx->output_ids;866 867 GGML_ASSERT(n_outputs <= ctx->output_size);868 869 output_pos.resize(n_outputs);870 871 // build a more compact representation of the output ids872 for (size_t i = 0; i < n_batch; ++i) {873 // map an output id to a position in the batch874 int32_t pos = output_ids[i];875 if (pos >= 0) {876 GGML_ASSERT((uint32_t) pos < n_outputs);877 output_pos[pos] = i;878 }879 }880 881 write(&n_outputs, sizeof(n_outputs));882 883 if (n_outputs) {884 write(output_pos.data(), n_outputs * sizeof(int32_t));885 }886 }887 888 void write_logits(const struct llama_context * ctx) {889 const uint64_t logits_size = std::min((uint64_t) ctx->logits_size, (uint64_t) ctx->n_outputs * ctx->model.vocab.n_tokens());890 891 write(&logits_size, sizeof(logits_size));892 893 if (logits_size) {894 write(ctx->logits, logits_size * sizeof(float));895 }896 }897 898 void write_embeddings(const struct llama_context * ctx) {899 const uint64_t embeddings_size = std::min((uint64_t) ctx->embd_size, (uint64_t) ctx->n_outputs * ctx->model.hparams.n_embd);900 901 write(&embeddings_size, sizeof(embeddings_size));902 903 if (embeddings_size) {904 write(ctx->embd, embeddings_size * sizeof(float));905 }906 }907 908 void write_kv_cache_meta(const llama_kv_cache & kv_self, const std::vector<std::pair<uint32_t, uint32_t>> & cell_ranges, llama_seq_id seq_id = -1) {909 for (const auto & range : cell_ranges) {910 for (uint32_t i = range.first; i < range.second; ++i) {911 const auto & cell = kv_self.cells[i];912 const llama_pos pos = cell.pos;913 const uint32_t n_seq_id = seq_id == -1 ? cell.seq_id.size() : 0;914 915 write(&pos, sizeof(pos));916 write(&n_seq_id, sizeof(n_seq_id));917 918 if (n_seq_id) {919 for (auto seq_id : cell.seq_id) {920 write(&seq_id, sizeof(seq_id));921 }922 }923 }924 }925 }926 927 void write_kv_cache_data(const struct llama_context * ctx, const std::vector<std::pair<uint32_t, uint32_t>> & cell_ranges) {928 const struct llama_kv_cache & kv_self = ctx->kv_self;929 const struct llama_hparams & hparams = ctx->model.hparams;930 931 const uint32_t v_trans = kv_self.v_trans ? 1 : 0;932 const uint32_t n_layer = hparams.n_layer;933 934 write(&v_trans, sizeof(v_trans));935 write(&n_layer, sizeof(n_layer));936 937 std::vector<uint8_t> tmp_buf;938 939 // Iterate and write all the keys first, each row is a cell940 // Get whole range at a time941 for (uint32_t il = 0; il < n_layer; ++il) {942 const uint32_t n_embd_k_gqa = hparams.n_embd_k_gqa(il) + hparams.n_embd_k_s();943 944 // Write key type945 const int32_t k_type_i = (int32_t)kv_self.k_l[il]->type;946 write(&k_type_i, sizeof(k_type_i));947 948 // Write row size of key949 const uint64_t k_size_row = ggml_row_size(kv_self.k_l[il]->type, n_embd_k_gqa);950 write(&k_size_row, sizeof(k_size_row));951 952 // Read each range of cells of k_size length each into tmp_buf and write out953 for (const auto & range : cell_ranges) {954 const size_t range_size = range.second - range.first;955 const size_t buf_size = range_size * k_size_row;956 write_tensor_data(kv_self.k_l[il], range.first * k_size_row, buf_size);957 }958 }959 960 if (!kv_self.v_trans) {961 for (uint32_t il = 0; il < n_layer; ++il) {962 const uint32_t n_embd_v_gqa = hparams.n_embd_v_gqa(il) + hparams.n_embd_v_s();963 964 // Write value type965 const int32_t v_type_i = (int32_t)kv_self.v_l[il]->type;966 write(&v_type_i, sizeof(v_type_i));967 968 // Write row size of value969 const uint64_t v_size_row = ggml_row_size(kv_self.v_l[il]->type, n_embd_v_gqa);970 write(&v_size_row, sizeof(v_size_row));971 972 // Read each range of cells of v_size length each into tmp_buf and write out973 for (const auto & range : cell_ranges) {974 const size_t range_size = range.second - range.first;975 const size_t buf_size = range_size * v_size_row;976 write_tensor_data(kv_self.v_l[il], range.first * v_size_row, buf_size);977 }978 }979 } else {980 // When v is transposed, we also need the element size and get the element ranges from each row981 const uint32_t kv_size = kv_self.size;982 for (uint32_t il = 0; il < n_layer; ++il) {983 const uint32_t n_embd_v_gqa = hparams.n_embd_v_gqa(il) + hparams.n_embd_v_s();984 985 // Write value type986 const int32_t v_type_i = (int32_t)kv_self.v_l[il]->type;987 write(&v_type_i, sizeof(v_type_i));988 989 // Write element size990 const uint32_t v_size_el = ggml_type_size(kv_self.v_l[il]->type);991 write(&v_size_el, sizeof(v_size_el));992 993 // Write GQA embedding size994 write(&n_embd_v_gqa, sizeof(n_embd_v_gqa));995 996 // For each row, we get the element values of each cell997 for (uint32_t j = 0; j < n_embd_v_gqa; ++j) {998 // Read each range of cells of v_size_el length each into tmp_buf and write out999 for (const auto & range : cell_ranges) {1000 const size_t range_size = range.second - range.first;1001 const size_t src_offset = (range.first + j * kv_size) * v_size_el;1002 const size_t buf_size = range_size * v_size_el;1003 write_tensor_data(kv_self.v_l[il], src_offset, buf_size);1004 }1005 }1006 }1007 }1008 }1009 1010 void write_kv_cache(const struct llama_context * ctx, llama_seq_id seq_id = -1) {1011 const struct llama_kv_cache & kv_self = ctx->kv_self;1012 std::vector<std::pair<uint32_t, uint32_t>> cell_ranges; // ranges, from inclusive, to exclusive1013 uint32_t cell_count = 0;1014 1015 // Count the number of cells with the specified seq_id1016 // Find all the ranges of cells with this seq id (or all, when -1)1017 uint32_t cell_range_begin = kv_self.size;1018 for (uint32_t i = 0; i < kv_self.size; ++i) {1019 const auto & cell = kv_self.cells[i];1020 if ((seq_id == -1 && !cell.is_empty()) || cell.has_seq_id(seq_id)) {1021 ++cell_count;1022 if (cell_range_begin == kv_self.size) {1023 cell_range_begin = i;1024 }1025 } else {1026 if (cell_range_begin != kv_self.size) {1027 cell_ranges.emplace_back(cell_range_begin, i);1028 cell_range_begin = kv_self.size;1029 }1030 }1031 }1032 if (cell_range_begin != kv_self.size) {1033 cell_ranges.emplace_back(cell_range_begin, kv_self.size);1034 }1035 1036 // DEBUG CHECK: Sum of cell counts in ranges should equal the total cell count1037 uint32_t cell_count_check = 0;1038 for (const auto & range : cell_ranges) {1039 cell_count_check += range.second - range.first;1040 }1041 GGML_ASSERT(cell_count == cell_count_check);1042 1043 write(&cell_count, sizeof(cell_count));1044 1045 write_kv_cache_meta(kv_self, cell_ranges, seq_id);1046 write_kv_cache_data(ctx, cell_ranges);1047 }1048};1049 1050struct llama_data_read {1051 virtual const uint8_t * read(size_t size) = 0;1052 virtual void read_to(void * dst, size_t size) = 0;1053 virtual size_t get_size_read() = 0;1054 virtual ~llama_data_read() = default;1055 1056 void read_string(std::string & str) {1057 uint32_t str_size;1058 read_to(&str_size, sizeof(str_size));1059 1060 str.assign((const char *) read(str_size), str_size);1061 }1062 1063 // validate model information1064 void read_model_info(const struct llama_context * ctx) {1065 const std::string cur_arch_str = llm_arch_name(ctx->model.arch);1066 1067 std::string arch_str;1068 read_string(arch_str);1069 if (cur_arch_str != arch_str) {1070 throw std::runtime_error(format("wrong model arch: '%s' instead of '%s'", arch_str.c_str(), cur_arch_str.c_str()));1071 }1072 // TODO: add more info which needs to be identical but which is not verified otherwise1073 }1074 1075 //void read_rng(std::mt19937 & rng) {1076 // std::string rng_str;1077 // read_string(rng_str);1078 1079 // std::istringstream rng_ss(rng_str);1080 // rng_ss >> rng;1081 1082 // if (rng_ss.fail()) {1083 // throw std::runtime_error("failed to load RNG state");1084 // }1085 //}1086 1087 void read_output_ids(struct llama_context * ctx) {1088 std::vector<int32_t> output_pos;1089 1090 uint32_t n_outputs;1091 read_to(&n_outputs, sizeof(n_outputs));1092 1093 if (n_outputs > llama_output_reserve(*ctx, n_outputs)) {1094 throw std::runtime_error("could not reserve outputs");1095 }1096 1097 if (n_outputs) {1098 output_pos.resize(n_outputs);1099 read_to(output_pos.data(), n_outputs * sizeof(int32_t));1100 1101 for (int32_t i = 0; i < (int32_t) output_pos.size(); ++i) {1102 int32_t id = output_pos[i];1103 if ((uint32_t) id >= ctx->cparams.n_batch) {1104 throw std::runtime_error(format("invalid output id, %d does not fit in batch size of %u", id, ctx->cparams.n_batch));1105 }1106 ctx->output_ids[id] = i;1107 }1108 1109 ctx->n_outputs = n_outputs;1110 }1111 }1112 1113 void read_logits(struct llama_context * ctx) {1114 uint64_t logits_size;1115 read_to(&logits_size, sizeof(logits_size));1116 1117 if (ctx->logits_size < logits_size) {1118 throw std::runtime_error("logits buffer too small");1119 }1120 1121 if (logits_size) {1122 read_to(ctx->logits, logits_size * sizeof(float));1123 }1124 }1125 1126 void read_embeddings(struct llama_context * ctx) {1127 uint64_t embeddings_size;1128 read_to(&embeddings_size, sizeof(embeddings_size));1129 1130 if (ctx->embd_size < embeddings_size) {1131 throw std::runtime_error("embeddings buffer too small");1132 }1133 1134 if (embeddings_size) {1135 read_to(ctx->embd, embeddings_size * sizeof(float));1136 }1137 }1138 1139 bool read_kv_cache_meta(struct llama_context * ctx, uint32_t cell_count, llama_seq_id dest_seq_id = -1) {1140 struct llama_kv_cache & kv_self = ctx->kv_self;1141 1142 if (dest_seq_id != -1) {1143 // single sequence1144 1145 llama_kv_cache_seq_rm(kv_self, dest_seq_id, -1, -1);1146 1147 llama_ubatch batch = ctx->sbatch.reserve_ubatch(cell_count, /* has_embd */ false);1148 batch.n_tokens = cell_count;1149 batch.n_seq_tokens = cell_count;1150 batch.n_seqs = 1;1151 1152 for (uint32_t i = 0; i < cell_count; ++i) {1153 llama_pos pos;1154 uint32_t n_seq_id;1155 1156 read_to(&pos, sizeof(pos));1157 read_to(&n_seq_id, sizeof(n_seq_id));1158 1159 if (n_seq_id != 0) {1160 LLAMA_LOG_ERROR("%s: invalid seq_id-agnostic kv cell\n", __func__);1161 return false;1162 }1163 1164 batch.pos[i] = pos;1165 }1166 batch.n_seq_id[0] = 1;1167 batch.seq_id[0] = &dest_seq_id;1168 if (!llama_kv_cache_find_slot(kv_self, batch)) {1169 LLAMA_LOG_ERROR("%s: failed to find available cells in kv cache\n", __func__);1170 return false;1171 }1172 1173 // DEBUG CHECK: kv_self.head should be our first cell, kv_self.head + cell_count - 1 should be our last cell (verify seq_id and pos values)1174 // Assume that this is one contiguous block of cells1175 GGML_ASSERT(kv_self.head + cell_count <= kv_self.size);1176 GGML_ASSERT(kv_self.cells[kv_self.head].pos == batch.pos[0]);1177 GGML_ASSERT(kv_self.cells[kv_self.head + cell_count - 1].pos == batch.pos[cell_count - 1]);1178 GGML_ASSERT(kv_self.cells[kv_self.head].has_seq_id(dest_seq_id));1179 GGML_ASSERT(kv_self.cells[kv_self.head + cell_count - 1].has_seq_id(dest_seq_id));1180 } else {1181 // whole KV cache restore1182 1183 if (cell_count > kv_self.size) {1184 LLAMA_LOG_ERROR("%s: not enough cells in kv cache\n", __func__);1185 return false;1186 }1187 1188 llama_kv_cache_clear(kv_self);1189 1190 for (uint32_t i = 0; i < cell_count; ++i) {1191 llama_kv_cell & cell = kv_self.cells[i];1192 1193 llama_pos pos;1194 uint32_t n_seq_id;1195 1196 read_to(&pos, sizeof(pos));1197 read_to(&n_seq_id, sizeof(n_seq_id));1198 1199 cell.pos = pos;1200 