diff --git a/src/models/glm5-next.cpp b/src/models/glm5-next.cpp index 5b11c77c34..cb1c6fed5e 100644 --- a/src/models/glm5-next.cpp +++ b/src/models/glm5-next.cpp @@ -273,7 +273,7 @@ public: ggml_tensor * pool_idxs = nullptr; // I32 [kpool, n_pool] member cells per pool, n_kv sentinel for the padded pools ggml_tensor * pool_mask = nullptr; // F32/F16 [n_pool, n_tokens] ggml_tensor * tail_idxs = nullptr; // I32 [kpool - 1, n_tokens] - ggml_tensor * gather_mask = nullptr; // F32 [n_sel, 1, 1, n_tokens] + ggml_tensor * gather_mask = nullptr; // F32 [n_sel, 1, 1, n_tokens] 0 for live selection slots, -inf for dead ones // n_new is never below 1, see build_inp_kpool ggml_tensor * new_pool_idxs = nullptr; // I32 [kpool, n_new] members of the pools completed this ubatch ggml_tensor * new_pool_rep = nullptr; // I64 [n_new] cell to write each new pooled key into @@ -330,12 +330,11 @@ llama_model_glm5_next::llm_graph_input_kpool * llama_model_glm5_next::graph::bui inp->n_sel = (uint32_t) n_sel; inp->gather = (int64_t) n_tokens <= max_ub && (int64_t) n_kv > n_sel; - if (inp->gather) { - inp->gather_mask = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_sel, 1, 1, n_tokens); - ggml_set_input(inp->gather_mask); - // Keep the mask allocated even when no op reads it, because set_input_kpool always fills it. - ggml_build_forward_expand(gf, inp->gather_mask); - } + // Both paths read the slot mask: gather adds it to the scores, scatter maps its dead slots to dump rows. + inp->gather_mask = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_sel, 1, 1, n_tokens); + ggml_set_input(inp->gather_mask); + // Keep the mask allocated even when no op reads it, because set_input_kpool always fills it. + ggml_build_forward_expand(gf, inp->gather_mask); } inp->n_new = n_new; @@ -892,13 +891,22 @@ ggml_tensor * llama_model_glm5_next::graph::build_kpool_select( 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 (visible pools, real tail cells) address disjoint cells. Each dead slot writes its own dump row + // n_kv + slot, so the scatter indices of a token are unique: idx = dump + live*(idx - dump), live = exp(mask). + GGML_ASSERT(inp_kpool->gather_mask->ne[0] == n_sel && inp_kpool->gather_mask->ne[3] == n_tokens); + ggml_tensor * live = ggml_exp(ctx0, ggml_reshape_2d(ctx0, inp_kpool->gather_mask, n_sel, n_tokens)); + ggml_tensor * dump = ggml_arange(ctx0, (float) n_kv, (float) (n_kv + n_sel), 1.0f); + 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)); sel = ggml_view_2d(ctx0, sel, n_kv, n_tokens, sel->nb[2], 0);