qwen4exp: fix tests

This commit is contained in:
Aman Gupta
2026-10-02 01:54:29 +08:00
parent e358d59178
commit 169ce4f27c
2 changed files with 38 additions and 21 deletions
+17 -18
View File
@@ -944,25 +944,24 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
}
GGML_ASSERT(i_new == n_new);
// 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;
// Padded entries re-pool cells whose pooled slot is never read: only the reps of complete pools are read.
// Each entry takes its own cell, entries sharing one would write it from several threads in the scatter.
if (n_new_g > n_new) {
std::vector<int64_t> reps(pcell, pcell + pool_end.size());
std::sort(reps.begin(), reps.end());
int64_t pad_cell = 0;
for (uint32_t i = n_new; i < n_new_g; ++i, ++pad_cell) {
while (std::binary_search(reps.begin(), reps.end(), pad_cell)) {
++pad_cell;
}
GGML_ASSERT(pad_cell < (int64_t) kv_size*n_stream_kv);
for (uint32_t k = 0; k < kpool; ++k) {
nidx[(size_t) i*kpool + k] = (int32_t) pad_cell;
}
if (nrep != nullptr) {
nrep[i] = pad_cell;
}
}
}
for (uint32_t i = n_new; i < n_new_g; ++i) {
for (uint32_t k = 0; k < kpool; ++k) {
nidx[(size_t) i*kpool + k] = (int32_t) pad_cell;
}
if (nrep != nullptr) {
nrep[i] = pad_cell;
}
}
+21 -3
View File
@@ -841,7 +841,7 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel(
GGML_ASSERT(n_sel == inp_kpool->n_sel);
// TODO: figure out to reduce the large copmute buffer that this creates
// scatter zeros for the selected cells into an all -inf row, the extra row n_kv takes the sentinels
// scatter zeros for the selected cells into an all -inf row, each dead slot into its own dump row n_kv + slot
// seeding from sel_idx ties the scatter storage lifetime to this layer
const int64_t n_kv = inp_kpool->n_kv;
@@ -849,13 +849,31 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel(
ggml_tensor * mask_seed = kq_mask->type == GGML_TYPE_F32 ? seed : ggml_cast(ctx0, seed, kq_mask->type);
mask_seed = ggml_fill(ctx0, mask_seed, -INFINITY);
ggml_tensor * mask_all = ggml_repeat_4d(ctx0, mask_seed, 1, n_kv + 1, n_tokens, 1);
mask_all = ggml_reshape_3d(ctx0, mask_all, 1, n_kv + 1, n_tokens);
ggml_tensor * mask_all = ggml_repeat_4d(ctx0, mask_seed, 1, n_kv + n_sel, n_tokens, 1);
mask_all = ggml_reshape_3d(ctx0, mask_all, 1, n_kv + n_sel, n_tokens);
ggml_tensor * zero_seed = ggml_fill(ctx0, seed, 0.0f);
ggml_tensor * zeros = ggml_repeat_4d(ctx0, zero_seed, 1, n_sel, n_tokens, 1);
zeros = ggml_reshape_3d(ctx0, zeros, 1, n_sel, n_tokens);
// live slots address disjoint cells, but padded pools and missing tail cells share the n_kv sentinel, and
// top_k fills a short selection with invisible pools that can overlap the tail, so the scatter would write
// some cells from several threads: map every dead slot to its own dump row, idx = dump + live*(idx - dump)
// a picked pool is live when visible: a visible score is a rectified sum >= 0, an invisible one is -inf
ggml_tensor * top_score = ggml_get_rows(ctx0, ggml_reshape_3d(ctx0, score, 1, n_pool, n_tokens), top_k); // [1, n_top_pool, n_tokens]
ggml_tensor * live_pool = ggml_clamp(ctx0, ggml_scale_bias(ctx0, top_score, 1.0f, 1.0f), 0.0f, 1.0f);
live_pool = ggml_reshape_2d(ctx0, ggml_repeat_4d(ctx0, live_pool, kpool, n_top_pool, n_tokens, 1), kpool*n_top_pool, n_tokens);
// a tail cell is live unless it is the n_kv sentinel
ggml_tensor * live_tail = ggml_cast(ctx0, inp_kpool->tail_idxs, GGML_TYPE_F32);
live_tail = ggml_clamp(ctx0, ggml_scale_bias(ctx0, live_tail, -1.0f, (float) n_kv), 0.0f, 1.0f);
ggml_tensor * live = ggml_concat(ctx0, live_pool, live_tail, 0); // [n_sel, n_tokens]
// dump rows n_kv + slot as a cumulative sum: the meta backend cannot split an arange, which has no source
ggml_tensor * dump = ggml_scale_bias(ctx0, ggml_cumsum(ctx0, ggml_fill(ctx0, live, 1.0f)), 1.0f, (float) (n_kv - 1));
ggml_tensor * idx_f = ggml_cast(ctx0, sel_idx, GGML_TYPE_F32);
idx_f = ggml_add(ctx0, ggml_mul(ctx0, ggml_sub(ctx0, idx_f, dump), live), dump);
sel_idx = ggml_cast(ctx0, idx_f, GGML_TYPE_I32);
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);