diff --git a/src/llama-memory-hybrid-idx.cpp b/src/llama-memory-hybrid-idx.cpp index 3b43df92d3..0bfba3ce0b 100644 --- a/src/llama-memory-hybrid-idx.cpp +++ b/src/llama-memory-hybrid-idx.cpp @@ -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 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; } } diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp index ab12c1d546..cf8a487241 100644 --- a/src/models/qwen4exp.cpp +++ b/src/models/qwen4exp.cpp @@ -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);