From 05af0d2b1398394cfa67e1918fee7feabccaa9bc Mon Sep 17 00:00:00 2001 From: Pascal Date: Wed, 30 Sep 2026 17:10:48 +0200 Subject: [PATCH] glm5-next: give dead indexer slots unique scatter rows (#29745) The sparse indexer mask is built with a set_rows scatter. Padded pools, absent sequences and missing tail cells all pointed to the same n_kv sentinel row, and invisible pools picked by top_k to fill the selection overlap the tail cells of the token, so several CPU threads wrote the same element (ThreadSanitizer data race in the sanitize CI). Allocate the slot mask for both selection paths and route every dead slot to its own dump row n_kv + slot. Live slots address disjoint cells, so the scatter indices of a token are unique. --- src/models/glm5-next.cpp | 26 +++++++++++++++++--------- 1 file changed, 17 insertions(+), 9 deletions(-) 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);