mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-03 03:17:32 -05:00
qwen4exp: fix tests
This commit is contained in:
@@ -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
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user