Team Ai
Datasetpublic

Brunobkr/llama.cpp_AlgMor24_github

ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes3.1kdownloads
llama-graph.cpp3781 linesDownload Raw Back to src
1#include "llama-graph.h"2 3#include "llama-impl.h"4#include "llama-model.h"5#include "llama-batch.h"6#include "llama-cparams.h"7#include "llama-sampler.h"8 9#include "llama-kv-cache.h"10#include "llama-kv-cache-iswa.h"11#include "llama-kv-cache-dsa.h"12#include "llama-kv-cache-msa.h"13#include "llama-kv-cache-dsv4.h"14#include "llama-memory-hybrid.h"15#include "llama-memory-hybrid-iswa.h"16#include "llama-memory-recurrent.h"17 18#include <cassert>19#include <cmath>20#include <cstring>21#include <numeric>22#include <sstream>23#include <string>24#include <unordered_set>25 26// dedup helpers27 28static ggml_tensor * build_attn_inp_kq_mask(29        ggml_context * ctx,30        const llama_kv_cache_context * mctx,31        const llama_ubatch & ubatch,32        const llama_cparams & cparams) {33    const auto n_kv     = mctx->get_n_kv();34    const auto n_tokens = ubatch.n_tokens;35    const auto n_stream = cparams.kv_unified ? 1 : ubatch.n_seqs_unq;36 37    // flash attention requires an f16 mask38    const auto type = cparams.flash_attn ? GGML_TYPE_F16 : GGML_TYPE_F32;39 40    ggml_tensor * res = ggml_new_tensor_4d(ctx, type, n_kv, n_tokens/n_stream, 1, n_stream);41    ggml_set_input(res);42    ggml_set_name(res, "attn_inp_kq_mask");43 44    return res;45}46 47static bool can_reuse_kq_mask(48        ggml_tensor * kq_mask,49        const llama_kv_cache_context * mctx,50        const llama_ubatch & ubatch,51        const llama_cparams & cparams) {52    const auto n_kv     = mctx->get_n_kv();53    const auto n_tokens = ubatch.n_tokens;54    const auto n_stream = cparams.kv_unified ? 1 : ubatch.n_seqs_unq;55 56    bool res = true;57 58    res &= (kq_mask->ne[0] == n_kv);59    res &= (kq_mask->ne[1] == n_tokens/n_stream);60    res &= (kq_mask->ne[2] == 1);61    res &= (kq_mask->ne[3] == n_stream);62 63    return res;64}65 66// impl67 68void llm_graph_input_embd::set_input(const llama_ubatch * ubatch) {69    if (ubatch->token) {70        const int64_t n_tokens = ubatch->n_tokens;71 72        ggml_backend_tensor_set(tokens, ubatch->token, 0, n_tokens*ggml_element_size(tokens));73    }74 75    if (ubatch->embd) {76        GGML_ASSERT(n_embd == embd->ne[0]);77 78        const int64_t n_tokens = ubatch->n_tokens;79 80        ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(embd));81    }82}83 84bool llm_graph_input_embd::can_reuse(const llm_graph_params & params) {85    bool res = true;86 87    res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);88    res &= (!params.ubatch.embd)  || (embd   &&   embd->ne[1] == params.ubatch.n_tokens);89 90    return res;91}92 93void llm_graph_input_embd_h::set_input(const llama_ubatch * ubatch) {94    const int64_t n_tokens = ubatch->n_tokens;95 96    if (ubatch->token) {97        ggml_backend_tensor_set(tokens, ubatch->token, 0, n_tokens*ggml_element_size(tokens));98    } else {99        // note: mtmd embedding input goes through here100        GGML_ASSERT(ubatch->embd);101        GGML_ASSERT(n_embd == embd->ne[0]);102 103        ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(h));104    }105 106    // TODO: extend llama_ubatch to differentiate between token embeddings and hidden states107    //       for now, we assume that the hidden state is always provided as an embedding108    //       ref: https://github.com/ggml-org/llama.cpp/pull/23643109    if (ubatch->embd) {110        GGML_ASSERT(n_embd == h->ne[0]);111 112        ggml_backend_tensor_set(h, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(h));113    }114}115 116bool llm_graph_input_embd_h::can_reuse(const llm_graph_params & params) {117    bool res = true;118 119    res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);120    res &= (!params.ubatch.embd)  || (embd   && embd->ne[1]   == params.ubatch.n_tokens);121    res &= (!params.ubatch.embd)  || (h      && h->ne[1]      == params.ubatch.n_tokens);122 123    return res;124}125 126void llm_graph_input_pos::set_input(const llama_ubatch * ubatch) {127    if (ubatch->pos && pos) {128        const int64_t n_tokens = ubatch->n_tokens;129 130        if (ubatch->token && n_pos_per_embd == 4) {131            // in case we're using M-RoPE with text tokens, convert the 1D positions to 4D132            // the 3 first dims are the same, and 4th dim is all 0133            std::vector<llama_pos> pos_data(n_tokens*n_pos_per_embd);134            // copy the first dimension135            for (int i = 0; i < n_tokens; ++i) {136                pos_data[               i] = ubatch->pos[i];137                pos_data[    n_tokens + i] = ubatch->pos[i];138                pos_data[2 * n_tokens + i] = ubatch->pos[i];139                pos_data[3 * n_tokens + i] = 0; // 4th dim is 0140            }141            ggml_backend_tensor_set(pos, pos_data.data(), 0, pos_data.size()*ggml_element_size(pos));142        } else {143            ggml_backend_tensor_set(pos, ubatch->pos, 0, n_tokens*n_pos_per_embd*ggml_element_size(pos));144        }145    }146}147 148bool llm_graph_input_pos::can_reuse(const llm_graph_params & params) {149    bool res = true;150 151    res &= pos->ne[0] == params.ubatch.n_tokens*n_pos_per_embd;152 153    return res;154}155 156void llm_graph_input_attn_temp::set_input(const llama_ubatch * ubatch) {157    if (ubatch->pos && attn_scale) {158        const int64_t n_tokens = ubatch->n_tokens;159 160        GGML_ASSERT(f_attn_temp_scale != 0.0f);161        GGML_ASSERT(n_attn_temp_floor_scale != 0);162 163        std::vector<float> attn_scale_data(n_tokens, 0.0f);164        for (int i = 0; i < n_tokens; ++i) {165            const float pos = ubatch->pos[i];166            attn_scale_data[i] = std::log(167                std::floor((pos + f_attn_temp_offset) / n_attn_temp_floor_scale) + 1.0168            ) * f_attn_temp_scale + 1.0;169        }170 171        ggml_backend_tensor_set(attn_scale, attn_scale_data.data(), 0, n_tokens*ggml_element_size(attn_scale));172    }173}174 175void llm_graph_input_pos_bucket::set_input(const llama_ubatch * ubatch) {176    if (pos_bucket) {177        const int64_t n_tokens = ubatch->n_tokens;178 179        GGML_ASSERT(ggml_backend_buffer_is_host(pos_bucket->buffer));180        GGML_ASSERT(!ubatch->equal_seqs()); // TODO: use ubatch->n_seqs instead of failing181 182        int32_t * data = (int32_t *) pos_bucket->data;183 184        for (int j = 0; j < n_tokens; ++j) {185            for (int i = 0; i < n_tokens; ++i) {186                data[j*n_tokens + i] = llama_relative_position_bucket(ubatch->pos[i], ubatch->pos[j], hparams.n_rel_attn_bkts, true);187            }188        }189    }190}191 192void llm_graph_input_pos_bucket_kv::set_input(const llama_ubatch * ubatch) {193    if (pos_bucket) {194        mctx->set_input_pos_bucket(pos_bucket, ubatch);195    }196}197 198void llm_graph_input_out_ids::set_input(const llama_ubatch * ubatch) {199    GGML_ASSERT(out_ids);200 201    const int64_t n_tokens = ubatch->n_tokens;202 203    GGML_ASSERT(ggml_backend_buffer_is_host(out_ids->buffer));204    int32_t * data = (int32_t *) out_ids->data;205 206    if (n_outputs == n_tokens) {207        for (int i = 0; i < n_tokens; ++i) {208            data[i] = i;209        }210 211        return;212    }213 214    GGML_ASSERT(ubatch->output);215 216    int n_outputs = 0;217 218    for (int i = 0; i < n_tokens; ++i) {219        if (ubatch->output[i]) {220            data[n_outputs++] = i;221        }222    }223}224 225bool llm_graph_input_out_ids::can_reuse(const llm_graph_params & params) {226    bool res = true;227 228    res &= n_outputs == params.n_outputs;229 230    return res;231}232 233void llm_graph_input_mean::set_input(const llama_ubatch * ubatch) {234    if (cparams.embeddings   &&235       (cparams.pooling_type == LLAMA_POOLING_TYPE_MEAN ||236        cparams.pooling_type == LLAMA_POOLING_TYPE_RANK )) {237 238        const int64_t n_tokens     = ubatch->n_tokens;239        const int64_t n_seq_tokens = ubatch->n_seq_tokens;240        const int64_t n_seqs_unq   = ubatch->n_seqs_unq;241 242        GGML_ASSERT(mean);243        GGML_ASSERT(ggml_backend_buffer_is_host(mean->buffer));244 245        float * data = (float *) mean->data;246        memset(mean->data, 0, n_tokens*n_seqs_unq*ggml_element_size(mean));247 248        std::vector<uint64_t> sums(n_seqs_unq, 0);249        for (int i = 0; i < n_tokens; i += n_seq_tokens) {250            for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {251                const llama_seq_id seq_id  = ubatch->seq_id[i][s];252                const int32_t      seq_idx = ubatch->seq_idx[seq_id];253 254                sums[seq_idx] += ubatch->n_seq_tokens;255            }256        }257 258        std::vector<float> div(n_seqs_unq, 0.0f);259        for (int s = 0; s < n_seqs_unq; ++s) {260            const uint64_t sum = sums[s];261            if (sum > 0) {262                div[s] = 1.0f/float(sum);263            }264        }265 266        for (int i = 0; i < n_tokens; i += n_seq_tokens) {267            for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {268                const llama_seq_id seq_id  = ubatch->seq_id[i][s];269                const int32_t      seq_idx = ubatch->seq_idx[seq_id];270 271                for (int j = 0; j < n_seq_tokens; ++j) {272                    data[seq_idx*n_tokens + i + j] = div[seq_idx];273                }274            }275        }276    }277}278 279void llm_graph_input_cls::set_input(const llama_ubatch * ubatch) {280    const int64_t n_tokens     = ubatch->n_tokens;281    const int64_t n_seqs_unq   = ubatch->n_seqs_unq;282 283    if (cparams.embeddings && (284        cparams.pooling_type == LLAMA_POOLING_TYPE_CLS  ||285        cparams.pooling_type == LLAMA_POOLING_TYPE_RANK ||286        cparams.pooling_type == LLAMA_POOLING_TYPE_LAST287    )) {288        GGML_ASSERT(cls);289        GGML_ASSERT(ggml_backend_buffer_is_host(cls->buffer));290 291        uint32_t * data = (uint32_t *) cls->data;292        memset(cls->data, 0, n_seqs_unq*ggml_element_size(cls));293 294        std::vector<int> target_pos(n_seqs_unq, -1);295        std::vector<int> target_row(n_seqs_unq, -1);296 297        const bool last = (298             cparams.pooling_type == LLAMA_POOLING_TYPE_LAST ||299            (cparams.pooling_type == LLAMA_POOLING_TYPE_RANK && (arch == LLM_ARCH_QWEN3 || arch == LLM_ARCH_QWEN3VL)) // qwen3 reranking & embedding models use last token300        );301 302        for (int i = 0; i < n_tokens; ++i) {303            const llama_pos pos = ubatch->pos[i];304 305            for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {306                const llama_seq_id seq_id  = ubatch->seq_id[i][s];307                const int32_t      seq_idx = ubatch->seq_idx[seq_id];308 309                if (310                    (target_pos[seq_idx] == -1) ||311                    ( last && pos >= target_pos[seq_idx]) ||312                    (!last && pos <  target_pos[seq_idx])313                ) {314                    target_pos[seq_idx] = pos;315                    target_row[seq_idx] = i;316                }317            }318        }319 320        for (int s = 0; s < n_seqs_unq; ++s) {321            if (target_row[s] >= 0) {322                data[s] = target_row[s];323            }324        }325    }326}327 328void llm_graph_input_rs::set_input(const llama_ubatch * ubatch) {329    GGML_UNUSED(ubatch);330 331    const int64_t n_rs = mctx->get_n_rs();332 333    if (s_copy) {334        GGML_ASSERT(ggml_backend_buffer_is_host(s_copy->buffer));335        int32_t * data = (int32_t *) s_copy->data;336 337        // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n338        for (uint32_t i = 0; i < n_rs; ++i) {339            data[i] = mctx->s_copy(i);340        }341    }342}343 344bool llm_graph_input_rs::can_reuse(const llm_graph_params & params) {345    const auto * mctx = static_cast<const llama_memory_recurrent_context *>(params.mctx);346 347    this->mctx = mctx;348 349    bool res = true;350 351    res &= s_copy->ne[0] == mctx->get_n_rs();352 353    res &= s_copy_main->ne[0]  == params.ubatch.n_seqs;354    res &= s_copy_extra->ne[0] == mctx->get_n_rs() - params.ubatch.n_seqs;355 356    res &= head == mctx->get_head();357    res &= rs_z == mctx->get_rs_z();358 359    return res;360}361 362void llm_graph_input_cross_embd::set_input(const llama_ubatch * ubatch) {363    GGML_UNUSED(ubatch);364 365    if (cross_embd && !cross->v_embd.empty()) {366        assert(cross_embd->type == GGML_TYPE_F32);367 368        ggml_backend_tensor_set(cross_embd, cross->v_embd.data(), 0, ggml_nbytes(cross_embd));369    }370}371 372template <typename T>373static void print_mask(const T * data, int64_t n_tokens, int64_t n_kv, int64_t n_swa, llama_swa_type swa_type) {374    LLAMA_LOG_DEBUG("%s: === Attention mask ===\n", __func__);375    const char * swa_type_str = "unknown";376 377    switch (swa_type) {378        case LLAMA_SWA_TYPE_NONE:      swa_type_str = "LLAMA_SWA_TYPE_NONE"; break;379        case LLAMA_SWA_TYPE_STANDARD:  swa_type_str = "LLAMA_SWA_TYPE_STANDARD"; break;380        case LLAMA_SWA_TYPE_CHUNKED:   swa_type_str = "LLAMA_SWA_TYPE_CHUNKED"; break;381        case LLAMA_SWA_TYPE_SYMMETRIC: swa_type_str = "LLAMA_SWA_TYPE_SYMMETRIC"; break;382    };383 384    LLAMA_LOG_DEBUG("%s: n_swa : %d, n_kv: %d, swa_type: %s\n", __func__, (int)n_swa, (int)n_kv, swa_type_str);385    LLAMA_LOG_DEBUG("%s: '0' = can attend, '∞' = masked\n", __func__);386    LLAMA_LOG_DEBUG("%s: Rows = query tokens, Columns = key/value tokens\n\n", __func__);387 388    LLAMA_LOG_DEBUG("    ");389    for (int j = 0; j < std::min((int64_t)20, n_kv); ++j) {390        LLAMA_LOG_DEBUG("%2d", j);391    }392    LLAMA_LOG_DEBUG("\n");393 394    for (int i = 0; i < std::min((int64_t)20, n_tokens); ++i) {395        LLAMA_LOG_DEBUG(" %2d ", i);396        for (int j = 0; j < std::min((int64_t)20, n_kv); ++j) {397            float val = llama_cast<float>(data[i * n_kv + j]);398            if (val == -INFINITY) {399                LLAMA_LOG_DEBUG(" ∞");400            } else {401                LLAMA_LOG_DEBUG(" 0");402            }403        }404        LLAMA_LOG_DEBUG("\n");405    }406}407 408void llm_graph_input_attn_no_cache::set_input(const llama_ubatch * ubatch) {409    const int64_t n_kv     = ubatch->n_tokens;410    const int64_t n_tokens = ubatch->n_tokens;411 412    const auto fill_mask = [&](auto * data, int64_t ne, int n_swa, llama_swa_type swa_type) {413        using T = std::remove_reference_t<decltype(*data)>;414        std::fill(data, data + ne, llama_cast<T>(-INFINITY));415 416        for (int i1 = 0; i1 < n_tokens; ++i1) {417            const llama_seq_id s1 = ubatch->seq_id[i1][0];418            const llama_pos    p1 = ubatch->pos[i1];419 420            const uint64_t idst = i1*n_kv;421 422            for (int i0 = 0; i0 < n_tokens; ++i0) {423                const llama_seq_id s0 = ubatch->seq_id[i0][0];424                const llama_pos p0    = ubatch->pos[i0];425 426                // mask different sequences427                if (s0 != s1) {428                    continue;429                }430 431                // mask future tokens432                if (cparams.causal_attn && p0 > p1) {433                    continue;434                }435 436                // apply SWA if any437                if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) {438                    continue;439                }440 441                data[idst + i0] = llama_cast<T>(hparams.use_alibi ? -std::abs(p0 - p1) : 0.0f);442            }443        }444 445        if (debug) {446            print_mask(data, n_tokens, n_kv, n_swa, swa_type);447        }448    };449 450    GGML_ASSERT(self_kq_mask);451    GGML_ASSERT(ggml_backend_buffer_is_host(self_kq_mask->buffer));452    if (self_kq_mask->type == GGML_TYPE_F16) {453        fill_mask((ggml_fp16_t *) self_kq_mask->data, ggml_nelements(self_kq_mask), 0, LLAMA_SWA_TYPE_NONE);454    } else {455        fill_mask((float       *) self_kq_mask->data, ggml_nelements(self_kq_mask), 0, LLAMA_SWA_TYPE_NONE);456    }457 458    if (hparams.swa_type != LLAMA_SWA_TYPE_NONE) {459        GGML_ASSERT(self_kq_mask_swa);460        GGML_ASSERT(ggml_backend_buffer_is_host(self_kq_mask_swa->buffer));461        if (self_kq_mask_swa->type == GGML_TYPE_F16) {462            fill_mask((ggml_fp16_t *) self_kq_mask_swa->data, ggml_nelements(self_kq_mask_swa), hparams.n_swa, hparams.swa_type);463        } else {464            fill_mask((float       *) self_kq_mask_swa->data, ggml_nelements(self_kq_mask_swa), hparams.n_swa, hparams.swa_type);465        }466    }467}468 469void llm_graph_input_attn_kv::set_input(const llama_ubatch * ubatch) {470    mctx->set_input_k_idxs(self_k_idxs, ubatch);471    mctx->set_input_v_idxs(self_v_idxs, ubatch);472 473    // the mask is left unallocated when the graph only stores K/V without attending474    // (e.g. DFlash's KV-injection pass)475    if (self_kq_mask && self_kq_mask->buffer) {476        mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);477    }478 479    if (self_k_rot && self_k_rot->buffer) {480        mctx->set_input_k_rot(self_k_rot);481    }482 483    if (self_v_rot && self_v_rot->buffer) {484        mctx->set_input_v_rot(self_v_rot);485    }486}487 488bool llm_graph_input_attn_kv::can_reuse(const llm_graph_params & params) {489    const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);490 491    this->mctx = mctx;492 493    bool res = true;494 495    res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;496  //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there497 498    res &= can_reuse_kq_mask(self_kq_mask, mctx, params.ubatch, params.cparams);499 500    return res;501}502 503void llm_graph_input_attn_k::set_input(const llama_ubatch * ubatch) {504    mctx->set_input_k_idxs(self_k_idxs, ubatch);505 506    mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);507}508 509bool llm_graph_input_attn_k::can_reuse(const llm_graph_params & params) {510    const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);511 512    this->mctx = mctx;513 514    bool res = true;515 516    res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;517 518    res &= can_reuse_kq_mask(self_kq_mask, mctx, params.ubatch, params.cparams);519 520    return res;521}522 523llm_graph_input_attn_kv_msa::llm_graph_input_attn_kv_msa(524        const llama_hparams & hparams,525        const llama_cparams & cparams,526        const llama_kv_cache_msa_context * mctx) :527    llm_graph_input_attn_kv(hparams, cparams, mctx->get_base()),528    mctx_msa(mctx) {529}530 531void llm_graph_input_attn_kv_msa::set_input(const llama_ubatch * ubatch) {532    llm_graph_input_attn_kv::set_input(ubatch);533 534    if (self_k_idxs_idx) {535        mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch);536    }537}538 539bool llm_graph_input_attn_kv_msa::can_reuse(const llm_graph_params & params) {540    mctx_msa = static_cast<const llama_kv_cache_msa_context *>(params.mctx);541 542    // the parent class operates on the base cache context543    this->mctx = mctx_msa->get_base();544 545    bool res = true;546 547    res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;548    if (self_k_idxs_idx) {549        res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens;550    }551 552    res &= can_reuse_kq_mask(self_kq_mask, this->mctx, params.ubatch, params.cparams);553 554    return res;555}556 557void llm_graph_input_attn_k_dsa::set_input(const llama_ubatch * ubatch) {558    mctx->get_mla()->set_input_k_idxs(self_k_idxs_mla, ubatch);559 560    mctx->get_mla()->set_input_kq_mask(self_kq_mask_mla, ubatch, cparams.causal_attn);561 562    mctx->get_lid()->set_input_k_idxs(self_k_idxs_lid, ubatch);563 564    mctx->get_lid()->set_input_kq_mask(self_kq_mask_lid, ubatch, cparams.causal_attn);565 566    mctx->get_lid()->set_input_k_rot(self_k_rot_lid);567}568 569bool llm_graph_input_attn_k_dsa::can_reuse(const llm_graph_params & params) {570    const auto * mctx = static_cast<const llama_kv_cache_dsa_context *>(params.mctx);571 572    this->mctx = mctx;573 574    bool res = true;575 576    res &= self_k_idxs_mla->ne[0] == params.ubatch.n_tokens;577    res &= self_k_idxs_lid->ne[0] == params.ubatch.n_tokens;578 579    res &= can_reuse_kq_mask(self_kq_mask_mla, mctx->get_mla(), params.ubatch, params.cparams);580    res &= can_reuse_kq_mask(self_kq_mask_lid, mctx->get_lid(), params.ubatch, params.cparams);581 582    return res;583}584 585void llm_graph_input_attn_kv_iswa::set_input(const llama_ubatch * ubatch) {586    // base tensors may not be allocated if there are no non-SWA attention layers587    if (self_k_idxs && self_k_idxs->buffer) {588        mctx->get_base()->set_input_k_idxs(self_k_idxs, ubatch);589        if (self_v_idxs) {590            mctx->get_base()->set_input_v_idxs(self_v_idxs, ubatch);591        }592    }593 594    // the kq mask guards on its own buffer: shared cells leave idxs unbacked while the mask stays live595    if (self_kq_mask && self_kq_mask->buffer) {596        mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);597    }598 599    // swa tensors may not be allocated if there are no SWA attention layers600    if (self_k_idxs_swa && self_k_idxs_swa->buffer) {601        mctx->get_swa()->set_input_k_idxs(self_k_idxs_swa, ubatch);602        if (self_v_idxs_swa) {603            mctx->get_swa()->set_input_v_idxs(self_v_idxs_swa, ubatch);604        }605    }606 607    if (self_kq_mask_swa && self_kq_mask_swa->buffer) {608        mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);609    }610 611    if (self_k_rot && self_k_rot->buffer) {612        mctx->get_base()->set_input_k_rot(self_k_rot);613    }614 615    if (self_v_rot && self_v_rot->buffer) {616        mctx->get_base()->set_input_v_rot(self_v_rot);617    }618 619    if (self_k_rot_swa && self_k_rot_swa->buffer) {620        mctx->get_swa()->set_input_k_rot(self_k_rot_swa);621    }622 623    if (self_v_rot_swa && self_v_rot_swa->buffer) {624        mctx->get_swa()->set_input_v_rot(self_v_rot_swa);625    }626}627 628bool llm_graph_input_attn_kv_iswa::can_reuse(const llm_graph_params & params) {629    const auto * mctx = static_cast<const llama_kv_cache_iswa_context *>(params.mctx);630 631    this->mctx = mctx;632 633    bool res = true;634 635    // base tensors may not be allocated if there are no non-SWA attention layers636    if (self_k_idxs && self_k_idxs->buffer) {637        res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;638      //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there639    }640 641    if (self_kq_mask && self_kq_mask->buffer) {642        res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams);643    }644 645    // swa tensors may not be allocated if there are no SWA attention layers646    if (self_k_idxs_swa && self_k_idxs_swa->buffer) {647        res &= self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;648      //res &= self_v_idxs_swa->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there649    }650 651    if (self_kq_mask_swa && self_kq_mask_swa->buffer) {652        res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams);653    }654 655    return res;656}657 658void llm_graph_input_attn_k_iswa::set_input(const llama_ubatch * ubatch) {659    // base tensors may not be allocated if there are no non-SWA attention layers660    if (self_k_idxs && self_k_idxs->buffer) {661        mctx->get_base()->set_input_k_idxs(self_k_idxs, ubatch);662    }663 664    // the kq mask guards on its own buffer: shared cells leave idxs unbacked while the mask stays live665    if (self_kq_mask && self_kq_mask->buffer) {666        mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);667    }668 669    // swa tensors may not be allocated if there are no SWA attention layers670    if (self_k_idxs_swa && self_k_idxs_swa->buffer) {671        mctx->get_swa()->set_input_k_idxs(self_k_idxs_swa, ubatch);672    }673 674    if (self_kq_mask_swa && self_kq_mask_swa->buffer) {675        mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);676    }677 678    if (self_k_rot && self_k_rot->buffer) {679        mctx->get_base()->set_input_k_rot(self_k_rot);680    }681 682    if (self_k_rot_swa && self_k_rot_swa->buffer) {683        mctx->get_swa()->set_input_k_rot(self_k_rot_swa);684    }685}686 687bool llm_graph_input_attn_k_iswa::can_reuse(const llm_graph_params & params) {688    const auto * mctx = static_cast<const llama_kv_cache_iswa_context *>(params.mctx);689 690    this->mctx = mctx;691 692    bool res = true;693 694    // base tensors may not be allocated if there are no non-SWA attention layers695    if (self_k_idxs && self_k_idxs->buffer) {696        res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;697    }698 699    if (self_kq_mask && self_kq_mask->buffer) {700        res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams);701    }702 703    // swa tensors may not be allocated if there are no SWA attention layers704    if (self_k_idxs_swa && self_k_idxs_swa->buffer) {705        res &= self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;706    }707 708    if (self_kq_mask_swa && self_kq_mask_swa->buffer) {709        res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams);710    }711 712    return res;713}714 715static void dsv4_set_i64(ggml_tensor * dst, const std::vector<int64_t> & src) {716    if (!dst || !dst->buffer) {717        return;718    }719 720    GGML_ASSERT(dst->ne[0] == (int64_t) src.size());721    ggml_backend_tensor_set(dst, src.data(), 0, src.size()*ggml_element_size(dst));722}723 724static void dsv4_set_i32(ggml_tensor * dst, const std::vector<int32_t> & src) {725    if (!dst || !dst->buffer) {726        return;727    }728 729    GGML_ASSERT(dst->ne[0] == (int64_t) src.size());730    ggml_backend_tensor_set(dst, src.data(), 0, src.size()*ggml_element_size(dst));731}732 733static void dsv4_set_kq_mask(734        ggml_tensor * dst,735        const llama_kv_cache_dsv4_context::comp_plan & plan,736        uint32_t n_tokens,737        int64_t n_stream) {738    if (!dst || !dst->buffer) {739        return;740    }741 742    GGML_ASSERT(dst->type == GGML_TYPE_F32 || dst->type == GGML_TYPE_F16);743    GGML_ASSERT(n_stream > 0);744    GGML_ASSERT(n_tokens%n_stream == 0);745    GGML_ASSERT(dst->ne[0] == plan.n_kv);746    GGML_ASSERT(dst->ne[1] == (int64_t) n_tokens/n_stream);747    GGML_ASSERT(dst->ne[2] == 1);748    GGML_ASSERT(dst->ne[3] == n_stream);749    GGML_ASSERT((int64_t) plan.n_visible.size() == (int64_t) n_tokens);750    GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));751 752    if (dst->type == GGML_TYPE_F32) {753        float * data = (float *) dst->data;754 755        for (int64_t i = 0; i < (int64_t) n_tokens; ++i) {756            const int32_t n_visible = plan.n_visible[i];757 758            for (int64_t j = 0; j < dst->ne[0]; ++j) {759                data[i*dst->ne[0] + j] = j < n_visible ? 0.0f : -INFINITY;760            }761        }762    } else if (dst->type == GGML_TYPE_F16) {763        ggml_fp16_t * data = (ggml_fp16_t *) dst->data;764        const ggml_fp16_t fp16_ninf = llama_cast<ggml_fp16_t>(-INFINITY);765        const ggml_fp16_t fp16_zero = llama_cast<ggml_fp16_t>(0.0f);766 767        for (int64_t i = 0; i < (int64_t) n_tokens; ++i) {768            const int32_t n_visible = plan.n_visible[i];769 770            for (int64_t j = 0; j < dst->ne[0]; ++j) {771                data[i*dst->ne[0] + j] = j < n_visible ? fp16_zero : fp16_ninf;772            }773        }774    }775}776 777static ggml_tensor * dsv4_build_raw_kq_mask(778        ggml_context * ctx,779        const llama_kv_cache_dsv4_raw_context * mctx,780        const llama_ubatch & ubatch,781        const llama_cparams & cparams,782        int64_t n_stream) {783    const auto n_kv     = mctx->get_n_kv();784    const auto n_tokens = ubatch.n_tokens;785 786    GGML_ASSERT(n_stream > 0);787    GGML_ASSERT(n_tokens%n_stream == 0);788 789    const auto type = cparams.flash_attn ? GGML_TYPE_F16 : GGML_TYPE_F32;790 791    ggml_tensor * res = ggml_new_tensor_4d(ctx, type, n_kv, n_tokens/n_stream, 1, n_stream);792    ggml_set_input(res);793    ggml_set_name(res, "attn_inp_kq_mask");794 795    return res;796}797 798static bool dsv4_can_reuse_raw_kq_mask(799        ggml_tensor * kq_mask,800        const llama_kv_cache_dsv4_raw_context * mctx,801        const llama_ubatch & ubatch,802        int64_t n_stream) {803    const auto n_kv     = mctx->get_n_kv();804    const auto n_tokens = ubatch.n_tokens;805 806    GGML_ASSERT(n_stream > 0);807 808    bool res = true;809 810    res &= (kq_mask->ne[0] == n_kv);811    res &= (kq_mask->ne[1] == n_tokens/n_stream);812    res &= (kq_mask->ne[2] == 1);813    res &= (kq_mask->ne[3] == n_stream);814 815    return res;816}817 818static std::string dsv4_plan_positions(const std::vector<int32_t> & values) {819    std::ostringstream ss;820    ss << "[";821    for (size_t i = 0; i < values.size(); ++i) {822        if (i > 0) {823            ss << ", ";824        }825        ss << values[i];826    }827    ss << "]";828    return ss.str();829}830 831static bool dsv4_compress_debug() {832    static const bool debug = []() {833        const char * env = getenv("LLAMA_DSV4_COMPRESS_DEBUG");834        return env && atoi(env) > 0;835    }();836 837    return debug;838}839 840static void dsv4_set_comp_inputs(841        const llm_graph_input_dsv4::comp_input & inp,842        const llama_kv_cache_dsv4_context::comp_plan & plan,843        const char * name,844        bool debug,845        uint32_t n_tokens,846        int64_t n_stream) {847    dsv4_set_i32(inp.state_pos, plan.state_pos);848    dsv4_set_i32(inp.state_persist_src_idxs, plan.state_persist_src_idxs);849    dsv4_set_i32(inp.state_persist_dst_idxs, plan.state_persist_dst_idxs);850    dsv4_set_i32(inp.state_restore_src_idxs, plan.state_restore_src_idxs);851    dsv4_set_i32(inp.state_restore_dst_idxs, plan.state_restore_dst_idxs);852    dsv4_set_i32(inp.state_snapshot_src_idxs, plan.state_snapshot_src_idxs);853    dsv4_set_i32(inp.state_snapshot_dst_idxs, plan.state_snapshot_dst_idxs);854    dsv4_set_i32(inp.state_read_idxs, plan.state_read_idxs);855    dsv4_set_i64(inp.state_write_idxs, plan.state_write_idxs);856    dsv4_set_i32(inp.state_write_pos, plan.state_write_pos);857    dsv4_set_kq_mask(inp.kq_mask, plan, n_tokens, n_stream);858 859    if (debug || dsv4_compress_debug()) {860        LLAMA_LOG_INFO("%s: %s n_tokens=%u, n_stream=%d, state_persist_dst=%s, state_write_pos=%s\n",861                __func__, name, n_tokens, (int) n_stream,862                dsv4_plan_positions(plan.state_persist_dst_idxs).c_str(),863                dsv4_plan_positions(plan.state_write_pos).c_str());864    }865}866 867static bool dsv4_can_reuse_tensor_1d(ggml_tensor * t, int64_t ne0) {868    return (t == nullptr && ne0 == 0) || (t != nullptr && t->ne[0] == ne0);869}870 871static bool dsv4_can_reuse_kq_mask(872        ggml_tensor * t,873        const llama_kv_cache_dsv4_context::comp_plan & plan,874        uint32_t n_tokens,875        int64_t n_stream) {876    if (plan.n_kv == 0) {877        return t == nullptr;878    }879 880    GGML_ASSERT(n_stream > 0);881 882    return t != nullptr &&883           t->ne[0] == plan.n_kv &&884           t->ne[1] == (int64_t) n_tokens/n_stream &&885           t->ne[2] == 1 &&886           t->ne[3] == n_stream;887}888 889static bool dsv4_can_reuse_comp_input(890        const llm_graph_input_dsv4::comp_input & inp,891        const llama_kv_cache_dsv4_context::comp_plan & plan,892        uint32_t n_tokens,893        int64_t n_stream) {894    bool res = true;895    res &= dsv4_can_reuse_tensor_1d(inp.state_pos, plan.state_pos.size());896    res &= dsv4_can_reuse_tensor_1d(inp.state_persist_src_idxs, plan.state_persist_src_idxs.size());897    res &= dsv4_can_reuse_tensor_1d(inp.state_persist_dst_idxs, plan.state_persist_dst_idxs.size());898    res &= dsv4_can_reuse_tensor_1d(inp.state_restore_src_idxs, plan.state_restore_src_idxs.size());899    res &= dsv4_can_reuse_tensor_1d(inp.state_restore_dst_idxs, plan.state_restore_dst_idxs.size());900    res &= dsv4_can_reuse_tensor_1d(inp.state_snapshot_src_idxs, plan.state_snapshot_src_idxs.size());901    res &= dsv4_can_reuse_tensor_1d(inp.state_snapshot_dst_idxs, plan.state_snapshot_dst_idxs.size());902    res &= dsv4_can_reuse_tensor_1d(inp.state_read_idxs, plan.state_read_idxs.size());903    res &= dsv4_can_reuse_tensor_1d(inp.state_write_idxs, plan.state_write_idxs.size());904    res &= dsv4_can_reuse_tensor_1d(inp.state_write_pos, plan.state_write_pos.size());905    res &= dsv4_can_reuse_kq_mask(inp.kq_mask, plan, n_tokens, n_stream);906 907    return res;908}909 910static ggml_tensor * dsv4_build_input_1d(911        ggml_context * ctx,912        ggml_type type,913        int64_t ne0,914        const std::string & name) {915    if (ne0 == 0) {916        return nullptr;917    }918 919    ggml_tensor * res = ggml_new_tensor_1d(ctx, type, ne0);920    ggml_set_input(res);921    ggml_set_name(res, name.c_str());922 923    return res;924}925 926static void dsv4_build_comp_inputs(927        ggml_context * ctx,928        llm_graph_input_dsv4::comp_input & inp,929        const llama_kv_cache_dsv4_context::comp_plan & plan,930        const char * name,931        const llama_cparams & cparams,932        int64_t n_stream) {933    inp.state_pos = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_pos.size(), std::string("dsv4_") + name + "_state_pos");934    inp.state_persist_src_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_persist_src_idxs.size(), std::string("dsv4_") + name + "_state_persist_src_idxs");935    inp.state_persist_dst_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_persist_dst_idxs.size(), std::string("dsv4_") + name + "_state_persist_dst_idxs");936    inp.state_restore_src_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_restore_src_idxs.size(), std::string("dsv4_") + name + "_state_restore_src_idxs");937    inp.state_restore_dst_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_restore_dst_idxs.size(), std::string("dsv4_") + name + "_state_restore_dst_idxs");938    inp.state_snapshot_src_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_snapshot_src_idxs.size(), std::string("dsv4_") + name + "_state_snapshot_src_idxs");939    inp.state_snapshot_dst_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_snapshot_dst_idxs.size(), std::string("dsv4_") + name + "_state_snapshot_dst_idxs");940    inp.state_read_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_read_idxs.size(), std::string("dsv4_") + name + "_state_read_idxs");941    inp.state_write_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I64, plan.state_write_idxs.size(), std::string("dsv4_") + name + "_state_write_idxs");942    inp.state_write_pos = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_write_pos.size(), std::string("dsv4_") + name + "_state_write_pos");943 944    if (plan.n_kv > 0) {945        const int64_t n_tokens = (int64_t) plan.n_visible.size();946 947        GGML_ASSERT(n_stream > 0);948        GGML_ASSERT(n_tokens%n_stream == 0);949 950        inp.kq_mask = ggml_new_tensor_4d(ctx, (strcmp(name, "lid") != 0 && cparams.flash_attn) || (strcmp(name, "lid") == 0 && cparams.fused_lid) ? GGML_TYPE_F16 : GGML_TYPE_F32, plan.n_kv, n_tokens/n_stream, 1, n_stream);951        ggml_set_input(inp.kq_mask);952        ggml_set_name(inp.kq_mask, (std::string("dsv4_") + name + "_kq_mask").c_str());953    }954}955 956void llm_graph_input_dsv4_raw::set_input(const llama_ubatch * ubatch) {957    if (self_k_idxs && self_k_idxs->buffer) {958        mctx->set_input_k_idxs(self_k_idxs);959    }960 961    if (self_kq_mask && self_kq_mask->buffer) {962        mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);963    }964 965    if (self_k_rot) {966        mctx->set_input_k_rot(self_k_rot);967    }968}969 970void llm_graph_input_dsv4::set_input(const llama_ubatch * ubatch) {971    const auto & plan_csa = mctx->get_csa_plan(*ubatch);972    const auto & plan_hca = mctx->get_hca_plan(*ubatch);973    const auto & plan_lid = mctx->get_lid_plan(*ubatch);974    const int64_t n_stream = plan_csa.n_stream;975 976    inp_raw->mctx = mctx->get_raw();977    inp_raw->set_input(ubatch);978 979    dsv4_set_comp_inputs(inp_csa, plan_csa, "csa", debug > 0, ubatch->n_tokens, n_stream);980    dsv4_set_comp_inputs(inp_hca, plan_hca, "hca", debug > 0, ubatch->n_tokens, n_stream);981    dsv4_set_comp_inputs(inp_lid, plan_lid, "lid", debug > 0, ubatch->n_tokens, n_stream);982 983    if (inp_csa.k_rot && inp_csa.k_rot->buffer) {984        mctx->get_csa()->set_input_k_rot(inp_csa.k_rot);985    }986 987    if (inp_hca.k_rot && inp_hca.k_rot->buffer) {988        mctx->get_hca()->set_input_k_rot(inp_hca.k_rot);989    }990 991    if (inp_lid.k_rot && inp_lid.k_rot->buffer) {992        mctx->get_lid()->set_input_k_rot(inp_lid.k_rot);993    }994}995 996bool llm_graph_input_dsv4::can_reuse(const llm_graph_params & params) {997    const auto * mctx = static_cast<const llama_kv_cache_dsv4_context *>(params.mctx);998 999    this->mctx = mctx;1000    inp_raw->mctx = mctx->get_raw();1001 1002    bool res = true;1003 1004    const auto & plan_csa = mctx->get_csa_plan(params.ubatch);1005    const auto & plan_hca = mctx->get_hca_plan(params.ubatch);1006    const auto & plan_lid = mctx->get_lid_plan(params.ubatch);1007    const int64_t n_stream = plan_csa.n_stream;1008 1009    const auto * raw_ctx = mctx->get_raw();1010    inp_raw->mctx = raw_ctx;1011 1012    if (inp_raw->self_k_idxs && inp_raw->self_k_idxs->buffer) {1013        res &= inp_raw->self_k_idxs->ne[0] == raw_ctx->get_n_write();1014    }1015    if (inp_raw->self_kq_mask && inp_raw->self_kq_mask->buffer) {1016        res &= dsv4_can_reuse_raw_kq_mask(inp_raw->self_kq_mask, raw_ctx, params.ubatch, n_stream);1017    }1018 1019    res &= dsv4_can_reuse_comp_input(inp_csa, plan_csa, params.ubatch.n_tokens, n_stream);1020    res &= dsv4_can_reuse_comp_input(inp_hca, plan_hca, params.ubatch.n_tokens, n_stream);1021    res &= dsv4_can_reuse_comp_input(inp_lid, plan_lid, params.ubatch.n_tokens, n_stream);1022 1023    return res;1024}1025 1026void llm_graph_input_attn_cross::set_input(const llama_ubatch * ubatch) {1027    GGML_ASSERT(cross_kq_mask);1028 1029    const int64_t n_enc    = cross_kq_mask->ne[0];1030    const int64_t n_tokens = ubatch->n_tokens;1031 1032    GGML_ASSERT(ggml_backend_buffer_is_host(cross_kq_mask->buffer));1033    GGML_ASSERT(!ubatch->equal_seqs()); // TODO: use ubatch->n_seqs instead of failing1034 1035    const auto fill_mask = [&](auto * data) {1036        using T = std::remove_reference_t<decltype(*data)>;1037        for (int i = 0; i < n_tokens; ++i) {1038            GGML_ASSERT(!cross->seq_ids_enc.empty() && "llama_encode must be called first");1039            for (int j = 0; j < n_enc; ++j) {1040                float f = -INFINITY;1041 1042                for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {1043                    const llama_seq_id seq_id = ubatch->seq_id[i][s];1044 1045                    if (cross->seq_ids_enc[j].find(seq_id) != cross->seq_ids_enc[j].end()) {1046                        f = 0.0f;1047                    }1048                }1049 1050                data[i*n_enc + j] = llama_cast<T>(f);1051            }1052        }1053    };1054 1055    if (cross_kq_mask->type == GGML_TYPE_F16) {1056        fill_mask((ggml_fp16_t *) cross_kq_mask->data);1057    } else {1058        fill_mask((float *) cross_kq_mask->data);1059    }1060}1061 1062void llm_graph_input_mem_hybrid::set_input(const llama_ubatch * ubatch) {1063    mctx->get_attn()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);1064    mctx->get_attn()->set_input_v_idxs(inp_attn->self_v_idxs, ubatch);1065 1066    mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);1067 1068    if (inp_attn->self_k_rot) {1069        mctx->get_attn()->set_input_k_rot(inp_attn->self_k_rot);1070    }1071 1072    if (inp_attn->self_v_rot) {1073        mctx->get_attn()->set_input_v_rot(inp_attn->self_v_rot);1074    }1075 1076    const int64_t n_rs = mctx->get_recr()->get_n_rs();1077 1078    if (inp_rs->s_copy) {1079        GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));1080        int32_t * data = (int32_t *) inp_rs->s_copy->data;1081 1082        // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n1083        for (uint32_t i = 0; i < n_rs; ++i) {1084            data[i] = mctx->get_recr()->s_copy(i);1085        }1086    }1087}1088 1089bool llm_graph_input_mem_hybrid::can_reuse(const llm_graph_params & params) {1090    const auto * mctx = static_cast<const llama_memory_hybrid_context *>(params.mctx);1091 1092    this->mctx = mctx;1093 1094    bool res = true;1095 1096    res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens;1097  //res &= inp_attn->self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there1098 1099    res &= can_reuse_kq_mask(inp_attn->self_kq_mask, mctx->get_attn(), params.ubatch, params.cparams);1100 1101    res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();1102 1103    res &= inp_rs->s_copy_main->ne[0]  == params.ubatch.n_seqs;1104    res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;1105 1106    res &= inp_rs->head == mctx->get_recr()->get_head();1107    res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();1108 1109    return res;1110}1111 1112// TODO: Hybrid input classes are a bit redundant.1113// Instead of creating a hybrid input, the graph can simply create 2 separate inputs.1114// Refactoring is required in the future.1115void llm_graph_input_mem_hybrid_k::set_input(const llama_ubatch * ubatch) {1116    mctx->get_attn()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);1117 1118    mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);1119 1120    const int64_t n_rs = mctx->get_recr()->get_n_rs();1121 1122    if (inp_rs->s_copy) {1123        GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));1124        int32_t * data = (int32_t *) inp_rs->s_copy->data;1125 1126        // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n1127        for (uint32_t i = 0; i < n_rs; ++i) {1128            data[i] = mctx->get_recr()->s_copy(i);1129        }1130    }1131}1132 1133bool llm_graph_input_mem_hybrid_k::can_reuse(const llm_graph_params & params) {1134    const auto * mctx = static_cast<const llama_memory_hybrid_context *>(params.mctx);1135 1136    this->mctx = mctx;1137 1138    bool res = true;1139 1140    res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens;1141 1142    res &= can_reuse_kq_mask(inp_attn->self_kq_mask, mctx->get_attn(), params.ubatch, params.cparams);1143 1144    res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();1145 1146    res &= inp_rs->s_copy_main->ne[0]  == params.ubatch.n_seqs;1147    res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;1148 1149    res &= inp_rs->head == mctx->get_recr()->get_head();1150    res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();1151 1152    return res;1153}1154 1155void llm_graph_input_mem_hybrid_iswa::set_input(const llama_ubatch * ubatch) {1156    const auto * attn_ctx = mctx->get_attn();1157 1158    // base tensors may not be allocated if there are no non-SWA attention layers1159    if (inp_attn->self_k_idxs && inp_attn->self_k_idxs->buffer) {1160        attn_ctx->get_base()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);1161        attn_ctx->get_base()->set_input_v_idxs(inp_attn->self_v_idxs, ubatch);1162    }1163 1164    if (inp_attn->self_kq_mask && inp_attn->self_kq_mask->buffer) {1165        attn_ctx->get_base()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);1166    }1167 1168    // swa tensors may not be allocated if there are no SWA attention layers1169    if (inp_attn->self_k_idxs_swa && inp_attn->self_k_idxs_swa->buffer) {1170        attn_ctx->get_swa()->set_input_k_idxs(inp_attn->self_k_idxs_swa, ubatch);1171        attn_ctx->get_swa()->set_input_v_idxs(inp_attn->self_v_idxs_swa, ubatch);1172    }1173 1174    if (inp_attn->self_kq_mask_swa && inp_attn->self_kq_mask_swa->buffer) {1175        attn_ctx->get_swa()->set_input_kq_mask(inp_attn->self_kq_mask_swa, ubatch, cparams.causal_attn);1176    }1177 1178    if (inp_attn->self_k_rot) {1179        attn_ctx->get_base()->set_input_k_rot(inp_attn->self_k_rot);1180    }1181 1182    if (inp_attn->self_v_rot) {1183        attn_ctx->get_base()->set_input_v_rot(inp_attn->self_v_rot);1184    }1185 1186    if (inp_attn->self_k_rot_swa) {1187        attn_ctx->get_swa()->set_input_k_rot(inp_attn->self_k_rot_swa);1188    }1189 1190    if (inp_attn->self_v_rot_swa) {1191        attn_ctx->get_swa()->set_input_v_rot(inp_attn->self_v_rot_swa);1192    }1193 1194    const int64_t n_rs = mctx->get_recr()->get_n_rs();1195 1196    if (inp_rs->s_copy) {1197        GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));1198        int32_t * data = (int32_t *) inp_rs->s_copy->data;1199 1200        // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n

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

Brunobkr/llama.cpp_AlgMor24_github · Team Ai