From 66e0c17ee1741fef493312e17fe60a5d2cf5f7d5 Mon Sep 17 00:00:00 2001 From: Aman Gupta Date: Thu, 1 Oct 2026 19:13:27 +0800 Subject: [PATCH] llama: fix qwen4exp (#29751) * llama: fix qwen4exp * qwen4exp: keep kq_mask input the same shape --- src/llama-hparams.h | 4 + src/llama-memory-hybrid-idx.cpp | 504 ++++++++------------------------ src/llama-memory-hybrid-idx.h | 32 +- src/models/models.h | 18 +- src/models/qwen4exp.cpp | 382 ++++++++++++------------ 5 files changed, 356 insertions(+), 584 deletions(-) diff --git a/src/llama-hparams.h b/src/llama-hparams.h index 8248add7d4..756007e1f7 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -284,6 +284,10 @@ struct llama_hparams { uint32_t indexer_top_k = 0; uint32_t indexer_kpool = 0; // k-pool size bool indexer_kpool_select_tail = true; + // head-size slots per cached indexer row, the last one holds the pooled key + uint32_t indexer_kpool_row = 3; + // pools are consecutive cells in sequence order, not runs of consecutive positions + bool indexer_kpool_by_order = false; // MSA uint32_t indexer_block_size = 0; uint32_t indexer_local_blocks = 0; diff --git a/src/llama-memory-hybrid-idx.cpp b/src/llama-memory-hybrid-idx.cpp index 32b0422536..de64a37000 100644 --- a/src/llama-memory-hybrid-idx.cpp +++ b/src/llama-memory-hybrid-idx.cpp @@ -53,8 +53,9 @@ llama_memory_hybrid_idx::llama_memory_hybrid_idx( mem_idx(filter_idx == nullptr ? nullptr : [&] { // MQA with a single key head of indexer_head_size, as llama_kv_cache_dsa shapes its own std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1); - // The glm5 next indexer caches key, gate and pooled values per token - hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size * (model.hparams.indexer_kpool > 0 ? 3 : 1); + // a k-pool indexer caches its per-token rows and the pooled key side by side + // (glm5-next: key | gate | pooled, qwen4exp: key | pooled) + hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size * (model.hparams.indexer_kpool > 0 ? model.hparams.indexer_kpool_row : 1); // the cached indexer keys are raw, rotation happens after pooling at read time, so a // K-shift must not rotate them while the stream copies in the same update still apply @@ -331,328 +332,6 @@ llama_kv_cache * llama_memory_hybrid_idx::get_mem_idx() const { return mem_idx.get(); } -void llama_memory_hybrid_idx::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, - bool causal_attn) const { - GGML_ASSERT(ratio > 0); - GGML_ASSERT(get_mem_idx() != nullptr); - - GGML_ASSERT(ggml_backend_buffer_is_host(cell_blk->buffer)); - - 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; - - 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; - - // 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); - const uint64_t slots_full = r == 64 ? ~uint64_t(0) : ((uint64_t(1) << r) - 1); - - // TODO: this runs per ubatch and is O(n_kv) per stream, about 865 us at 33k context. the cost - // is the per-cell scan rather than these allocations, so hoisting them buys nothing - std::vector blk_of(n_kv); - std::vector cell_grp(n_kv); - std::vector grp_head(n_blocks); - std::vector grp_next; - std::vector grp_first; - std::vector grp_slot0; - std::vector grp_slots; - std::vector grp_bid; - std::vector bid_idx; - std::vector bid_cell; - std::vector bid_slot0; - - 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); - - std::fill(cur_blk_cells, cur_blk_cells + r*n_blocks, 0); - - bid_idx .clear(); - bid_cell .clear(); - bid_slot0.clear(); - - int n_seq_present = 0; - - for (int sq = 0; sq < LLAMA_MAX_SEQ && n_seq_present < 2; ++sq) { - if (cells.seq_pos_min(sq) >= 0) { - n_seq_present++; - } - } - - const bool one_seq = n_seq_present <= 1; - - // a cell no block covers needs its own -inf, which a per-block bias cannot carry - // every cache path keeps the position below the cell window, so this stays false - bool oor = false; - - bool dup = false; - - bool ranked = false; - - auto group_cells = [&]() { - // -1 means no usable block: an incomplete or short group cannot be pooled - std::fill(blk_of.begin(), blk_of.end(), -1); - std::fill(cell_grp.begin(), cell_grp.end(), -1); - std::fill(grp_head.begin(), grp_head.end(), -1); - - grp_next .clear(); - grp_first.clear(); - grp_slot0.clear(); - grp_slots.clear(); - grp_bid .clear(); - - oor = false; - dup = false; - - for (int64_t j = 0; j < n_kv; ++j) { - if (cells.is_empty(j)) { - continue; - } - - const int64_t idx = ranked ? rank[j] : cells.pos_get(j); - const int64_t pb = idx/r; - - if (pb >= n_blocks) { - oor = true; - continue; - } - - int32_t g = -1; - - for (int32_t c = grp_head[pb]; c >= 0; c = grp_next[c]) { - if (one_seq || cells.seq_get_all((uint32_t) grp_first[c]) == cells.seq_get_all((uint32_t) j)) { - g = c; - break; - } - } - - if (g < 0) { - g = (int32_t) grp_first.size(); - - grp_next .push_back(grp_head[pb]); - grp_first.push_back((int32_t) j); - grp_slot0.push_back(-1); - grp_slots.push_back(0); - grp_bid .push_back(-1); - - grp_head[pb] = g; - } - - const uint64_t bit = uint64_t(1) << (idx%r); - - dup |= (grp_slots[g] & bit) != 0; - - cell_grp[j] = g; - grp_slots[g] |= bit; - - if (idx%r == 0) { - grp_slot0[g] = (int32_t) j; - } - } - }; - - group_cells(); - - // mrope repeats one position across an image, so rank cells instead of using the position - if (dup && ubatch->is_pos_2d() && one_seq) { - order.clear(); - order.reserve(n_kv); - - for (int64_t j = 0; j < n_kv; ++j) { - if (!cells.is_empty(j)) { - order.push_back((int32_t) j); - } - } - - // same total order the mrope causal mask uses: pos, then ext.y, then ext.x - std::sort(order.begin(), order.end(), [&cells](int32_t a, int32_t b) { - const llama_pos pa = cells.pos_get(a); - const llama_pos pb = cells.pos_get(b); - - if (pa != pb) { - return pa < pb; - } - - const auto & ea = cells.ext_get(a); - - return cells.ext_get(b).is_2d_gt(ea.x, ea.y); - }); - - rank.assign(n_kv, -1); - - for (int64_t k = 0; k < (int64_t) order.size(); ++k) { - rank[order[k]] = (int32_t) k; - } - - ranked = true; - - group_cells(); - } - - GGML_ASSERT((!blk_bias || !oor) && "qsa: cell position runs past the cell window"); - - int32_t n_bid = 0; - - for (int64_t pb = 0; pb < n_blocks; ++pb) { - for (int32_t g = grp_head[pb]; g >= 0; g = grp_next[g]) { - if (grp_slots[g] != slots_full) { - continue; - } - - grp_bid[g] = n_bid++; - - bid_idx .push_back((int32_t) (pb*r)); - bid_cell .push_back(grp_first[g]); - bid_slot0.push_back(grp_slot0[g]); - } - } - - GGML_ASSERT(n_bid <= n_blocks); - - for (int32_t b = 0; b < n_bid; ++b) { - int32_t sec_pos[4] = { bid_idx[b], bid_idx[b], bid_idx[b], bid_idx[b] }; - - if (ranked) { - const int32_t c = bid_slot0[b]; - const llama_pos p = cells.pos_get(c); - const auto & e = cells.ext_get(c); - - sec_pos[0] = p; - sec_pos[1] = e.y; - sec_pos[2] = e.x; - sec_pos[3] = p; - } - - for (int64_t sec = 0; sec < 4; ++sec) { - dst_blk_pos[sec*(n_blocks*n_ns) + s*n_blocks + b] = sec_pos[sec]; - } - } - - // unpooled cells all point at one spare block. a spare block exists only when some - // cell is unpooled: n_bid == n_blocks means every cell sits in a full block. - const bool have_dead = n_bid < n_blocks; - const int32_t dead_bid = have_dead ? n_bid : n_blocks - 1; - - for (int64_t j = 0; j < n_kv; ++j) { - const int32_t g = cell_grp[j]; - - blk_of[j] = g < 0 ? -1 : grp_bid[g]; - - 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; - } - - cur_cell_blk[j] = blk_of[j] < 0 ? dead_bid : blk_of[j]; - } - - 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]; - - int64_t q = ubatch->pos[i]; - - if (ranked) { - const llama_pos qt = ubatch->pos[i]; - const llama_pos qy = ubatch->pos[i + n_tokens]; - const llama_pos qx = ubatch->pos[i + n_tokens*2]; - - int64_t lo = 0; - int64_t hi = (int64_t) order.size(); - - while (lo < hi) { - const int64_t mid = (lo + hi)/2; - const int32_t c = order[mid]; - const llama_pos pc = cells.pos_get(c); - - if (pc < qt || (pc == qt && !cells.ext_get(c).is_2d_gt(qx, qy))) { - lo = mid + 1; - } else { - hi = mid; - } - } - - q = lo - 1; - } - - // the tail is an incomplete block and is always visible, as in the reference - const int64_t tail_start = (q + 1)/r*r; - - if (blk_bias) { - // a block sits wholly inside or outside the tail, so one value covers it - // the caller adds the attention mask, which drops empty, foreign and, when causal, future cells - float * cur_blk_bias = dst_bias + i*n_blocks; - - for (int64_t b = 0; b < n_blocks; ++b) { - if (b >= n_bid || !cells.seq_has((uint32_t) bid_cell[b], seq_id)) { - cur_blk_bias[b] = -INFINITY; - continue; - } - - // finite, so it can never meet a -inf and produce a nan - cur_blk_bias[b] = (causal_attn && bid_idx[b] >= tail_start) ? 1e9f : 0.0f; - } - - // the spare block holds the unpooled cells, which are the incomplete tail, so - // it gets the tail value. it must stay finite: a sequence with fewer than - // `ratio` cells owns no full block, and a row of -inf only gives a nan. - if (have_dead) { - cur_blk_bias[dead_bid] = 1e9f; - } - - continue; - } - - float * cur_bias = dst_bias + i*n_kv; - - for (int64_t j = 0; j < n_kv; ++j) { - float v = -INFINITY; - - if (!cells.is_empty(j) && cells.seq_has(j, seq_id)) { - const int64_t idx = ranked ? rank[j] : cells.pos_get(j); - - if (!causal_attn) { - // every visible block competes on score and the unpooled cells are always selected - v = blk_of[j] < 0 ? 1e9f : 0.0f; - } else if (idx <= q) { - // finite, so it can never meet a -inf and produce a nan - v = idx >= tail_start ? 1e9f : (blk_of[j] < 0 ? -INFINITY : 0.0f); - } - } - - cur_bias[j] = v; - } - } - } -} - // // llama_memory_hybrid_idx_context // @@ -697,6 +376,7 @@ struct llama_memory_hybrid_idx_context::kpool_state { uint32_t n_pool_real = 0; uint32_t n_new = 0; + uint32_t n_new_g = 1; // graph size of the new pool list, stable across decode steps bool cache_safe = true; }; @@ -707,6 +387,13 @@ uint32_t kpool_pad(uint32_t n_pool) { return std::max(64u, GGML_PAD(n_pool + 1, 64u)); } +// Rank of (pos, cell) in a sequence's cells sorted by position then cell, or -1 when absent. +// In order mode the rank alone places a token: cells sharing a position (M-RoPE images) have distinct ranks. +int64_t kpool_rank(const std::vector> & cells, llama_pos pos, uint32_t cell) { + auto it = std::lower_bound(cells.begin(), cells.end(), std::make_pair(pos, cell)); + return it != cells.end() && it->second == cell && it->first == pos ? it - cells.begin() : -1; +} + } llama_memory_hybrid_idx::~llama_memory_hybrid_idx() = default; @@ -787,24 +474,31 @@ const llama_memory_hybrid_idx::kpool_layout & llama_memory_hybrid_idx::kpool_lay // Pools start at the first valid token size_t j = sq.j_next; - while (j + kpool <= sq.cells.size()) { - const llama_pos p0 = sq.cells[j].first; - if ((p0 - sq.pos_min) % (llama_pos) kpool != 0) { - ++j; - continue; - } - bool ok = true; - for (uint32_t k = 1; k < kpool; ++k) { - if (sq.cells[j + k].first != p0 + (llama_pos) k) { - ok = false; - break; - } - } - if (ok) { + if (hparams_idx.indexer_kpool_by_order) { + // consecutive cells in sequence order, whatever their positions + for (; j + kpool <= sq.cells.size(); j += kpool) { sq.pools.push_back((uint32_t) j); - j += kpool; - } else { - ++j; + } + } else { + while (j + kpool <= sq.cells.size()) { + const llama_pos p0 = sq.cells[j].first; + if ((p0 - sq.pos_min) % (llama_pos) kpool != 0) { + ++j; + continue; + } + bool ok = true; + for (uint32_t k = 1; k < kpool; ++k) { + if (sq.cells[j + k].first != p0 + (llama_pos) k) { + ok = false; + break; + } + } + if (ok) { + sq.pools.push_back((uint32_t) j); + j += kpool; + } else { + ++j; + } } } sq.j_next = j; @@ -835,7 +529,8 @@ llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context(llama_memory_hy const uint64_t n_pool_max = uint64_t(idx->get_size() / mem->get_kpool()) * idx->get_n_seq_max(); GGML_ASSERT(n_pool_max <= UINT32_MAX - 64); st.n_pool_real = std::max(st.n_pool_real, uint32_t(n_pool_max)); - st.n_new = st.n_pool_real; + st.n_new = st.n_pool_real; + st.n_new_g = std::max(st.n_new, 1u); kpool_st = std::make_unique(std::move(st)); i_kpool = 0; } @@ -860,6 +555,7 @@ llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context( llama_memory_hybrid_context(mem, std::move(sinfos_attn), ubatches), mem(mem), ns_ubatch(llama_memory_hybrid_idx_ns(sinfos_idx)), + sinfos_kpool(mem->get_mem_idx() != nullptr && mem->get_kpool() > 0 && mem->get_kpool_by_order() ? sinfos_idx : slot_info_vec_t()), ctx_idx(mem->get_mem_idx() == nullptr ? nullptr : new llama_kv_cache_context(mem->get_mem_idx(), std::move(sinfos_idx), ubatches)) { // Sequence edits force the touched positions to re-pool. @@ -918,29 +614,17 @@ uint32_t llama_memory_hybrid_idx_context::get_n_stream() const { return ns_ubatch[i_cur]; } -void llama_memory_hybrid_idx_context::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, - bool causal_attn) const { - GGML_ASSERT(mem != nullptr); - - mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, causal_attn); -} - llama_memory_hybrid_idx_context::kpool_access::kpool_access(ggml_context * ctx, ggml_tensor * k, int64_t n_embd) : ctx(ctx) { - GGML_ASSERT(k->ne[0] == 3*n_embd); + // rows are the per-token part (glm5-next: key | gate, qwen4exp: key), then the pooled key + const int64_t n_tok = k->ne[0] - n_embd; + GGML_ASSERT(n_tok > 0 && n_tok % n_embd == 0); const int64_t n_cells = k->ne[1]*k->ne[2]; // Pool indices can refer to other streams. Revisit these full-storage views if that changes: // https://github.com/ggml-org/llama.cpp/pull/27773#discussion_r4130905603 - key_gate = ggml_view_2d(ctx, k, 2*n_embd, n_cells, k->nb[1], 0); - pooled = ggml_view_2d(ctx, k, n_embd, n_cells, k->nb[1], ggml_row_size(k->type, 2*n_embd)); + key_gate = ggml_view_2d(ctx, k, n_tok, n_cells, k->nb[1], 0); + pooled = ggml_view_2d(ctx, k, n_embd, n_cells, k->nb[1], ggml_row_size(k->type, n_tok)); } ggml_tensor * llama_memory_hybrid_idx_context::kpool_access::gather_key_gate(ggml_tensor * idxs) const { @@ -972,7 +656,7 @@ ggml_tensor * llama_memory_hybrid_idx_context::gather_mla_rows( return ggml_get_rows(ctx, rows, ggml_reshape_1d(ctx, idxs, n_rows)); } -// k-pool DSA indexer (glm5-next) +// k-pool DSA indexer (glm5-next, qwen4exp QSA) // Sizes only, used by the full cache context so get_n_kpool() works during graph reserve. llama_memory_hybrid_idx_context::kpool_state llama_memory_hybrid_idx_context::kpool_build_sizes() const { @@ -1034,7 +718,7 @@ void llama_memory_hybrid_idx_context::kpool_build_state(const llama_ubatch & uba } auto first = std::lower_bound(sq.pools.begin(), sq.pools.end(), stale_from, - [&](uint32_t j, llama_pos p) { return sq.cells[j].first + (llama_pos) kpool <= p; }); + [&](uint32_t j, llama_pos p) { return sq.cells[j + kpool - 1].first < p; }); for (auto it = first; it != sq.pools.end(); ++it) { mark(pool_start[s] + (uint32_t) (it - sq.pools.begin())); } @@ -1043,26 +727,45 @@ void llama_memory_hybrid_idx_context::kpool_build_state(const llama_ubatch & uba if (!st.cache_safe) { std::fill(st.is_new.begin(), st.is_new.end(), st.generation); - st.n_new = st.n_pool_real; + st.n_new = st.n_pool_real; + st.n_new_g = std::max(st.n_new, 1u); return; } + // in order mode a token's cell gives its rank, and the rank its pool: positions cannot, as an image shares one + const bool by_order = mem->get_kpool_by_order(); + const auto * sinfo = by_order ? &sinfos_kpool[i_cur] : nullptr; + const uint32_t n_tps = by_order ? (uint32_t) sinfo->size() : 0; + for (uint32_t i = 0; i < ubatch.n_tokens; ++i) { const llama_pos p = ubatch.pos[i]; for (int32_t k = 0; k < ubatch.n_seq_id[i]; ++k) { const llama_seq_id s = ubatch.seq_id[i][k]; const auto & sq = lay.seqs[s]; + if (by_order) { + const int64_t r = kpool_rank(sq.cells, p, sinfo->idxs[i / n_tps][i % n_tps]); + GGML_ASSERT(r >= 0); + if ((size_t) r / kpool < sq.pools.size()) { + mark(pool_start[s] + (uint32_t) (r / kpool)); + } + continue; + } auto it = std::upper_bound(sq.pools.begin(), sq.pools.end(), p, [&](llama_pos pos, uint32_t j) { return pos < sq.cells[j].first; }); if (it == sq.pools.begin()) { continue; } --it; - if (p < sq.cells[*it].first + (llama_pos) kpool) { + if (p <= sq.cells[*it + kpool - 1].first) { mark(pool_start[s] + (uint32_t) (it - sq.pools.begin())); } } } + + // a ubatch touches at most t_s/kpool + 1 pools of a sequence with t_s tokens: pad to that bound so the + // graph keeps its shape as the count moves, e.g. between 0 and n_seq while several sequences decode + const uint32_t bound = ubatch.n_tokens/kpool + ubatch.n_seqs_unq; + st.n_new_g = std::max({st.n_new, 1u, std::min(bound, kpool_pad(st.n_pool_real) - 1)}); } const llama_memory_hybrid_idx_context::kpool_state & llama_memory_hybrid_idx_context::kpool_cur() const { @@ -1076,7 +779,7 @@ uint32_t llama_memory_hybrid_idx_context::get_n_kpool() const { } uint32_t llama_memory_hybrid_idx_context::get_n_kpool_new() const { - return kpool_cur().n_new; + return kpool_cur().n_new_g; } bool llama_memory_hybrid_idx_context::get_kpool_cache_safe() const { @@ -1085,7 +788,7 @@ bool llama_memory_hybrid_idx_context::get_kpool_cache_safe() const { void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, ggml_tensor * pool_idxs, ggml_tensor * pool_mask, ggml_tensor * tail_idxs, ggml_tensor * gather_mask, bool gather, ggml_tensor * new_pool_idxs, ggml_tensor * new_pool_rep, - const llama_ubatch * ubatch) const { + const llama_ubatch * ubatch, ggml_tensor * new_pool_pos) const { GGML_ASSERT(mem != nullptr && mem->get_mem_idx() != nullptr); GGML_ASSERT(ggml_backend_buffer_is_host(pool_cells->buffer)); GGML_ASSERT(ggml_backend_buffer_is_host(pool_idxs->buffer)); @@ -1101,8 +804,10 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, const uint32_t n_tokens = ubatch->n_tokens; const uint32_t n_pool = (uint32_t) pool_cells->ne[0]; const uint32_t n_new = st.n_new; - // the graph always pools at least one entry, see build_inp_kpool - const uint32_t n_new_g = std::max(n_new, 1u); + // the graph always pools at least one entry, padded to a stable bound, see kpool_build_state + const uint32_t n_new_g = st.n_new_g; + + const bool by_order = mem->get_kpool_by_order(); GGML_ASSERT(n_pool == kpool_pad(st.n_pool_real)); GGML_ASSERT(st.is_new.size() == st.n_pool_real); @@ -1116,6 +821,10 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, GGML_ASSERT(ggml_backend_buffer_is_host(new_pool_rep->buffer)); GGML_ASSERT(new_pool_rep->ne[0] == (int64_t) n_new_g); } + if (new_pool_pos != nullptr) { + GGML_ASSERT(ggml_backend_buffer_is_host(new_pool_pos->buffer)); + GGML_ASSERT(new_pool_pos->ne[0] == 4*(int64_t) n_new_g); + } const uint32_t kv_size = mem->get_mem_idx()->get_size(); const uint32_t n_stream_kv = mem->get_mem_idx()->get_n_stream(); @@ -1142,6 +851,19 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, dummy_cell = gcell(sq, it->second); } + // in order mode a token sees the pools and the tail up to its own rank in the sequence, which its cell pins down + std::vector rank; + if (by_order) { + const auto & sinfo = sinfos_kpool[i_cur]; + const uint32_t n_tps = (uint32_t) sinfo.size(); + + rank.resize(n_tokens); + for (uint32_t i = 0; i < n_tokens; ++i) { + rank[i] = kpool_rank(lay.seqs[ubatch->seq_id[i][0]].cells, ubatch->pos[i], sinfo.idxs[i / n_tps][i % n_tps]); + GGML_ASSERT(rank[i] >= 0); + } + } + // Gather maps padding to a real cell and masks it separately. const int32_t sentinel = gather ? (int32_t) dummy_cell : (int32_t) n_kv; @@ -1168,6 +890,11 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, int32_t * pidx = (int32_t *) pool_idxs->data; int32_t * nidx = (int32_t *) new_pool_idxs->data; int64_t * nrep = new_pool_rep != nullptr ? (int64_t *) new_pool_rep->data : nullptr; + int32_t * npos = new_pool_pos != nullptr ? (int32_t *) new_pool_pos->data : nullptr; + + if (npos != nullptr) { + std::fill(npos, npos + 4*n_new_g, 0); + } uint32_t i_new = 0; for (llama_seq_id s = 0; s < LLAMA_MAX_SEQ; ++s) { @@ -1198,6 +925,15 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, if (nrep != nullptr) { nrep[i_new] = gcell(sq, rep); } + if (npos != nullptr) { + // a pooled key is rotated to the M-RoPE position of its first member + const uint32_t c = sq.cells[j].second; + const auto & e = mem->get_mem_idx()->get_cells(s).ext_get(c); + npos[0*n_new_g + i_new] = sq.cells[j].first; + npos[1*n_new_g + i_new] = e.y; + npos[2*n_new_g + i_new] = e.x; + npos[3*n_new_g + i_new] = sq.cells[j].first; + } ++i_new; } @@ -1206,14 +942,25 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, } GGML_ASSERT(i_new == n_new); - // A ubatch that completes no pool re-pools the cell of its first token. That cell cannot belong to - // a complete pool here, else the pool would be marked new, so the write never touches a cached key. - if (n_new == 0) { + // Padded entries re-pool a cell whose pooled slot is never read. With no new pool that is the cell of the + // first token: it cannot belong to a complete pool, else the pool would be marked new. Otherwise it is the + // first member of a pool, which is never a pool's rep. + int64_t pad_cell = dummy_cell; + if (n_new > 0) { + for (llama_seq_id s = 0; s < LLAMA_MAX_SEQ; ++s) { + const auto & sq = lay.seqs[s]; + if (!sq.pools.empty()) { + pad_cell = gcell(sq, sq.cells[sq.pools[0]].second); + break; + } + } + } + for (uint32_t i = n_new; i < n_new_g; ++i) { for (uint32_t k = 0; k < kpool; ++k) { - nidx[k] = (int32_t) dummy_cell; + nidx[(size_t) i*kpool + k] = (int32_t) pad_cell; } if (nrep != nullptr) { - nrep[0] = dummy_cell; + nrep[i] = pad_cell; } } @@ -1240,7 +987,8 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, const uint32_t p0 = seq_pool_start[s]; const uint32_t p1 = p0 + (uint32_t) lay.seqs[s].pools.size(); - const uint32_t nv = (uint32_t) (std::upper_bound(pool_end.begin() + p0, pool_end.begin() + p1, p) - (pool_end.begin() + p0)); + const uint32_t nv = by_order ? std::min(p1 - p0, (uint32_t) ((rank[i] + 1)/kpool)) : + (uint32_t) (std::upper_bound(pool_end.begin() + p0, pool_end.begin() + p1, p) - (pool_end.begin() + p0)); std::fill(row + p0, row + p0 + nv, keep); // Finite visible pools occupy the first min(nv, n_top) ranked slots. @@ -1264,12 +1012,18 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, const llama_pos p = ubatch->pos[i]; const auto & sq = lay.seqs[s]; - const uint32_t n_tail = (uint32_t) ((p - sq.pos_min + 1) % (llama_pos) kpool); + const uint32_t n_tail = by_order ? + (uint32_t) ((rank[i] + 1) % kpool) : + (uint32_t) ((p - sq.pos_min + 1) % (llama_pos) kpool); for (uint32_t k = 0; k < kpool - 1; ++k) { int32_t cell = sentinel; bool real = false; - if (k < n_tail) { + if (k < n_tail && by_order) { + const uint32_t c = sq.cells[rank[i] - k].second; + cell = (int32_t) (gather ? gcell(sq, c) : (int64_t) c); + real = true; + } else if (k < n_tail) { const llama_pos pt = p - (llama_pos) k; auto it = std::lower_bound(sq.cells.begin(), sq.cells.end(), std::make_pair(pt, 0u)); if (it != sq.cells.end() && it->first == pt) { diff --git a/src/llama-memory-hybrid-idx.h b/src/llama-memory-hybrid-idx.h index 66953cacf4..b954d9f7a0 100644 --- a/src/llama-memory-hybrid-idx.h +++ b/src/llama-memory-hybrid-idx.h @@ -80,23 +80,13 @@ public: llama_kv_cache * get_mem_idx() const; // nullptr when the model carries no indexer - // block-compressed sparse attention (qwen4exp QSA) over the cells of the indexer cache. - // Blocks cut the position line, not the cell array, so no caller assumes a contiguous layout: - // cell_blk I32 [n_kv, ns] block each cell belongs to - // blk_cells I32 [ratio*n_blocks, ns] cells making up each block - // blk_pos I32 [4*n_blocks*ns] mrope position rows of each block's first token - // bias F32 [n_kv, n_tokens/ns, ns] -inf where invisible, large where always visible - // blk_bias asks for the bias per block instead: [n_blocks, n_tokens/ns, ns] - // the caller then adds the attention mask, the only part of the bias that varies within a block - // causal_attn selects the rule: causal forces the query's own block on, non-causal lets every visible block compete on score - 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, bool causal_attn) const; - // The model's indexer pool size. uint32_t get_kpool() const { return hparams_idx.indexer_kpool; } - // Which cells of a sequence make up which pool of kpool consecutive positions. + // Whether pools are kpool consecutive cells in sequence order (qwen4exp) instead of kpool consecutive positions. + bool get_kpool_by_order() const { return hparams_idx.indexer_kpool_by_order; } + + // Which cells of a sequence make up which pool of kpool consecutive positions (or cells, in order mode). // It is kept here because it outlives the batch: pools are fixed by the positions relative to the // sequence's first one, so a ubatch only ever appends to it. Sequence edits drop it, see mem_idx_stale. struct kpool_layout; @@ -203,18 +193,16 @@ public: // streams in the current slot info, the `ns` of get_k/get_v; 1 if unified uint32_t get_n_stream() const; - // glm5-next, complete pools of kpool consecutive positions per sequence, scored as whole pools. + // glm5-next and qwen4exp, complete pools of kpool cells per sequence, scored as whole pools. uint32_t get_n_kpool () const; // Padded pool count, where the last pool is always unused. - uint32_t get_n_kpool_new() const; // Exact count of pools completed by the current ubatch. + uint32_t get_n_kpool_new() const; // Pools to re-pool this ubatch, padded to a stable bound, never below 1. bool get_kpool_cache_safe() const; kpool_access get_kpool_access(ggml_context * ctx, int32_t il, int64_t n_embd) const; ggml_tensor * gather_mla_rows(ggml_context * ctx, ggml_tensor * idxs, int64_t n_rows, int64_t n_embd, int32_t il) const; + // new_pool_pos (I32 [4*n_new]): M-RoPE position of each new pool's first member, for pooled keys rotated at pooling time void set_input_kpool(ggml_tensor * pool_cells, ggml_tensor * pool_idxs, ggml_tensor * pool_mask, ggml_tensor * tail_idxs, ggml_tensor * gather_mask, bool gather, ggml_tensor * new_pool_idxs, ggml_tensor * new_pool_rep, - const llama_ubatch * ubatch) const; - 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, bool causal_attn) const; + const llama_ubatch * ubatch, ggml_tensor * new_pool_pos = nullptr) const; private: llama_memory_hybrid_idx * mem = nullptr; @@ -223,6 +211,10 @@ private: // declared first, so it is initialised while sinfos_idx is still intact const std::vector ns_ubatch; + // the indexer cells of each ubatch, kept for pools in cache order (qwen4exp): token s*n + i of ubatch u + // sits in cell idxs[s][i] of stream strm[s] of sinfos_kpool[u], and several cells can share a position + const slot_info_vec_t sinfos_kpool; + // null unless the model has an indexer const llama_memory_context_ptr ctx_idx; diff --git a/src/models/models.h b/src/models/models.h index 0e7d59a73b..898d22f6b6 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -2388,7 +2388,7 @@ struct llama_model_qwen35 : public llama_model_base { struct llama_model_qwen4exp : public llama_model_base { llama_model_qwen4exp(const struct llama_model_params & params) : llama_model_base(params) {} - class llm_graph_input_qsa; + class llm_graph_input_kpool; void load_arch_hparams(llama_model_loader & ml) override; void load_arch_tensors(llama_model_loader & ml) override; @@ -2415,28 +2415,30 @@ struct llama_model_qwen4exp : public llama_model_base { ggml_tensor * build_layer_attn( llm_graph_input_attn_kv * inp_attn, const llama_memory_hybrid_idx_context * mctx_hyb, + llm_graph_input_kpool * inp_kpool, ggml_tensor * cur, ggml_tensor * inp_pos, int * sections, int il); - // dense self-attention restricted to the cells that top_k names + // dense self-attention over the cells the QSA mask keeps ggml_tensor * build_attn_qsa( llm_graph_input_attn_kv * inp, ggml_tensor * q_cur, ggml_tensor * k_cur, ggml_tensor * v_cur, - ggml_tensor * top_k, + ggml_tensor * sel, + int64_t n_sel, float kq_scale, int il); - // the QSA cache layout inputs do not depend on the layer, only on its compress ratio, - // so the layers sharing a ratio share one input set - std::map qsa_inps; + // the QSA layers share one set of k-pool inputs, see llama_memory_hybrid_idx + llm_graph_input_kpool * build_inp_kpool(const llama_memory_hybrid_idx_context * mctx_hyb); - // QSA: token indices this layer's queries may attend to, or nullptr for dense - ggml_tensor * build_qsa_top_k( + // QSA: the additive mask [n_kv, n_tokens] of the top blocks and the tail, kq_mask included + ggml_tensor * build_qsa_sel( const llama_memory_hybrid_idx_context * mctx_hyb, + llm_graph_input_kpool * inp_kpool, ggml_tensor * cur, ggml_tensor * inp_pos, ggml_tensor * kq_mask, diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp index 168bc52945..768ade04d1 100644 --- a/src/models/qwen4exp.cpp +++ b/src/models/qwen4exp.cpp @@ -64,6 +64,27 @@ void llama_model_qwen4exp::load_arch_hparams(llama_model_loader & ml) { qwen4exp_require_nonzero(ml, LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k); ml.get_key_or_arr(LLM_KV_ATTENTION_COMPRESS_RATIOS, hparams.dsv4_compress_ratios, hparams.n_layer_all, false); + // QSA pools the indexer keys of blocks of compress_ratio cells, one block size for the whole model + hparams.indexer_kpool = 0; + for (uint32_t il = 0; il < hparams.n_layer(); ++il) { + const uint32_t r = hparams.dsv4_compress_ratios[il]; + if (r == 0) { + continue; + } + if (hparams.indexer_kpool != 0 && r != hparams.indexer_kpool) { + throw std::runtime_error(format("QSA layers must share one compress ratio, got %u and %u", hparams.indexer_kpool, r)); + } + hparams.indexer_kpool = r; + } + if (hparams.indexer_kpool == 1 || (hparams.indexer_kpool > 0 && hparams.indexer_top_k % hparams.indexer_kpool != 0)) { + throw std::runtime_error(format("QSA needs a compress ratio above 1 that divides the budget, got %u and %u", + hparams.indexer_kpool, hparams.indexer_top_k)); + } + // the reference groups the visible tokens in cache order and always keeps the tail + hparams.indexer_kpool_row = 2; // raw key | pooled key + hparams.indexer_kpool_by_order = true; + hparams.indexer_kpool_select_tail = true; + // PLE n-gram hash embeddings; if the key group is absent every field stays zero hparams.is_ple_impl.reset(); hparams.ple_n_heads = 0; @@ -378,6 +399,13 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa "the indexer cache must track the attention cache cell for cell"); } + // the QSA layers share one set of k-pool inputs + // the CUDA lightning indexer takes 32 or 64 heads, QSA has a few, so it scores with plain ops + llm_graph_input_kpool * inp_kpool = nullptr; + if (mctx_idx && hparams.indexer_kpool > 0) { + inp_kpool = build_inp_kpool(mctx_hyb); + } + ggml_tensor * inp_pos = build_inp_pos(); ggml_tensor * inp_out_ids = build_inp_out_ids(); @@ -416,7 +444,7 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa if (hparams.is_recr(il)) { cur = build_layer_attn_linear(inp->get_recr(), cur, il); } else { - cur = build_layer_attn(inp->get_attn(), mctx_hyb, cur, inp_pos, sections, il); + cur = build_layer_attn(inp->get_attn(), mctx_hyb, inp_kpool, cur, inp_pos, sections, il); } if (il == n_layer - 1 && inp_out_ids) { @@ -490,17 +518,16 @@ ggml_tensor * llama_model_qwen4exp::graph::build_norm_gated( return ggml_mul(ctx0, normalized, gated); } -// QSA attends to a budget of whole blocks of compress_ratio tokens, plus the incomplete tail -// one mean-pooled indexer key scores each block; set_input resolves the cache layout -class llama_model_qwen4exp::llm_graph_input_qsa : public llm_graph_input_i { +// QSA k-pool inputs, shared by the QSA layers: blocks of compress_ratio cells in sequence order, see llama_memory_hybrid_idx +class llama_model_qwen4exp::llm_graph_input_kpool : public llm_graph_input_i { public: - llm_graph_input_qsa(const llama_memory_hybrid_idx_context * mctx, uint32_t ratio, bool blk_bias, bool causal_attn) : - mctx(mctx), ratio(ratio), blk_bias(blk_bias), causal_attn(causal_attn) {} - virtual ~llm_graph_input_qsa() = default; + llm_graph_input_kpool(const llama_memory_hybrid_idx_context * mctx, uint32_t kpool) : mctx(mctx), kpool(kpool) {} + virtual ~llm_graph_input_kpool() = default; 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, causal_attn); + mctx->set_input_kpool(pool_cells, pool_idxs, pool_mask, tail_idxs, nullptr, false, new_pool_idxs, new_pool_rep, + ubatch, new_pool_pos); } bool can_reuse(const llm_graph_params & params) override { @@ -511,44 +538,85 @@ public: return false; } - const int64_t n_kv = idx->get_n_kv(); - const int64_t n_stream = mctx->get_n_stream(); - const int64_t n_blocks = (n_kv + ratio - 1)/ratio; - bool res = true; - res &= params.ubatch.n_tokens % n_stream == 0; - - 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; - res &= bias->ne[0] == (blk_bias ? n_blocks : n_kv); - res &= bias->ne[1] == params.ubatch.n_tokens/n_stream; + res &= k_idxs->ne[0] == params.ubatch.n_tokens; + res &= pool_cells->ne[0] == mctx->get_n_kpool(); + res &= pool_mask->ne[1] == params.ubatch.n_tokens; + res &= tail_idxs->ne[1] == params.ubatch.n_tokens; + // the scatter mask shape follows n_kv + res &= n_kv == idx->get_n_kv(); + res &= n_new == mctx->get_n_kpool_new(); + res &= cache_safe == mctx->get_kpool_cache_safe(); return res; } - // per stream: a cell index names a different token in each stream - ggml_tensor * k_idxs = nullptr; // I32 [n_tokens] - ggml_tensor * cell_blk = nullptr; // I32 [n_kv, n_stream] - ggml_tensor * blk_cells = nullptr; // I32 [ratio*n_blocks, n_stream] - 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] + ggml_tensor * k_idxs = nullptr; // I64 [n_tokens] + ggml_tensor * pool_cells = nullptr; // I32 [n_pool] cell caching each block's pooled key + ggml_tensor * pool_idxs = nullptr; // I32 [kpool, n_pool] member cells per block, n_kv sentinel for the padded blocks + ggml_tensor * pool_mask = nullptr; // F32 [n_pool, n_tokens] + ggml_tensor * tail_idxs = nullptr; // I32 [kpool - 1, n_tokens] + ggml_tensor * new_pool_idxs = nullptr; // I32 [kpool, n_new] members of the blocks to re-pool this ubatch + ggml_tensor * new_pool_rep = nullptr; // I64 [n_new] cell to write each new pooled key into + ggml_tensor * new_pool_pos = nullptr; // I32 [4*n_new] M-RoPE position of each new block's first member const llama_memory_hybrid_idx_context * mctx; - const uint32_t ratio; - - // the per-cell half of the bias is the attention mask, so only the per-block half is uploaded - const bool blk_bias; - - // this is fixed for the graph's lifetime, as causal_attn is part of the reuse key (llm_graph_params::allow_reuse) - const bool causal_attn; + const uint32_t kpool; + uint32_t n_new = 0; // padded to a stable bound, never below 1 + uint32_t n_sel = 0; + uint32_t n_kv = 0; + bool cache_safe = true; }; -ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k( +llama_model_qwen4exp::llm_graph_input_kpool * llama_model_qwen4exp::graph::build_inp_kpool(const llama_memory_hybrid_idx_context * mctx_hyb) { + const auto * mctx_idx = mctx_hyb->get_idx(); + GGML_ASSERT(mctx_idx != nullptr); + + const uint32_t kpool = hparams.indexer_kpool; + const uint32_t n_pool = mctx_hyb->get_n_kpool(); + + auto inp = std::make_unique(mctx_hyb, kpool); + + inp->k_idxs = mctx_idx->build_input_k_idxs(ctx0, ubatch); + inp->pool_cells = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_pool); + inp->pool_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool, n_pool); + inp->pool_mask = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_pool, n_tokens); + inp->tail_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool - 1, n_tokens); + ggml_set_input(inp->pool_cells); + ggml_set_input(inp->pool_idxs); + ggml_set_input(inp->pool_mask); + ggml_set_input(inp->tail_idxs); + + // set_input fills them all, so keep them allocated even when no op reads them + ggml_build_forward_expand(gf, inp->pool_cells); + ggml_build_forward_expand(gf, inp->pool_idxs); + ggml_build_forward_expand(gf, inp->pool_mask); + ggml_build_forward_expand(gf, inp->tail_idxs); + + inp->n_kv = mctx_idx->get_n_kv(); + inp->n_new = mctx_hyb->get_n_kpool_new(); + inp->cache_safe = mctx_hyb->get_kpool_cache_safe(); + // the top blocks plus the tail + inp->n_sel = kpool*std::min(n_pool, hparams.indexer_top_k / kpool) + kpool - 1; + + inp->new_pool_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool, inp->n_new); + ggml_set_input(inp->new_pool_idxs); + if (inp->cache_safe) { + inp->new_pool_rep = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, inp->n_new); + ggml_set_input(inp->new_pool_rep); + } + inp->new_pool_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*inp->n_new); + ggml_set_input(inp->new_pool_pos); + + return (llm_graph_input_kpool *) res->add_input(std::move(inp)); +} + +// QSA attends to the top blocks of compress_ratio cells plus the incomplete tail, like the glm5-next k-pool indexer +// a block is scored by one pooled key: the mean of its raw indexer keys, normed and rotated to its first member +ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel( const llama_memory_hybrid_idx_context * mctx_hyb, + llm_graph_input_kpool * inp_kpool, ggml_tensor * cur, ggml_tensor * inp_pos, ggml_tensor * kq_mask, @@ -556,89 +624,57 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k( int il) { const llama_kv_cache_context * mctx_idx = mctx_hyb->get_idx(); - const int64_t idx_dim = hparams.indexer_head_size; - const int64_t n_idx_h = hparams.indexer_n_head; - const int64_t r = hparams.dsv4_compress_ratios[il]; - const int64_t n_kv = mctx_idx->get_n_kv(); + const int64_t idx_dim = hparams.indexer_head_size; + const int64_t n_idx_h = hparams.indexer_n_head; + const int64_t kpool = inp_kpool->kpool; + const int64_t n_pool = inp_kpool->pool_cells->ne[0]; + const int64_t n_new = inp_kpool->n_new; - GGML_ASSERT(r > 0); + GGML_ASSERT(hparams.dsv4_compress_ratios[il] == kpool); - const int64_t n_blocks = (n_kv + r - 1)/r; - - // build_attn_qsa and the KQ mask need the tokens to divide evenly across the streams - const int64_t n_stream = mctx_hyb->get_n_stream(); - GGML_ASSERT(n_tokens % n_stream == 0); - const int64_t n_tps = n_tokens/n_stream; - - // only the "which block is visible" half of the bias varies per block - // the rest is the visible/not test the attention mask already carries, so upload the per-block half only: 1/ratio of the cells - // alibi writes distances instead of a mask, so it opts out - // the mask also holds an mrope rule for the query's own position, but only 2d image positions can differ there - const bool blk_bias = kq_mask != nullptr && - kq_mask->ne[0] == n_kv && kq_mask->ne[1] == n_tps && kq_mask->ne[3] == n_stream && - !hparams.use_alibi; - - // nothing above depends on the layer, so the layers sharing a ratio share one input set - llm_graph_input_qsa * inp = nullptr; - - const auto it = qsa_inps.find((uint32_t) r); - if (it != qsa_inps.end()) { - inp = it->second; - } else { - auto qsa = std::make_unique(mctx_hyb, (uint32_t) r, blk_bias, cparams.causal_attn); - - 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); - - inp = qsa.get(); - res->add_input(std::move(qsa)); - qsa_inps.emplace((uint32_t) r, inp); - } - - // cached indexer keys are raw: pooling precedes norm and rotation, so apply neither + // cache rows store raw key | pooled key: pooling precedes norm and rotation, so the raw key gets neither ggml_tensor * k_raw = build_lora_mm(model.layers[il].index_k_proj, cur); - k_raw = ggml_reshape_3d(ctx0, k_raw, idx_dim, 1, n_tokens); cb(k_raw, "indexer_k_raw", il); - ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, k_raw, inp->k_idxs, il)); + ggml_tensor * pzero = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, idx_dim, n_tokens), 0.0f); + ggml_tensor * packed = ggml_reshape_3d(ctx0, ggml_concat(ctx0, k_raw, pzero, 0), 2*idx_dim, 1, n_tokens); + ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, packed, inp_kpool->k_idxs, il)); - // one key head, so rows are contiguous. get_k gives [idx_dim, n_head_kv, n_kv, n_stream]. - 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); + // the raw keys and the persistent pooled slots, see llama_memory_hybrid_idx::mem_idx_stale + auto kpool_cache = mctx_hyb->get_kpool_access(ctx0, il, idx_dim); - // 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); + // pool only the blocks this ubatch completes or regroups + ggml_tensor * rows = kpool_cache.gather_key_gate(ggml_reshape_1d(ctx0, inp_kpool->new_pool_idxs, kpool*n_new)); + rows = ggml_reshape_3d(ctx0, rows, idx_dim, kpool, n_new); - // 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; + // mean over the members; kpool is small, so summing slices beats a transpose plus sum_rows + ggml_tensor * pooled_new = nullptr; + for (int64_t i = 0; i < kpool; ++i) { + ggml_tensor * slice = ggml_view_2d(ctx0, rows, idx_dim, n_new, rows->nb[2], i*rows->nb[1]); + pooled_new = pooled_new ? ggml_add(ctx0, pooled_new, slice) : ggml_cont(ctx0, slice); } - pooled = ggml_scale(ctx0, pooled, 1.0f/(float) r); - cb(pooled, "indexer_k_pooled", il); + pooled_new = ggml_scale(ctx0, pooled_new, 1.0f/(float) kpool); + pooled_new = build_norm(pooled_new, model.layers[il].index_k_norm, nullptr, LLM_NORM_RMS, 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, + pooled_new = ggml_reshape_3d(ctx0, pooled_new, idx_dim, 1, n_new); + pooled_new = ggml_rope_multi(ctx0, pooled_new, inp_kpool->new_pool_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); + pooled_new = ggml_reshape_2d(ctx0, pooled_new, idx_dim, n_new); + cb(pooled_new, "indexer_pool_k_new", il); + + ggml_tensor * pooled = nullptr; + if (inp_kpool->cache_safe) { + // write before the pool gather + ggml_build_forward_expand(gf, kpool_cache.scatter_pooled(pooled_new, inp_kpool->new_pool_rep)); + pooled = kpool_cache.gather_pooled(inp_kpool->pool_cells); + } else { + // shared cells re-pool every pool, in layout order + GGML_ASSERT(n_new < n_pool); + ggml_tensor * pad = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, idx_dim, n_pool - n_new), 0.0f); + pooled = ggml_concat(ctx0, pooled_new, pad, 1); + } + pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, 1, n_pool); cb(pooled, "indexer_k", il); ggml_tensor * q = build_lora_mm(model.layers[il].index_q_proj, cur); @@ -649,67 +685,73 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k( ext_factor, attn_factor, beta_fast, beta_slow); cb(q, "indexer_q", il); - // rectify each head dot product before the sum, as in the DeepSeek lightning indexer - // mul_mat matches ne[2], so the queries of stream s only meet the blocks of stream s - ggml_tensor * score = ggml_mul_mat(ctx0, pooled, - ggml_reshape_3d(ctx0, q, idx_dim, n_idx_h*n_tps, n_stream)); - score = ggml_reshape_4d(ctx0, score, n_blocks, n_idx_h, n_tps, n_stream); - score = ggml_relu(ctx0, score); + // the reference sums the rectified head scores unweighted, scaled by 1/sqrt(head_dim) + // one product for all heads, then the heads are summed as slices, so nothing is transposed + ggml_tensor * kq = ggml_mul_mat(ctx0, + ggml_reshape_2d(ctx0, pooled, idx_dim, n_pool), + ggml_reshape_2d(ctx0, q, idx_dim, n_idx_h*n_tokens)); // [n_pool, n_idx_h*n_tokens] + kq = ggml_relu(ctx0, ggml_reshape_3d(ctx0, kq, n_pool, n_idx_h, n_tokens)); - // the heads sit side by side on ne[1] and there are only a few of them - ggml_tensor * summed = nullptr; + ggml_tensor * score = nullptr; for (int64_t h = 0; h < n_idx_h; ++h) { - ggml_tensor * slice = ggml_view_3d(ctx0, score, n_blocks, n_tps, n_stream, - score->nb[2], score->nb[3], h*score->nb[1]); - summed = summed ? ggml_add(ctx0, summed, slice) : ggml_cont(ctx0, slice); + ggml_tensor * slice = ggml_view_2d(ctx0, kq, n_pool, n_tokens, kq->nb[2], h*kq->nb[1]); + score = score ? ggml_add(ctx0, score, slice) : ggml_cont(ctx0, slice); } - - score = summed; + score = ggml_scale(ctx0, score, 1.0f/sqrtf((float) idx_dim)); + score = ggml_add(ctx0, score, inp_kpool->pool_mask); // [n_pool, n_tokens] cb(score, "indexer_score", il); - // one value per block, so it is cheaper to bias here than after the cells are expanded - if (blk_bias) { - score = ggml_add(ctx0, score, inp->bias); - } - - // every token of a block gets the block score; the budget is whole blocks, so top-k cuts on a block boundary - ggml_tensor * expanded = ggml_get_rows(ctx0, - ggml_cont(ctx0, ggml_permute(ctx0, score, 1, 0, 2, 3)), inp->cell_blk); - expanded = ggml_cont(ctx0, ggml_permute(ctx0, expanded, 1, 0, 2, 3)); - - if (blk_bias) { - // flash attention keeps the mask in f16; the scores are f32 - ggml_tensor * mask = kq_mask->type == GGML_TYPE_F32 ? kq_mask : ggml_cast(ctx0, kq_mask, GGML_TYPE_F32); - expanded = ggml_add(ctx0, expanded, ggml_reshape_3d(ctx0, mask, n_kv, n_tps, n_stream)); - } else { - expanded = ggml_add(ctx0, expanded, inp->bias); - } - cb(expanded, "indexer_score_tokens", il); - - // the reference returns indexer_top_k + compress_ratio - 1: whole blocks plus the tail - const int64_t width = std::min(n_kv, (int64_t) hparams.indexer_top_k + r - 1); - - ggml_tensor * top_k = ggml_cont(ctx0, ggml_top_k(ctx0, expanded, width)); - - // build_attn_qsa reads [n_top_k, n_batch, 1, n_stream], matching the KQ mask. - top_k = ggml_reshape_4d(ctx0, top_k, width, n_tps, 1, n_stream); + const int64_t n_top_pool = std::min(n_pool, hparams.indexer_top_k / kpool); + ggml_tensor * top_k = ggml_top_k(ctx0, score, n_top_pool); // [n_top_pool, n_tokens], unordered cb(top_k, "indexer_top_k", il); - return top_k; + // the top blocks, then the incomplete tail with n_kv for missing cells + ggml_tensor * sel_idx = ggml_get_rows(ctx0, inp_kpool->pool_idxs, + ggml_reshape_1d(ctx0, top_k, n_top_pool*n_tokens)); // [kpool, n_top_pool*n_tokens] + sel_idx = ggml_reshape_2d(ctx0, sel_idx, kpool*n_top_pool, n_tokens); + sel_idx = ggml_concat(ctx0, sel_idx, inp_kpool->tail_idxs, 0); + const int64_t n_sel = sel_idx->ne[0]; + GGML_ASSERT(n_sel == inp_kpool->n_sel); + + // scatter zeros for the selected cells into an all -inf row, the extra row n_kv takes the sentinels + // seeding from sel_idx ties the scatter storage lifetime to this layer + const int64_t n_kv = inp_kpool->n_kv; + + ggml_tensor * seed = ggml_cast(ctx0, ggml_view_1d(ctx0, sel_idx, 1, 0), GGML_TYPE_F32); + + ggml_tensor * mask_seed = kq_mask->type == GGML_TYPE_F32 ? seed : ggml_cast(ctx0, seed, kq_mask->type); + mask_seed = ggml_fill(ctx0, mask_seed, -INFINITY); + ggml_tensor * mask_all = ggml_repeat_4d(ctx0, mask_seed, 1, n_kv + 1, n_tokens, 1); + mask_all = ggml_reshape_3d(ctx0, mask_all, 1, n_kv + 1, n_tokens); + + ggml_tensor * zero_seed = ggml_fill(ctx0, seed, 0.0f); + ggml_tensor * zeros = ggml_repeat_4d(ctx0, zero_seed, 1, n_sel, n_tokens, 1); + zeros = ggml_reshape_3d(ctx0, zeros, 1, n_sel, n_tokens); + + ggml_tensor * sel = ggml_set_rows(ctx0, mask_all, zeros, ggml_reshape_3d(ctx0, sel_idx, n_sel, n_tokens, 1)); + + GGML_ASSERT(kq_mask->ne[0] == n_kv && kq_mask->ne[1]*kq_mask->ne[2]*kq_mask->ne[3] == n_tokens); + const size_t row = sel->nb[2]; + sel = ggml_view_4d(ctx0, sel, n_kv, kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3], + row, row*kq_mask->ne[1], row*kq_mask->ne[1]*kq_mask->ne[2], 0); + sel = ggml_add(ctx0, sel, kq_mask); + cb(sel, "indexer_sel", il); + + return sel; } -// Dense GQA self-attention restricted to the cells that top_k names. -// The mask build below copies the MLA sparse path in llm_graph_context::build_attn. +// Dense GQA self-attention over the cells that the QSA mask keeps. ggml_tensor * llama_model_qwen4exp::graph::build_attn_qsa( llm_graph_input_attn_kv * inp, ggml_tensor * q_cur, ggml_tensor * k_cur, ggml_tensor * v_cur, - ggml_tensor * top_k, + ggml_tensor * sel, + int64_t n_sel, float kq_scale, int il) { // rotate q/k/v before they reach a quantized cache, as the dense path does. the indexer - // has already scored with its own query in build_qsa_top_k, so top_k is unaffected. + // has already scored with its own query in build_qsa_sel, so the selection is unaffected. if (inp->self_k_rot) { q_cur = llama_mul_mat_hadamard(ctx0, q_cur, inp->self_k_rot); k_cur = llama_mul_mat_hadamard(ctx0, k_cur, inp->self_k_rot); @@ -737,39 +779,16 @@ ggml_tensor * llama_model_qwen4exp::graph::build_attn_qsa( ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, v_cur, v_idxs, il)); } + // the selection mask already carries the causal mask ggml_tensor * kq_mask = inp->get_kq_mask(); - - // prepare new kq mask - starts filled with -INFINITY - ggml_tensor * kq_mask_all = ggml_fill(ctx0, kq_mask, -INFINITY); - - // reshape KQ mask into tensor with rows of size 1: - // [n_kv, n_batch, 1, n_stream] -> [1, n_kv, n_batch, n_stream] - kq_mask_all = ggml_view_4d(ctx0, kq_mask_all, 1, kq_mask_all->ne[0], kq_mask_all->ne[1], kq_mask_all->ne[3], kq_mask_all->nb[0], kq_mask_all->nb[1], kq_mask_all->nb[2], 0); - - // reshape top_k indices: [n_top_k, n_batch, 1, n_stream] -> [n_top_k, n_batch, n_stream, 1] - ggml_tensor * top_k_3d = ggml_view_4d(ctx0, top_k, top_k->ne[0], top_k->ne[1], top_k->ne[3], 1, top_k->nb[1], top_k->nb[2], top_k->ne[3]*top_k->nb[3], 0); - - // prepare zero-filled tensor with rows of size 1: [1, n_top_k, n_batch, n_stream] - // this will be our source of zero values for unmasking top k mask elements - ggml_tensor * zeros = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, top_k_3d->ne[0], top_k_3d->ne[1], top_k_3d->ne[2]); - zeros = ggml_fill(ctx0, zeros, 0.0f); - - // modify KQ mask by unmasking elements that are in top_k indices - // ggml_set_rows([1, n_kv, n_batch, n_stream], [1, n_top_k, n_batch, n_stream], [n_top_k, n_batch, n_stream, 1]) - ggml_tensor * kq_mask_top_k = ggml_set_rows(ctx0, kq_mask_all, zeros, top_k_3d); - - // reshape to restore the original shape of KQ mask: - // [1, n_kv, n_batch, n_stream] -> [n_kv, n_batch, 1, n_stream] - kq_mask_top_k = ggml_view_4d(ctx0, kq_mask_top_k, kq_mask_top_k->ne[1], kq_mask_top_k->ne[2], 1, kq_mask_top_k->ne[3], kq_mask_top_k->nb[2], kq_mask_top_k->nb[3], kq_mask_top_k->nb[3], 0); - - // combine with the original kq mask - kq_mask_top_k = ggml_add(ctx0, kq_mask_top_k, kq_mask); + ggml_tensor * mask = ggml_reshape_4d(ctx0, sel, kq_mask->ne[0], kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3]); + cb(mask, "kq_mask_qsa", il); ggml_tensor * q = q_cur; ggml_tensor * k = mctx_cur->get_k(ctx0, il); ggml_tensor * v = mctx_cur->get_v(ctx0, il); - ggml_tensor * cur = build_attn_mha(q, k, v, nullptr, kq_mask_top_k, nullptr, nullptr, top_k->ne[0], kq_scale, il); + ggml_tensor * cur = build_attn_mha(q, k, v, nullptr, mask, nullptr, nullptr, n_sel, kq_scale, il); cb(cur, "kqv_out", il); // the rotation is its own inverse, so undo it on the value side of the output @@ -783,6 +802,7 @@ ggml_tensor * llama_model_qwen4exp::graph::build_attn_qsa( ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn( llm_graph_input_attn_kv * inp, const llama_memory_hybrid_idx_context * mctx_hyb, + llm_graph_input_kpool * inp_kpool, ggml_tensor * cur, ggml_tensor * inp_pos, int * sections, @@ -791,9 +811,9 @@ ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn( GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); // indexer reads the same block input as q/k/v; no cache or no ratio means dense - const bool qsa = mctx_hyb->get_idx() != nullptr && hparams.dsv4_compress_ratios[il] > 0; + const bool qsa = inp_kpool != nullptr && hparams.dsv4_compress_ratios[il] > 0; - ggml_tensor * top_k = qsa ? build_qsa_top_k(mctx_hyb, cur, inp_pos, inp->get_kq_mask(), sections, il) : nullptr; + ggml_tensor * sel = qsa ? build_qsa_sel(mctx_hyb, inp_kpool, cur, inp_pos, inp->get_kq_mask(), sections, il) : nullptr; // Qwen3Next uses a single Q projection that outputs query + gate ggml_tensor * Qcur_full = build_lora_mm(model.layers[il].wq, cur, model.layers[il].wq_s); // [ (n_embd_head * 2) * n_head, n_tokens ] @@ -845,8 +865,8 @@ ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn( const float kq_scale = hparams.f_attention_scale == 0.0f ? 1.0f / sqrtf(float(n_embd_head)) : hparams.f_attention_scale; - if (top_k) { - cur = build_attn_qsa(inp, Qcur, Kcur, Vcur, top_k, kq_scale, il); + if (sel) { + cur = build_attn_qsa(inp, Qcur, Kcur, Vcur, sel, inp_kpool->n_sel, kq_scale, il); } else { cur = build_attn(inp, nullptr, nullptr, nullptr,