CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

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

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