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