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