Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
llama-context.cpp1776 linesDownload Raw Back to src
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 

Showing the first 1,200 of 1776 lines. Download the file for the rest.