qwen4exp : optimize mask constructions

This commit is contained in:
Georgi Gerganov
2026-10-01 23:10:42 +03:00
parent 169ce4f27c
commit ce5b0e0fbc
2 changed files with 12 additions and 8 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 = {
+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