Commit 66e0c17ee for llama.cpp
commit 66e0c17ee1741fef493312e17fe60a5d2cf5f7d5
Author: Aman Gupta <amangupta052@gmail.com>
Date: Thu Oct 1 19:13:27 2026 +0800
llama: fix qwen4exp (#29751)
* llama: fix qwen4exp
* qwen4exp: keep kq_mask input the same shape
diff --git a/src/llama-hparams.h b/src/llama-hparams.h
index 8248add7d..756007e1f 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 32b042253..de64a3700 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<int32_t> blk_of(n_kv);
- std::vector<int32_t> cell_grp(n_kv);
- std::vector<int32_t> grp_head(n_blocks);
- std::vector<int32_t> grp_next;
- std::vector<int32_t> grp_first;
- std::vector<int32_t> grp_slot0;
- std::vector<uint64_t> grp_slots;
- std::vector<int32_t> grp_bid;
- std::vector<int32_t> bid_idx;
- std::vector<int32_t> bid_cell;
- std::vector<int32_t> bid_slot0;
-
- std::vector<int32_t> order;
- std::vector<int32_t> 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<uint32_t>(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<std::pair<llama_pos, uint32_t>> & 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;
+ 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);
}
- bool ok = true;
- for (uint32_t k = 1; k < kpool; ++k) {
- if (sq.cells[j + k].first != p0 + (llama_pos) k) {
- ok = false;
- break;
+ } 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;
}
- }
- 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<kpool_state>(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<int64_t> 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 66953cacf..b954d9f7a 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<uint32_t> 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 0e7d59a73..898d22f6b 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<uint32_t, llm_graph_input_qsa *> 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 168bc5294..768ade04d 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;
+ 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;
+};
- // the per-cell half of the bias is the attention mask, so only the per-block half is uploaded
- const bool blk_bias;
+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<llm_graph_input_kpool>(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<uint32_t>(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);
- // 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;
-};
+ return (llm_graph_input_kpool *) res->add_input(std::move(inp));
+}
-ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
+// 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();
-
- GGML_ASSERT(r > 0);
-
- 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;
+ 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;
- // 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;
+ GGML_ASSERT(hparams.dsv4_compress_ratios[il] == kpool);
- // 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<llm_graph_input_qsa>(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);
- }
+ const int64_t n_top_pool = std::min<int64_t>(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);
- // 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));
+ // 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);
- 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);
+ // 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;
- // the reference returns indexer_top_k + compress_ratio - 1: whole blocks plus the tail
- const int64_t width = std::min<int64_t>(n_kv, (int64_t) hparams.indexer_top_k + r - 1);
+ ggml_tensor * seed = ggml_cast(ctx0, ggml_view_1d(ctx0, sel_idx, 1, 0), GGML_TYPE_F32);
- ggml_tensor * top_k = ggml_cont(ctx0, ggml_top_k(ctx0, expanded, width));
+ 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);
- // 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);
- cb(top_k, "indexer_top_k", il);
+ 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 top_k;
+ 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,