qwen4exp : optimize mask constructions (#29824)

* qwen4exp : optimize mask constructions

* cont : apply the same change for GLM5-next
This commit is contained in:
Georgi Gerganov
2026-10-02 11:08:50 +03:00
committed by GitHub
parent 631109b34d
commit 4e2713c162
3 changed files with 19 additions and 15 deletions
+3 -1
View File
@@ -273,7 +273,7 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) {
ggml_is_contiguous(node->src[3]), node->src[3]->name);
}
if (node) {
GGML_LOG_DEBUG("%s: node - %4s [%5lld, %5lld, %5lld, %5lld] [%5lld, %5lld, %5lld, %5lld], 1, %s\n", __func__, ggml_type_name(node->type), ne0, ne1, ne2, ne3, nb0, nb1, nb2, nb3,
GGML_LOG_DEBUG("%s: node - %4s [%5lld, %5lld, %5lld, %5lld] [%5lld, %5lld, %5lld, %5lld], 1, %s\n", __func__, ggml_type_name(node->type), ne0, ne1, ne2, ne3, nb0, nb1, nb2, nb3,
node->name);
}
}
@@ -649,6 +649,8 @@ int ggml_metal_op_repeat(ggml_metal_op_t ctx, int idx) {
GGML_TENSOR_LOCALS( int32_t, ne, op, ne);
GGML_TENSOR_LOCALS(uint64_t, nb, op, nb);
// TODO: optimize for degenerate cases such as ggml_nelements(op->src[0]) == 1 and others
auto pipeline = ggml_metal_library_get_pipeline_repeat(lib, op->type);
ggml_metal_kargs_repeat args = {
+7 -7
View File
@@ -886,16 +886,16 @@ ggml_tensor * llama_model_glm5_next::graph::build_kpool_select(
return sel_idx;
}
// Tie scatter storage lifetime to this layer's selected indices.
ggml_tensor * seed = ggml_cast(ctx0, ggml_view_1d(ctx0, sel_idx, 1, 0), GGML_TYPE_F32);
ggml_build_forward_expand(gf, sel_idx);
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 + n_sel, n_tokens, 1);
ggml_tensor * mask_all = ggml_new_tensor_4d(ctx0, kq_mask->type, n_kv + n_sel, 1, 1, 1);
mask_all = ggml_fill(ctx0, mask_all, -INFINITY);
mask_all = ggml_repeat_4d(ctx0, mask_all, n_kv + n_sel, n_tokens, 1, 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);
ggml_tensor * zeros = ggml_new_tensor_4d(ctx0, kq_mask->type, n_sel, 1, 1, 1);
zeros = ggml_fill(ctx0, zeros, 0.0f);
zeros = ggml_repeat_4d(ctx0, zeros, n_sel, n_tokens, 1, 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
+9 -7
View File
@@ -840,20 +840,22 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel(
const int64_t n_sel = sel_idx->ne[0];
GGML_ASSERT(n_sel == inp_kpool->n_sel);
ggml_build_forward_expand(gf, sel_idx);
// TODO: figure out to reduce the large copmute buffer that this creates
// 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;
ggml_tensor * seed = ggml_cast(ctx0, ggml_view_1d(ctx0, sel_idx, 1, 0), GGML_TYPE_F32);
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 + n_sel, n_tokens, 1);
ggml_tensor * mask_all = ggml_new_tensor_4d(ctx0, kq_mask->type, n_kv + n_sel, 1, 1, 1);
mask_all = ggml_fill(ctx0, mask_all, -INFINITY);
mask_all = ggml_repeat_4d(ctx0, mask_all, n_kv + n_sel, n_tokens, 1, 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);
ggml_tensor * zeros = ggml_new_tensor_4d(ctx0, kq_mask->type, n_sel, 1, 1, 1);
zeros = ggml_fill(ctx0, zeros, 0.0f);
zeros = ggml_repeat_4d(ctx0, zeros, n_sel, n_tokens, 1, 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