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,