diff --git a/src/llama-memory-hybrid-idx.cpp b/src/llama-memory-hybrid-idx.cpp index 3972ce9ce293..9087b3562ae8 100644 --- a/src/llama-memory-hybrid-idx.cpp +++ b/src/llama-memory-hybrid-idx.cpp @@ -5,6 +5,7 @@ #include "llama-io.h" #include "llama-model.h" +#include "ggml-backend.h" #include #include @@ -65,7 +66,89 @@ llama_memory_hybrid_idx::llama_memory_hybrid_idx( model, hparams_idx, type_k, type_v, v_trans, offload, unified, kv_size, n_seq_max, n_pad, n_swa, swa_type, nullptr, filter_idx, nullptr, nullptr, "idx_"); - }()) {} + }()) { + // [TAG_QSA_POOLED_CACHE] one f32 row per position block per layer; single-stream memories + // only (a unified cache shares one stream; block rows are position-indexed) + if (mem_idx && mem_idx->get_n_stream() == 1) { + uint32_t ratio = 0; + for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { + // only the dense-attention layers own a QSA cache: a recurrent layer may still + // carry a nonzero ratio in the metadata, but it is not a pooled-cache layer + if (!model.hparams.is_recr(il) && model.hparams.dsv4_compress_ratios[il] > 0) { + ratio = model.hparams.dsv4_compress_ratios[il]; + break; + } + } + + const uint32_t idx_dim = model.hparams.indexer_head_size; + + if (ratio > 0 && idx_dim > 0) { + // + 1 so a partial trailing block has a slot, + 1 dustbin row for padded writes + pooled_rows = kv_size/ratio + 2; + pooled_ratio = ratio; + + // one context+buffer per device: the indexer caches of the QSA layers are spread + // across the layer-split devices, and a row written by a device that does not own + // it would travel the inter-GPU link every decode step + std::vector bufts; + std::vector> per_buf_tensors; + + for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { + // the idx cache is filtered to the dense-attention layers; get_k_storage on a + // recurrent layer is out of range even when it carries a ratio + if (model.hparams.is_recr(il) || model.hparams.dsv4_compress_ratios[il] == 0) { + continue; + } + ggml_tensor * k = mem_idx->get_k_storage((int32_t) il); + if (k == nullptr) { + continue; + } + + const ggml_backend_buffer_type_t buft = ggml_backend_buffer_get_type(k->buffer); + + size_t ci = SIZE_MAX; + for (size_t j = 0; j < bufts.size(); ++j) { + if (bufts[j] == buft) { + ci = j; + break; + } + } + if (ci == SIZE_MAX) { + ci = bufts.size(); + bufts.push_back(buft); + per_buf_tensors.emplace_back(); + + ggml_init_params ip = { + /*.mem_size =*/ 2*model.hparams.n_layer()*ggml_tensor_overhead(), + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + pooled_ctxs.emplace_back(ggml_init(ip)); + } + + ggml_tensor * t = ggml_new_tensor_2d(pooled_ctxs[ci].get(), GGML_TYPE_F32, idx_dim, pooled_rows); + ggml_format_name(t, "idx_pooled_l%u", il); + pooled_k[(int32_t) il] = t; + per_buf_tensors[ci].push_back(t); + } + + size_t total_bytes = 0; + for (size_t ci = 0; ci < bufts.size(); ++ci) { + pooled_bufs.emplace_back(ggml_backend_alloc_ctx_tensors_from_buft(pooled_ctxs[ci].get(), bufts[ci])); + GGML_ASSERT(pooled_bufs.back() && "failed to allocate the pooled indexer key cache"); + // stale rows are read (and masked); they must be finite, never uninitialized + ggml_backend_buffer_clear(pooled_bufs.back().get(), 0); + total_bytes += ggml_backend_buffer_get_size(pooled_bufs.back().get()); + } + + if (!pooled_k.empty()) { + LLAMA_LOG_INFO("%s: pooled indexer key cache, %zu layers x %u rows on %zu buffers, %.2f MiB\n", + __func__, pooled_k.size(), pooled_rows, pooled_bufs.size(), + total_bytes/1024.0/1024.0); + } + } + } +} llama_memory_context_ptr llama_memory_hybrid_idx::init_batch(llama_batch_allocr & balloc, uint32_t n_ubatch, bool embd_all) { // note: repeats llama_memory_hybrid::init_batch, as the indexer needs the attention slot infos that the base context hides @@ -146,6 +229,9 @@ void llama_memory_hybrid_idx::clear(bool data) { if (mem_idx) { mem_idx->clear(data); } + + // [TAG_QSA_POOLED_CACHE] + pooled_reset(-1); } bool llama_memory_hybrid_idx::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { @@ -158,6 +244,9 @@ bool llama_memory_hybrid_idx::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_po mem_idx->seq_rm(seq_id, p0, p1); } + // [TAG_QSA_POOLED_CACHE] + pooled_rm(seq_id, p0, p1); + return get_mem_attn()->seq_rm(seq_id, p0, p1); } @@ -167,6 +256,10 @@ void llama_memory_hybrid_idx::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_i if (mem_idx) { mem_idx->seq_cp(seq_id_src, seq_id_dst, p0, p1); } + + // [TAG_QSA_POOLED_CACHE] rows are shared in the single-stream cache; the copy's blocks + // are refilled from its own cells on its first ubatch + pooled_reset(seq_id_dst); } void llama_memory_hybrid_idx::seq_keep(llama_seq_id seq_id) { @@ -175,6 +268,11 @@ void llama_memory_hybrid_idx::seq_keep(llama_seq_id seq_id) { if (mem_idx) { mem_idx->seq_keep(seq_id); } + + // [TAG_QSA_POOLED_CACHE] only seq_id's rows survive as trusted + const int64_t keep = pooled_w.count(seq_id) ? pooled_w[seq_id] : 0; + pooled_w.clear(); + pooled_w[seq_id] = keep; } void llama_memory_hybrid_idx::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) { @@ -183,6 +281,9 @@ void llama_memory_hybrid_idx::seq_add(llama_seq_id seq_id, llama_pos p0, llama_p if (mem_idx) { mem_idx->seq_add(seq_id, p0, p1, shift); } + + // [TAG_QSA_POOLED_CACHE] shifting positions remaps every block + pooled_reset(seq_id); } void llama_memory_hybrid_idx::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) { @@ -191,6 +292,9 @@ void llama_memory_hybrid_idx::seq_div(llama_seq_id seq_id, llama_pos p0, llama_p if (mem_idx) { mem_idx->seq_div(seq_id, p0, p1, d); } + + // [TAG_QSA_POOLED_CACHE] + pooled_reset(seq_id); } std::map llama_memory_hybrid_idx::memory_breakdown() const { @@ -248,6 +352,14 @@ void llama_memory_hybrid_idx::state_read(llama_io_read_i & io, llama_seq_id seq_ throw; } + + // [TAG_QSA_POOLED_CACHE] a full restore rewrites the indexer cells with arbitrary + // content, so no pooled row can be trusted; the next ubatch refills the whole range. + // A PARTIAL_ONLY restore (speculative checkpoint replay) leaves the cells untouched + // and its rollback arrives through seq_rm, which already clamped the watermark. + if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) { + pooled_reset(seq_id); + } } void llama_memory_hybrid_idx::state_drop(llama_seq_id seq_id) { @@ -264,12 +376,50 @@ void llama_memory_hybrid_idx::state_drop(llama_seq_id seq_id) { if (mem_idx) { mem_idx->seq_rm(seq_id, -1, -1); } + + // [TAG_QSA_POOLED_CACHE] + pooled_reset(seq_id); } llama_kv_cache * llama_memory_hybrid_idx::get_mem_idx() const { return mem_idx.get(); } +ggml_tensor * llama_memory_hybrid_idx::get_pooled_k(int32_t il) const { + const auto it = pooled_k.find(il); + return it == pooled_k.end() ? nullptr : it->second; +} + +int64_t & llama_memory_hybrid_idx::pooled_valid(llama_seq_id seq_id) const { + return pooled_w[seq_id]; +} + +void llama_memory_hybrid_idx::pooled_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { + if (pooled_k.empty()) { + return; + } + + if (p0 <= 0 && p1 < 0) { + pooled_w[seq_id] = 0; + return; + } + + // blocks at or beyond the first removed position lose members; earlier rows keep their + // content (removals only ever drop the tail or a middle range, never rewrite the prefix) + const int64_t blk = pooled_ratio > 0 ? std::max(p0, 0)/pooled_ratio : 0; + + auto & w = pooled_w[seq_id]; + w = std::min(w, blk); +} + +void llama_memory_hybrid_idx::pooled_reset(llama_seq_id seq_id) { + if (seq_id < 0) { + pooled_w.clear(); + } else { + pooled_w[seq_id] = 0; + } +} + void llama_memory_hybrid_idx::set_input_qsa( ggml_tensor * cell_blk, ggml_tensor * blk_cells, @@ -277,7 +427,10 @@ void llama_memory_hybrid_idx::set_input_qsa( ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio, - bool blk_bias) const { + bool blk_bias, + ggml_tensor * dirty_cells, + ggml_tensor * dirty_pos, + ggml_tensor * dirty_rows) const { GGML_ASSERT(ratio > 0); GGML_ASSERT(get_mem_idx() != nullptr); @@ -285,18 +438,26 @@ void llama_memory_hybrid_idx::set_input_qsa( const int64_t n_kv = cell_blk->ne[0]; const int64_t n_ns = cell_blk->ne[1]; // streams in this ubatch - const int64_t n_blocks = blk_pos->ne[0]/(4*n_ns); const int64_t n_tokens = ubatch->n_tokens; const int64_t r = ratio; + // same formula as the graph; blk_pos may be null on the pooled path + const int64_t n_blocks = (n_kv + r - 1)/r; GGML_ASSERT(n_tokens % n_ns == 0); const int64_t n_tps = n_tokens/n_ns; // tokens per stream int32_t * dst_cell_blk = (int32_t *) cell_blk->data; - int32_t * dst_blk_cells = (int32_t *) blk_cells->data; - int32_t * dst_blk_pos = (int32_t *) blk_pos->data; float * dst_bias = (float *) bias->data; + // [TAG_QSA_POOLED_CACHE] the pooled path drops blk_cells/blk_pos from the graph (the dirty + // tables replace them), so they may be null here; the block map is still needed for the + // dirty fill, so it is built in local buffers either way + int32_t * dst_blk_cells = blk_cells != nullptr ? (int32_t *) blk_cells->data : nullptr; + int32_t * dst_blk_pos = blk_pos != nullptr ? (int32_t *) blk_pos->data : nullptr; + + std::vector loc_blk_cells(r*n_blocks); + std::vector loc_blk_pos(4*n_blocks); + // a block is keyed on (sequence set, index bucket): a unified cache counts every sequence // from zero, so the bucket alone would pool two sequences into one block GGML_ASSERT(r <= 64); @@ -319,17 +480,15 @@ void llama_memory_hybrid_idx::set_input_qsa( std::vector order; std::vector rank; - std::fill(dst_blk_pos, dst_blk_pos + 4*n_blocks*n_ns, 0); - for (int64_t s = 0; s < n_ns; ++s) { // ubatch index s*n_tps belongs to this stream; ask which cells array it uses const llama_seq_id seq_of_stream = ubatch->seq_id[s*n_tps][0]; const auto & cells = get_mem_idx()->get_cells(seq_of_stream); - int32_t * cur_cell_blk = dst_cell_blk + s*n_kv; - int32_t * cur_blk_cells = dst_blk_cells + s*(r*n_blocks); + int32_t * cur_cell_blk = dst_cell_blk + s*n_kv; - std::fill(cur_blk_cells, cur_blk_cells + r*n_blocks, 0); + std::fill(loc_blk_cells.begin(), loc_blk_cells.end(), 0); + std::fill(loc_blk_pos.begin(), loc_blk_pos.end(), 0); bid_idx .clear(); bid_cell .clear(); @@ -488,7 +647,7 @@ void llama_memory_hybrid_idx::set_input_qsa( } for (int64_t sec = 0; sec < 4; ++sec) { - dst_blk_pos[sec*(n_blocks*n_ns) + s*n_blocks + b] = sec_pos[sec]; + loc_blk_pos[sec*n_blocks + b] = sec_pos[sec]; } } @@ -505,12 +664,75 @@ void llama_memory_hybrid_idx::set_input_qsa( if (blk_of[j] >= 0) { const int64_t idx = ranked ? rank[j] : cells.pos_get(j); - cur_blk_cells[blk_of[j]*r + (idx%r)] = (int32_t) j; + loc_blk_cells[blk_of[j]*r + (idx%r)] = (int32_t) j; } cur_cell_blk[j] = blk_of[j] < 0 ? dead_bid : blk_of[j]; } + if (dst_blk_cells != nullptr) { + std::copy(loc_blk_cells.begin(), loc_blk_cells.end(), dst_blk_cells + s*(r*n_blocks)); + } + + if (dst_blk_pos != nullptr) { + for (int64_t sec = 0; sec < 4; ++sec) { + std::copy(loc_blk_pos.begin() + sec*n_blocks, loc_blk_pos.begin() + (sec + 1)*n_blocks, + dst_blk_pos + sec*(n_blocks*n_ns) + s*n_blocks); + } + } + + // [TAG_QSA_POOLED_CACHE] resolve which blocks the graph must (re)pool this ubatch: + // the range from the sequence's watermark to its last complete block. Complete blocks + // are immutable, so rows below the watermark stay valid; rollbacks arrive as + // seq_rm/state_read, which clamp the watermark before this runs. + if (dirty_cells != nullptr) { + GGML_ASSERT(n_ns == 1 && "the pooled cache path is single-stream only"); + + const int64_t n_dirty_max = dirty_rows->ne[0]; + const int64_t dustbin = (int64_t) get_pooled_rows() - 1; + + int32_t * dst_d_cells = (int32_t *) dirty_cells->data; + int32_t * dst_d_pos = (int32_t *) dirty_pos->data; + int64_t * dst_d_rows = (int64_t *) dirty_rows->data; + + // the bids are the complete blocks, pushed in position-block order: the last + // bid's block ends the complete range + const int64_t n_complete = n_bid > 0 ? (int64_t) bid_idx[n_bid - 1]/r + 1 : 0; + + auto & w = pooled_valid(seq_of_stream); + w = std::min(w, n_complete); + + const int64_t n_dirty = n_complete - w; + GGML_ASSERT(n_dirty <= n_dirty_max && "dirty tables sized at graph build; see qsa_pooled_n_dirty_max"); + + // position block -> bid: an incomplete block below the complete end pools nothing + // this time and keeps its stale row, masked by the bias + std::vector pb_bid(n_complete > 0 ? (size_t) n_complete : 1u, -1); + for (int32_t t = 0; t < n_bid; ++t) { + const int64_t pb = bid_idx[t]/r; + if (pb < n_complete) { + pb_bid[pb] = t; + } + } + + for (int64_t i = 0; i < n_dirty_max; ++i) { + const bool live = i < n_dirty; + const int64_t b = w + i; + const int32_t t = live && b < n_complete ? pb_bid[b] : -1; + + dst_d_rows[i] = live ? b : dustbin; + + for (int64_t sec = 0; sec < 4; ++sec) { + dst_d_pos[sec*n_dirty_max + i] = t >= 0 ? loc_blk_pos[sec*n_blocks + t] : 0; + } + for (int64_t j = 0; j < r; ++j) { + dst_d_cells[i*r + j] = t >= 0 ? loc_blk_cells[t*r + j] : 0; + } + } + + w = n_complete; + } + for (int64_t ii = 0; ii < n_tps; ++ii) { const int64_t i = s*n_tps + ii; const llama_seq_id seq_id = ubatch->seq_id[i][0]; @@ -676,8 +898,45 @@ void llama_memory_hybrid_idx_context::set_input_qsa( ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio, - bool blk_bias) const { + bool blk_bias, + ggml_tensor * dirty_cells, + ggml_tensor * dirty_pos, + ggml_tensor * dirty_rows) const { + GGML_ASSERT(mem != nullptr); + + mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, + dirty_cells, dirty_pos, dirty_rows); +} + +ggml_tensor * llama_memory_hybrid_idx_context::get_pooled_k(int32_t il) const { + return mem != nullptr && get_idx() != nullptr ? mem->get_pooled_k(il) : nullptr; +} + +uint32_t llama_memory_hybrid_idx_context::get_pooled_rows() const { + return mem != nullptr ? mem->get_pooled_rows() : 0; +} + +uint32_t llama_memory_hybrid_idx_context::qsa_pooled_n_dirty_max(const llama_ubatch & ubatch, uint32_t ratio) const { + GGML_ASSERT(ratio > 0); GGML_ASSERT(mem != nullptr); - mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias); + // the reserve pass builds worst-case graphs from a mock ubatch with no seq/pos data; + // give it the per-ubatch bound (the refill after a state load resizes on a live ubatch) + if (ubatch.seq_id == nullptr || ubatch.seq_id[0] == nullptr || ubatch.pos == nullptr) { + return (ubatch.n_tokens + ratio - 1)/ratio + 1; + } + + // single-stream memories only (get_pooled_k gates the callers); like the block tables, + // the watermark follows the first token's sequence + const llama_seq_id seq = ubatch.seq_id[0][0]; + + llama_pos q_max = -1; + for (uint32_t i = 0; i < ubatch.n_tokens; ++i) { + q_max = std::max(q_max, ubatch.pos[i]); + } + + const int64_t n_complete = (int64_t) (q_max + 1)/ratio; + const int64_t w = std::min(mem->pooled_valid(seq), n_complete); + + return (uint32_t) std::max(1, n_complete - w); } diff --git a/src/llama-memory-hybrid-idx.h b/src/llama-memory-hybrid-idx.h index 705189e7eb58..6b49e2369906 100644 --- a/src/llama-memory-hybrid-idx.h +++ b/src/llama-memory-hybrid-idx.h @@ -2,7 +2,11 @@ #include "llama-memory-hybrid.h" +#include "ggml-backend.h" + +#include #include +#include #include // @@ -85,7 +89,22 @@ class llama_memory_hybrid_idx : public llama_memory_hybrid { // the caller then adds the attention mask, the only part of the bias that varies within a block void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos, ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio, - bool blk_bias) const; + bool blk_bias, + ggml_tensor * dirty_cells = nullptr, + ggml_tensor * dirty_pos = nullptr, + ggml_tensor * dirty_rows = nullptr) const; + + // [TAG_QSA_POOLED_CACHE] cache of the indexer's block summary keys (mean-pooled, + // normalized, roped), one f32 row per position block, written by the graph via set_rows. + // Only complete blocks are scored and a complete block's members never change, so rows are + // write-once per content epoch. Validity is a per-sequence block watermark; seq_rm clamps + // it and replay recomputes the range. Rows at or beyond the watermark may hold stale but + // finite data and are masked by the -inf bias. Single-stream memories only. + ggml_tensor * get_pooled_k(int32_t il) const; // nullptr: no indexer / multi-stream + uint32_t get_pooled_rows() const { return pooled_rows; } // rows per stream, incl. trailing dustbin row + + // blocks of seq_id whose pooled rows are known valid; mutable like a cache's bookkeeping + int64_t & pooled_valid(llama_seq_id seq_id) const; private: // forget seq_id (all of it if seq_id < 0) in every cache at once, so a failed restore cannot leave the caches out of step @@ -97,6 +116,22 @@ class llama_memory_hybrid_idx : public llama_memory_hybrid { llama_hparams hparams_idx; const std::unique_ptr mem_idx; + + // [TAG_QSA_POOLED_CACHE] storage + watermarks; empty unless the model has an indexer + // one buffer per device: each layer's rows must live with that layer's indexer cache, + // or every decode step would copy the rows across the (slow) inter-GPU links + std::vector pooled_ctxs; + std::vector pooled_bufs; + std::map pooled_k; + + uint32_t pooled_rows = 0; + uint32_t pooled_ratio = 0; + + mutable std::unordered_map pooled_w; + + // clamp helpers, one per llama_memory_i operation that can invalidate rows + void pooled_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1); + void pooled_reset(llama_seq_id seq_id); // -1 resets every sequence }; class llama_memory_hybrid_idx_context : public llama_memory_hybrid_context { @@ -141,9 +176,28 @@ class llama_memory_hybrid_idx_context : public llama_memory_hybrid_context { // streams in the current slot info, the `ns` of get_k/get_v; 1 if unified uint32_t get_n_stream() const; + // [TAG_QSA_POOLED_CACHE] the dirty_* tensors are optional: when given, the fill also + // resolves which blocks must be (re)pooled this ubatch — the range from the sequence's + // watermark to its last complete block — and advances the watermark. + // dirty_cells I32 [ratio*n_dirty_max, ns] cells of each block to (re)pool, 0-padded + // dirty_pos I32 [4*n_dirty_max*ns] mrope position rows of those blocks + // dirty_rows I64 [n_dirty_max*ns] pooled-cache rows to write, dustbin-padded void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos, ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio, - bool blk_bias) const; + bool blk_bias, + ggml_tensor * dirty_cells = nullptr, + ggml_tensor * dirty_pos = nullptr, + ggml_tensor * dirty_rows = nullptr) const; + + // [TAG_QSA_POOLED_CACHE] pooled tensor for il, or nullptr when the cache is unavailable + // (no indexer, multi-stream memory, or a non-batch context) + ggml_tensor * get_pooled_k(int32_t il) const; + + uint32_t get_pooled_rows() const; + + // capacity the dirty tables need for this ubatch: completed blocks plus pending refill + // below the watermark; stable at 1 during steady decode so graph reuse holds + uint32_t qsa_pooled_n_dirty_max(const llama_ubatch & ubatch, uint32_t ratio) const; private: const llama_memory_hybrid_idx * mem = nullptr; diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp index 8ace95f73475..9fb78deb3759 100644 --- a/src/models/qwen4exp.cpp +++ b/src/models/qwen4exp.cpp @@ -478,7 +478,8 @@ class llama_model_qwen4exp::llm_graph_input_qsa : public llm_graph_input_i { void set_input(const llama_ubatch * ubatch) override { mctx->get_idx()->set_input_k_idxs(k_idxs, ubatch); - mctx->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias); + mctx->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, + dirty_cells, dirty_pos, dirty_rows); } bool can_reuse(const llm_graph_params & params) override { @@ -500,11 +501,21 @@ class llama_model_qwen4exp::llm_graph_input_qsa : public llm_graph_input_i { res &= k_idxs->ne[0] == params.ubatch.n_tokens; res &= cell_blk->ne[0] == n_kv; res &= cell_blk->ne[1] == n_stream; - res &= blk_cells->ne[0] == (int64_t) ratio*n_blocks; - res &= blk_pos->ne[0] == 4*n_blocks*n_stream; + if (blk_cells != nullptr) { + res &= blk_cells->ne[0] == (int64_t) ratio*n_blocks; + } + if (blk_pos != nullptr) { + res &= blk_pos->ne[0] == 4*n_blocks*n_stream; + } res &= bias->ne[0] == (blk_bias ? n_blocks : n_kv); res &= bias->ne[1] == params.ubatch.n_tokens/n_stream; + // [TAG_QSA_POOLED_CACHE] the dirty tables must hold this ubatch's (re)pool range; + // steady decode needs at most one block, so the capacity is stable at 1 there + if (dirty_rows != nullptr) { + res &= dirty_rows->ne[0] == (int64_t) mctx->qsa_pooled_n_dirty_max(params.ubatch, ratio); + } + return res; } @@ -515,6 +526,11 @@ class llama_model_qwen4exp::llm_graph_input_qsa : public llm_graph_input_i { ggml_tensor * blk_pos = nullptr; // I32 [4*n_blocks*n_stream] ggml_tensor * bias = nullptr; // F32 [n_blocks or n_kv, n_tokens/n_stream, n_stream] + // [TAG_QSA_POOLED_CACHE] present only when the pooled cache path is active + ggml_tensor * dirty_cells = nullptr; // I32 [ratio*n_dirty_max, 1] + ggml_tensor * dirty_pos = nullptr; // I32 [4*n_dirty_max] + ggml_tensor * dirty_rows = nullptr; // I64 [n_dirty_max] + const llama_memory_hybrid_idx_context * mctx; const uint32_t ratio; @@ -564,15 +580,36 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k( qsa->k_idxs = mctx_idx->build_input_k_idxs(ctx0, ubatch); qsa->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, n_stream); - qsa->blk_cells = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, r*n_blocks, n_stream); - qsa->blk_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*n_blocks*n_stream); qsa->bias = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, blk_bias ? n_blocks : n_kv, n_tps, n_stream); ggml_set_input(qsa->cell_blk); - ggml_set_input(qsa->blk_cells); - ggml_set_input(qsa->blk_pos); ggml_set_input(qsa->bias); + // [TAG_QSA_POOLED_CACHE] complete blocks' summaries are cached; per ubatch only the + // freshly completed blocks (plus any pending refill after a full state load) are + // pooled/normed/roped, and the score reads the cache. The full-recompute tables are + // then dead graph inputs, so they are not created at all (an unreferenced input is + // never allocated, and filling it would write through a null pointer). + // Kill switch for A/B testing. + if (mctx_hyb->get_pooled_k(il) != nullptr && n_stream == 1 && + getenv("LLAMA_QSA_NO_POOLED_CACHE") == nullptr) { + const int64_t n_dirty_max = mctx_hyb->qsa_pooled_n_dirty_max(ubatch, (uint32_t) r); + + qsa->dirty_cells = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, r*n_dirty_max, 1); + qsa->dirty_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*n_dirty_max); + qsa->dirty_rows = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, n_dirty_max); + + ggml_set_input(qsa->dirty_cells); + ggml_set_input(qsa->dirty_pos); + ggml_set_input(qsa->dirty_rows); + } else { + qsa->blk_cells = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, r*n_blocks, n_stream); + qsa->blk_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*n_blocks*n_stream); + + ggml_set_input(qsa->blk_cells); + ggml_set_input(qsa->blk_pos); + } + inp = qsa.get(); res->add_input(std::move(qsa)); qsa_inps.emplace((uint32_t) r, inp); @@ -589,32 +626,75 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k( ggml_tensor * k_all = mctx_idx->get_k(ctx0, il); k_all = ggml_view_3d(ctx0, k_all, idx_dim, n_kv, n_stream, k_all->nb[2], k_all->nb[3], 0); - // gathers per stream: blk_cells row s indexes stream s's own cells - ggml_tensor * members = ggml_get_rows(ctx0, k_all, inp->blk_cells); - members = ggml_reshape_4d(ctx0, members, idx_dim, r, n_blocks, n_stream); - - // mean over the block members; r is small, so summing slices beats a transpose plus sum_rows ggml_tensor * pooled = nullptr; - for (int64_t i = 0; i < r; ++i) { - ggml_tensor * slice = ggml_cont(ctx0, - ggml_view_3d(ctx0, members, idx_dim, n_blocks, n_stream, - members->nb[2], members->nb[3], i*members->nb[1])); - pooled = pooled ? ggml_add(ctx0, pooled, slice) : slice; - } - pooled = ggml_scale(ctx0, pooled, 1.0f/(float) r); - cb(pooled, "indexer_k_pooled", il); - - // count blocks along ne1: rms_norm launches gridDim.y = ne2, capped at 65535, and 262144/4 = 65536 - pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, n_blocks*n_stream, 1); - pooled = build_norm(pooled, model.layers[il].index_k_norm, nullptr, LLM_NORM_RMS, il); - // rope wants [n_dims, n_head, n_tokens]: lay every stream's blocks flat, split after. - pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, 1, n_blocks*n_stream); - pooled = ggml_rope_multi(ctx0, pooled, inp->blk_pos, nullptr, - n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, - ext_factor, attn_factor, beta_fast, beta_slow); - pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, n_blocks, n_stream); - cb(pooled, "indexer_k", il); + if (inp->dirty_rows != nullptr) { + // [TAG_QSA_POOLED_CACHE] pool only this ubatch's dirty blocks and scatter them into + // the cache; the score then reads the cache. Rows of incomplete blocks hold stale + // (finite) data and are masked by the -inf bias, exactly like the garbage partial + // pools of the full recompute below. + ggml_tensor * store = mctx_hyb->get_pooled_k(il); + GGML_ASSERT(store != nullptr); + + const int64_t n_dirty_max = inp->dirty_rows->ne[0]; + + ggml_tensor * members = ggml_get_rows(ctx0, k_all, inp->dirty_cells); + members = ggml_reshape_4d(ctx0, members, idx_dim, r, n_dirty_max, 1); + + ggml_tensor * fresh = nullptr; + for (int64_t i = 0; i < r; ++i) { + ggml_tensor * slice = ggml_cont(ctx0, + ggml_view_3d(ctx0, members, idx_dim, n_dirty_max, 1, + members->nb[2], members->nb[3], i*members->nb[1])); + fresh = fresh ? ggml_add(ctx0, fresh, slice) : slice; + } + fresh = ggml_scale(ctx0, fresh, 1.0f/(float) r); + cb(fresh, "indexer_k_pooled", il); + + // count dirty blocks along ne1: rms_norm launches gridDim.y = ne2, capped at 65535 + fresh = ggml_reshape_3d(ctx0, fresh, idx_dim, n_dirty_max, 1); + fresh = build_norm(fresh, model.layers[il].index_k_norm, nullptr, LLM_NORM_RMS, il); + + fresh = ggml_reshape_3d(ctx0, fresh, idx_dim, 1, n_dirty_max); + fresh = ggml_rope_multi(ctx0, fresh, inp->dirty_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + fresh = ggml_reshape_2d(ctx0, fresh, idx_dim, n_dirty_max); + + ggml_tensor * store_view = ggml_view_2d(ctx0, store, + idx_dim, store->ne[1], store->nb[1], 0); + ggml_build_forward_expand(gf, ggml_set_rows(ctx0, store_view, fresh, inp->dirty_rows)); + + pooled = ggml_view_3d(ctx0, store, idx_dim, n_blocks, 1, + store->nb[1], store->nb[1]*n_blocks, 0); + cb(pooled, "indexer_k", il); + } else { + // gathers per stream: blk_cells row s indexes stream s's own cells + ggml_tensor * members = ggml_get_rows(ctx0, k_all, inp->blk_cells); + members = ggml_reshape_4d(ctx0, members, idx_dim, r, n_blocks, n_stream); + + // mean over the block members; r is small, so summing slices beats a transpose plus sum_rows + for (int64_t i = 0; i < r; ++i) { + ggml_tensor * slice = ggml_cont(ctx0, + ggml_view_3d(ctx0, members, idx_dim, n_blocks, n_stream, + members->nb[2], members->nb[3], i*members->nb[1])); + pooled = pooled ? ggml_add(ctx0, pooled, slice) : slice; + } + pooled = ggml_scale(ctx0, pooled, 1.0f/(float) r); + cb(pooled, "indexer_k_pooled", il); + + // count blocks along ne1: rms_norm launches gridDim.y = ne2, capped at 65535, and 262144/4 = 65536 + pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, n_blocks*n_stream, 1); + pooled = build_norm(pooled, model.layers[il].index_k_norm, nullptr, LLM_NORM_RMS, il); + + // rope wants [n_dims, n_head, n_tokens]: lay every stream's blocks flat, split after. + pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, 1, n_blocks*n_stream); + pooled = ggml_rope_multi(ctx0, pooled, inp->blk_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, n_blocks, n_stream); + cb(pooled, "indexer_k", il); + } ggml_tensor * q = build_lora_mm(model.layers[il].index_q_proj, cur); q = ggml_reshape_3d(ctx0, q, idx_dim, n_idx_h, n_tokens);