From 8a27bef8d42c9948e24601ab58e1de1614eb5722 Mon Sep 17 00:00:00 2001 From: Kawrakow Date: Tue, 28 Jul 2026 07:51:50 +0300 Subject: [PATCH] DS4 refactoring (cont'd) (#2194) --- src/graphs/build_deepseek4.cpp | 642 ++++++++++++++++----------------- src/llama.cpp | 6 +- 2 files changed, 313 insertions(+), 335 deletions(-) diff --git a/src/graphs/build_deepseek4.cpp b/src/graphs/build_deepseek4.cpp index 5bbc7a5f..3525736d 100644 --- a/src/graphs/build_deepseek4.cpp +++ b/src/graphs/build_deepseek4.cpp @@ -616,39 +616,6 @@ static ggml_tensor * dsv4_build_attn( return ggml_cont_2d(ctx, cur, cur->ne[0] * cur->ne[1], cur->ne[2] * cur->ne[3]); } -static ggml_tensor * build_hc_sinkhorn( - ggml_context * ctx0, - const llama_hparams & hparams, - ggml_tensor * comb) { - comb = ggml_soft_max(ctx0, comb); - - ggml_tensor * eps = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, 1); - eps = ggml_fill(ctx0, eps, hparams.dsv4_hc_eps); - comb = ggml_add(ctx0, comb, eps); - - auto norm_cols = [&]() { - ggml_tensor * comb_src_dst = ggml_cont(ctx0, ggml_permute(ctx0, comb, 1, 0, 2, 3)); - ggml_tensor * col_sum = ggml_sum_rows(ctx0, comb_src_dst); - col_sum = ggml_add(ctx0, col_sum, eps); - col_sum = ggml_permute(ctx0, col_sum, 1, 0, 2, 3); - comb = ggml_div(ctx0, comb, col_sum); - }; - - auto norm_rows = [&]() { - ggml_tensor * row_sum = ggml_sum_rows(ctx0, comb); - row_sum = ggml_add(ctx0, row_sum, eps); - comb = ggml_div(ctx0, comb, row_sum); - }; - - norm_cols(); - for (uint32_t i = 1; i < hparams.dsv4_hc_sinkhorn_iters; ++i) { - norm_rows(); - norm_cols(); - } - - return comb; -} - static ggml_tensor * build_hc_pre( ggml_context * ctx0, llm_build_context & llm, @@ -1039,6 +1006,307 @@ static void ds4_build_comp(ggml_tensor * cur, llm_build_context & llm, ggml_cont llm.cb(state_score_write, (tag + "_score_state_persist").c_str(), il); } +static ggml_tensor * ds4_attention(ggml_cgraph * gf, ggml_context * ctx0, llm_build_context & llm, ggml_tensor * inpL, + ggml_tensor ** append_csa_state, ggml_tensor ** append_csa_score, + ggml_tensor ** append_lid_state, ggml_tensor ** append_lid_score, + ggml_tensor * inp_pos, ggml_tensor * KQ_mask, int il) { + + ggml_tensor * residual = inpL; + ggml_tensor * post = nullptr; + ggml_tensor * comb = nullptr; + + const auto & model = llm.model; + const auto & layer = model.layers[il]; + const auto & hparams = model.hparams; + const auto & cparams = llm.cparams; + const auto & cb = llm.cb; + + auto & lctx = llm.lctx; + auto & kv_self = llm.kv_self; + + const int64_t n_embd_head = hparams.n_embd_head_k(0); + const int64_t n_embd_head_rope = hparams.n_rot; + const int64_t n_embd_head_nope = n_embd_head - n_embd_head_rope; + const int64_t hc = hparams.dsv4_hc_mult; + + const auto n_tokens = llm.n_tokens; + const auto n_head = llm.n_head; + const auto n_kv = llm.n_kv; + + ggml_tensor * cur = build_hc_pre(ctx0, llm, hparams, llm.n_embd, hparams.f_norm_rms_eps, inpL, + layer.hc_attn_fn, + layer.hc_attn_scale, + layer.hc_attn_base, + &post, &comb, llm.cb, il); + llm.cb(cur, "hc_attn_pre", il); + + cur = llm.llm_build_norm(ctx0, cur, hparams, layer.attn_norm, nullptr, LLM_NORM_RMS, llm.cb, il); + cb(cur, "attn_norm", il); + + ggml_tensor * qr = llm.llm_build_lora_mm(llm.lctx, ctx0, layer.wq_a, cur); + cb(qr, "qr", il); + + qr = llm.llm_build_norm(ctx0, qr, hparams, layer.attn_q_a_norm, nullptr, LLM_NORM_RMS, cb, il); + cb(qr, "qr_norm", il); + + const int64_t ratio = hparams.dsv4_compress_ratios[il]; + const bool use_compress_rope = ratio != 0; + const float freq_base_l = use_compress_rope ? hparams.dsv4_compress_rope_base : llm.freq_base; + const float freq_scale_l = use_compress_rope ? llm.freq_scale : 1.0f; + const float ext_factor_l = use_compress_rope ? llm.ext_factor : 0.0f; + const float attn_factor_l = dsv4_rope_attn_factor(freq_scale_l, ext_factor_l); + const float beta_fast_l = use_compress_rope ? llm.beta_fast : 0.0f; + const float beta_slow_l = use_compress_rope ? llm.beta_slow : 0.0f; + const int32_t n_ctx_orig_l = use_compress_rope ? llm.n_ctx_orig : 0; + + auto build_rope = [&] (int nhead, ggml_tensor * qin, ggml_tensor * wq, ggml_tensor * norm, const std::string & tag) { + auto q = llm.llm_build_lora_mm(llm.lctx, ctx0, wq, qin); + cb(q, (tag + "_b").c_str(), il); + q = ggml_reshape_2d(ctx0, q, n_embd_head, nhead * n_tokens); + q = llm.llm_build_norm(ctx0, q, hparams, norm, nullptr, LLM_NORM_RMS, cb, il); + cb(q, (tag + "_norm").c_str(), il); + q = ggml_reshape_3d(ctx0, q, n_embd_head, nhead, n_tokens); + ggml_tensor * q_nope = ggml_view_3d(ctx0, q, n_embd_head_nope, nhead, n_tokens, + ggml_row_size(q->type, n_embd_head), + ggml_row_size(q->type, n_embd_head) * nhead, + 0); + ggml_tensor * q_rope = ggml_view_3d(ctx0, q, n_embd_head_rope, nhead, n_tokens, + ggml_row_size(q->type, n_embd_head), + ggml_row_size(q->type, n_embd_head) * nhead, + ggml_row_size(q->type, n_embd_head_nope)); + q_rope = ggml_rope_ext(ctx0, q_rope, inp_pos, nullptr, n_embd_head_rope, llm.rope_type, n_ctx_orig_l, + freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); + cb(q_rope, (tag + "_rope").c_str(), il); + q = ggml_concat(ctx0, q_nope, q_rope, 0); + cb(q, tag.c_str(), il); + return q; + }; + + auto q = build_rope(n_head, qr, layer.wq_b, nullptr, "q"); + + auto kv = build_rope(1, cur, layer.wkv_latent, layer.attn_kv_norm, "kv"); + + if (cparams.k_cache_hadamard) { + if (int block_size = lctx.model.hadamard_size_k(il); block_size > 0) { + q = ggml_hadamard(ctx0, q, block_size); + kv = ggml_hadamard(ctx0, kv, block_size); + cb(q, "q_hadamard", il); + cb(kv, "kv_hadamard", il); + } + } + const float kq_scale = 1.0f / std::sqrt(float(n_embd_head)); + + if (ratio == llama_context::dsv4_runtime::CSA_RATIO && + lctx.dsv4.inputs.csa.state_pos != nullptr && + lctx.dsv4.csa_plan.state_pos.size() > 0) { + + ds4_build_comp(cur, llm, ctx0, lctx.dsv4.inputs.csa, lctx.dsv4.csa_plan, + layer.attn_comp_wkv, layer.attn_comp_wgate, + layer.attn_comp_ape, layer.attn_comp_norm, + lctx.dsv4.cache.csa_state_kv[il], lctx.dsv4.cache.csa_state_score[il], lctx.dsv4.cache.csa_k[il], + append_csa_state, append_csa_score, + n_embd_head, il, false, "csa", gf, false); + + + ds4_build_comp(cur, llm, ctx0, lctx.dsv4.inputs.lid, lctx.dsv4.lid_plan, + layer.indexer_comp_wkv, layer.indexer_comp_wgate, + layer.indexer_comp_ape, layer.indexer_comp_norm, + lctx.dsv4.cache.lid_state_kv[il], lctx.dsv4.cache.lid_state_score[il], lctx.dsv4.cache.lid_k[il], + append_lid_state, append_lid_score, + hparams.indexer_head_size, il, true, "lid", gf, false); + + } + + if (ratio == llama_context::dsv4_runtime::HCA_RATIO && + lctx.dsv4.inputs.hca.state_pos != nullptr && + lctx.dsv4.hca_plan.state_pos.size() > 0) { + + ds4_build_comp(cur, llm, ctx0, lctx.dsv4.inputs.hca, lctx.dsv4.hca_plan, + layer.attn_comp_wkv, layer.attn_comp_wgate, + layer.attn_comp_ape, layer.attn_comp_norm, + lctx.dsv4.cache.hca_state_kv[il], lctx.dsv4.cache.hca_state_score[il], lctx.dsv4.cache.hca_k[il], + nullptr, nullptr, + n_embd_head, il, false, "hca", gf, true); + + } + + ggml_tensor * raw_k_write = nullptr; + if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_write_idxs != nullptr) { + raw_k_write = dsv4_raw_cpy_k(&lctx, ctx0, kv_self.k_l[il], kv, + lctx.dsv4.inputs.raw_k_write_src_idxs, lctx.dsv4.inputs.raw_k_write_idxs, gf, n_embd_head, cb, il); + if (raw_k_write != nullptr) { + cb(raw_k_write, "dsv4_raw_k_write", il); + } + } + if (raw_k_write == nullptr) { + llm.llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, kv, nullptr, n_tokens, llm.kv_head, cb, il); + } + if (il < (int64_t) kv_self.v_l.size() && kv_self.v_l[il] != nullptr) { + llm.llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, nullptr, kv, n_tokens, llm.kv_head, cb, il); + } + + ggml_tensor * raw_k = nullptr; + if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_read_idxs != nullptr) { + raw_k = dsv4_raw_get_k(&lctx, ctx0, kv_self.k_l[il], lctx.dsv4.inputs.raw_k_read_idxs, n_embd_head, cb, il); + } + if (raw_k == nullptr) { + raw_k = ggml_view_3d(ctx0, kv_self.k_l[il], + n_embd_head, hparams.n_head_kv(il), n_kv, + ggml_row_size(kv_self.k_l[il]->type, n_embd_head), + ggml_row_size(kv_self.k_l[il]->type, n_embd_head) * hparams.n_head_kv(il), + 0); + } + cb(raw_k, "raw_k", il); + + const int64_t raw_kq_n_kv = raw_k != nullptr && lctx.dsv4.raw.n_kv > 0 + ? lctx.dsv4.raw.n_kv + : (raw_k != nullptr ? raw_k->ne[2] * raw_k->ne[3] : n_kv); + const int64_t raw_attn_n_kv = raw_kq_n_kv > 0 ? std::max(256, GGML_PAD(raw_kq_n_kv, 256)) : raw_kq_n_kv; + if (raw_k != nullptr && raw_k->ne[3] == 1) { + raw_k = dsv4_pad_raw_k_to(ctx0, raw_k, raw_attn_n_kv); + } + ggml_tensor * raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, + lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, raw_k->ne[3], cb, il); + cb(raw_mask, "raw_mask_view", il); + raw_mask = dsv4_pad_mask_tokens(ctx0, raw_mask, n_tokens); + raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); + cb(raw_mask, "dsv4_raw_mask_padded", il); + ggml_tensor * attn = nullptr; + + if (hparams.n_swa > 0) { + constexpr int k_fa_chunk = 256; + int n_swa = hparams.n_swa; + int ntokens = std::max(k_fa_chunk, int(q->ne[2])); + int nton = k_fa_chunk*((ntokens + n_swa + k_fa_chunk - 1)/k_fa_chunk); + int first = raw_k->ne[2] - nton; + if (first > 0) { + raw_k = ggml_view_4d(ctx0, raw_k, raw_k->ne[0], raw_k->ne[1], nton, raw_k->ne[3], + raw_k->nb[1], raw_k->nb[2], raw_k->nb[3], raw_k->nb[2]*first); + raw_mask = ggml_view_4d(ctx0, raw_mask, nton, raw_mask->ne[1], raw_mask->ne[2], raw_mask->ne[3], + raw_mask->nb[1], raw_mask->nb[2], raw_mask->nb[3], raw_mask->nb[0]*first); + } + } + + auto build_the_attn = [&] (ggml_tensor * raw_k, ggml_tensor * raw_mask, ggml_tensor * extra_mask, + ggml_tensor * cache, const auto & extra_ctx, + const std::string & tag, int n_swa_eff) { + auto n_stream = std::max(1, lctx.dsv4.cache.n_stream); + auto extra_k = dsv4_comp_get_k(ctx0, cache, extra_ctx, n_embd_head, cache->ne[1]/n_stream); + if (cparams.flash_attn) { + extra_mask = dsv4_pad_mask_tokens(ctx0, extra_mask, n_tokens); + } + raw_k = dsv4_repeat_streams(ctx0, raw_k, extra_k->ne[3]); + if (!cparams.flash_attn) { + raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, + lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, extra_k->ne[3], cb, il); + raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); + } + if (cparams.flash_attn && extra_mask->type != GGML_TYPE_F16) { + extra_mask = ggml_cast(ctx0, extra_mask, GGML_TYPE_F16); + } + if (raw_mask->type != extra_mask->type) { + raw_mask = ggml_cast(ctx0, raw_mask, extra_mask->type); + } + if (raw_k->type != extra_k->type) { + extra_k = ggml_cast(ctx0, extra_k, raw_k->type); + } + ggml_tensor * k_all = ggml_concat(ctx0, raw_k, extra_k, 2); + ggml_tensor * kq_mask = ggml_concat(ctx0, raw_mask, extra_mask, 0); + cb(extra_k, (tag + "_k").c_str(), il); + cb(k_all, (tag + "_k_all").c_str(), il); + cb(kq_mask, (tag + "_kq_mask").c_str(), il); + + auto attn = dsv4_build_attn(ctx0, hparams, cparams, q, k_all, k_all, kq_mask, + model.layers[il].attn_sinks, kq_scale, cb, il, n_swa_eff, gf); + return attn; + //return std::make_pair(k_all, kq_mask); + }; + + auto num_streams = [] (const auto & comp) { + int n_stream = comp.sinfo.n_stream(); + return std::max(1, n_stream); + }; + + if (ratio == llama_context::dsv4_runtime::CSA_RATIO && + lctx.dsv4.inputs.csa.kq_mask != nullptr && + lctx.dsv4.csa_plan.n_kv > 0 && + lctx.dsv4.lid_plan.n_kv > 0 && + !cparams.k_cache_hadamard) { + auto csa_mask = lctx.dsv4.inputs.csa.kq_mask; + if (hparams.indexer_top_k < lctx.dsv4.inputs.csa.kq_mask->ne[0]) { + auto top_k = dsv4_build_lid_top_k(ctx0, llm, qr, cur, inp_pos, il, gf, cb); + csa_mask = build_top_k_mask(ctx0, + dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.csa.kq_mask, nullptr, + lctx.dsv4.csa_plan.n_kv, n_tokens, num_streams(lctx.dsv4.csa_ctx), cb, il), + top_k); + cb(csa_mask, "csa_mask", il); + } + int n_csa = hparams.n_swa + hparams.indexer_top_k; + attn = build_the_attn(raw_k, raw_mask, csa_mask, lctx.dsv4.cache.csa_k[il], lctx.dsv4.csa_ctx, "csa", n_csa); + cb(attn, "attn_csa", il); + } else if (ratio == llama_context::dsv4_runtime::HCA_RATIO && + lctx.dsv4.inputs.hca.kq_mask != nullptr && + lctx.dsv4.hca_plan.n_kv > 0 && + std::any_of(lctx.dsv4.hca_plan.n_visible.begin(), lctx.dsv4.hca_plan.n_visible.end(), + [](int32_t n_visible) { return n_visible > 0; }) && + !cparams.k_cache_hadamard) { + ggml_tensor * hca_mask = dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.hca.kq_mask, nullptr, + lctx.dsv4.hca_plan.n_kv, n_tokens, num_streams(lctx.dsv4.hca_ctx), cb, il); + int n_hca = hparams.n_swa + (n_kv + llama_context::dsv4_runtime::HCA_RATIO - 1)/llama_context::dsv4_runtime::HCA_RATIO; + attn = build_the_attn(raw_k, raw_mask, hca_mask, lctx.dsv4.cache.hca_k[il], lctx.dsv4.hca_ctx, "hca", n_hca); + cb(attn, "attn_hca", il); + } else { + attn = dsv4_build_attn(ctx0, hparams, cparams, q, raw_k, raw_k, raw_mask, model.layers[il].attn_sinks, kq_scale, cb, il, -1, gf); + cb(attn, "attn_raw", il); + } + ggml_build_forward_expand(gf, attn); + + attn = ggml_reshape_3d(ctx0, attn, n_embd_head, n_head, n_tokens); + ggml_tensor * attn_nope = ggml_view_3d(ctx0, attn, n_embd_head_nope, n_head, n_tokens, + ggml_row_size(attn->type, n_embd_head), + ggml_row_size(attn->type, n_embd_head) * n_head, + 0); + ggml_tensor * attn_pe = ggml_view_3d(ctx0, attn, n_embd_head_rope, n_head, n_tokens, + ggml_row_size(attn->type, n_embd_head), + ggml_row_size(attn->type, n_embd_head) * n_head, + ggml_row_size(attn->type, n_embd_head_nope)); + attn_pe = ggml_rope_back(ctx0, attn_pe, inp_pos, nullptr, n_embd_head_rope, llm.rope_type, n_ctx_orig_l, + freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); + cb(attn_pe, "attn_derope", il); + attn = ggml_concat(ctx0, attn_nope, attn_pe, 0); + cb(attn, "attn", il); + + const int64_t o_group_dim = layer.wo_a->ne[0]; + const int64_t n_groups = (n_head * n_embd_head) / o_group_dim; + const int64_t o_lora_rank = layer.wo_b->ne[0] / n_groups; + + GGML_ASSERT((n_head * n_embd_head) % o_group_dim == 0); + GGML_ASSERT(layer.wo_b->ne[0] % n_groups == 0); + + attn = ggml_reshape_3d(ctx0, attn, o_group_dim, n_groups, n_tokens); + attn = ggml_permute(ctx0, attn, 0, 2, 1, 3); + + ggml_tensor * oa = ggml_mul_mat(ctx0, + ggml_reshape_3d(ctx0, layer.wo_a, layer.wo_a->ne[0], o_lora_rank, n_groups), + attn); + cb(oa, "attn_wo_a", il); + oa = ggml_permute(ctx0, oa, 0, 2, 1, 3); + if (n_tokens == 1) { + oa = ggml_reshape_2d(ctx0, oa, o_lora_rank * n_groups, n_tokens); + } else { + oa = ggml_cont_2d(ctx0, oa, o_lora_rank * n_groups, n_tokens); + } + + cur = llm.llm_build_lora_mm(lctx, ctx0, layer.wo_b, oa); + cb(cur, "attn_out", il); + + inpL = llm.build_mhc_post(cur, post, residual, comb, llm.n_embd, hc, true); + cb(inpL, "hc_attn_post", il); + + return inpL; + +} + ggml_cgraph * llm_build_context::build_deepseek4() { ggml_cgraph * gf = new_graph_custom(); @@ -1046,8 +1314,6 @@ ggml_cgraph * llm_build_context::build_deepseek4() { GGML_ABORT("DeepSeek4 MTP execution is not implemented"); } - //printf("================================================================= %s\n", __func__); - const int64_t n_embd_head = hparams.n_embd_head_k(0); const int64_t n_embd_head_rope = hparams.n_rot; const int64_t n_embd_head_nope = n_embd_head - n_embd_head_rope; @@ -1076,307 +1342,15 @@ ggml_cgraph * llm_build_context::build_deepseek4() { ggml_tensor * append_lid_score = nullptr; for (int il = 0; il < n_layer; ++il) { - ggml_tensor * residual = inpL; - ggml_tensor * post = nullptr; - ggml_tensor * comb = nullptr; - ggml_tensor * cur = build_hc_pre(ctx0, *this, hparams, n_embd, hparams.f_norm_rms_eps, - inpL, - model.layers[il].hc_attn_fn, - model.layers[il].hc_attn_scale, - model.layers[il].hc_attn_base, - &post, &comb, cb, il); - cb(cur, "hc_attn_pre", il); + auto cur = ds4_attention(gf, ctx0, *this, inpL, + &append_csa_state, &append_csa_score, + &append_lid_state, &append_lid_score, + inp_pos, KQ_mask, il); + inpL = cur; - cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].attn_norm, nullptr, LLM_NORM_RMS, cb, il); - cb(cur, "attn_norm", il); - - ggml_tensor * qr = llm_build_lora_mm(lctx, ctx0, model.layers[il].wq_a, cur); - cb(qr, "qr", il); - - qr = llm_build_norm(ctx0, qr, hparams, model.layers[il].attn_q_a_norm, nullptr, LLM_NORM_RMS, cb, il); - cb(qr, "qr_norm", il); - - const int64_t ratio = hparams.dsv4_compress_ratios[il]; - const bool use_compress_rope = ratio != 0; - const float freq_base_l = use_compress_rope ? hparams.dsv4_compress_rope_base : freq_base; - const float freq_scale_l = use_compress_rope ? freq_scale : 1.0f; - const float ext_factor_l = use_compress_rope ? ext_factor : 0.0f; - const float attn_factor_l = dsv4_rope_attn_factor(freq_scale_l, ext_factor_l); - const float beta_fast_l = use_compress_rope ? beta_fast : 0.0f; - const float beta_slow_l = use_compress_rope ? beta_slow : 0.0f; - const int32_t n_ctx_orig_l = use_compress_rope ? n_ctx_orig : 0; - - ggml_tensor * q = llm_build_lora_mm(lctx, ctx0, model.layers[il].wq_b, qr); - cb(q, "q_b", il); - q = ggml_reshape_3d(ctx0, q, n_embd_head, n_head, n_tokens); - q = ggml_rms_norm(ctx0, q, hparams.f_norm_rms_eps); - cb(q, "q_b", il); - - ggml_tensor * q_nope = ggml_view_3d(ctx0, q, n_embd_head_nope, n_head, n_tokens, - ggml_row_size(q->type, n_embd_head), - ggml_row_size(q->type, n_embd_head) * n_head, - 0); - ggml_tensor * q_pe = ggml_view_3d(ctx0, q, n_embd_head_rope, n_head, n_tokens, - ggml_row_size(q->type, n_embd_head), - ggml_row_size(q->type, n_embd_head) * n_head, - ggml_row_size(q->type, n_embd_head_nope)); - q_pe = ggml_rope_ext(ctx0, q_pe, inp_pos, nullptr, n_embd_head_rope, rope_type, n_ctx_orig_l, - freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); - cb(q_pe, "q_pe", il); - q = ggml_concat(ctx0, q_nope, q_pe, 0); - cb(q, "q", il); - - ggml_tensor * kv = llm_build_lora_mm(lctx, ctx0, model.layers[il].wkv_latent, cur); - cb(kv, "wkv", il); - kv = llm_build_norm(ctx0, kv, hparams, model.layers[il].attn_kv_norm, nullptr, LLM_NORM_RMS, cb, il); - cb(kv, "kv_norm", il); - kv = ggml_reshape_3d(ctx0, kv, n_embd_head, 1, n_tokens); - cb(kv, "kv_norm", il); - - ggml_tensor * kv_nope = ggml_view_3d(ctx0, kv, n_embd_head_nope, 1, n_tokens, - ggml_row_size(kv->type, n_embd_head), - ggml_row_size(kv->type, n_embd_head), - 0); - ggml_tensor * kv_pe = ggml_view_3d(ctx0, kv, n_embd_head_rope, 1, n_tokens, - ggml_row_size(kv->type, n_embd_head), - ggml_row_size(kv->type, n_embd_head), - ggml_row_size(kv->type, n_embd_head_nope)); - kv_pe = ggml_rope_ext(ctx0, kv_pe, inp_pos, nullptr, n_embd_head_rope, rope_type, n_ctx_orig_l, - freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); - cb(kv_pe, "kv_pe", il); - kv = ggml_concat(ctx0, kv_nope, kv_pe, 0); - cb(kv, "kv", il); - - if (cparams.k_cache_hadamard) { - if (int block_size = lctx.model.hadamard_size_k(il); block_size > 0) { - q = ggml_hadamard(ctx0, q, block_size); - kv = ggml_hadamard(ctx0, kv, block_size); - cb(q, "q_hadamard", il); - cb(kv, "kv_hadamard", il); - } - } - const float kq_scale = 1.0f / std::sqrt(float(n_embd_head)); - - if (ratio == llama_context::dsv4_runtime::CSA_RATIO && - lctx.dsv4.inputs.csa.state_pos != nullptr && - lctx.dsv4.csa_plan.state_pos.size() > 0) { - - ds4_build_comp(cur, *this, ctx0, lctx.dsv4.inputs.csa, lctx.dsv4.csa_plan, - model.layers[il].attn_comp_wkv, model.layers[il].attn_comp_wgate, - model.layers[il].attn_comp_ape, model.layers[il].attn_comp_norm, - lctx.dsv4.cache.csa_state_kv[il], lctx.dsv4.cache.csa_state_score[il], lctx.dsv4.cache.csa_k[il], - &append_csa_state, &append_csa_score, - n_embd_head, il, false, "csa", gf, false); - - - ds4_build_comp(cur, *this, ctx0, lctx.dsv4.inputs.lid, lctx.dsv4.lid_plan, - model.layers[il].indexer_comp_wkv, model.layers[il].indexer_comp_wgate, - model.layers[il].indexer_comp_ape, model.layers[il].indexer_comp_norm, - lctx.dsv4.cache.lid_state_kv[il], lctx.dsv4.cache.lid_state_score[il], lctx.dsv4.cache.lid_k[il], - &append_lid_state, &append_lid_score, - hparams.indexer_head_size, il, true, "lid", gf, false); - - } - - if (ratio == llama_context::dsv4_runtime::HCA_RATIO && - lctx.dsv4.inputs.hca.state_pos != nullptr && - lctx.dsv4.hca_plan.state_pos.size() > 0) { - - ds4_build_comp(cur, *this, ctx0, lctx.dsv4.inputs.hca, lctx.dsv4.hca_plan, - model.layers[il].attn_comp_wkv, model.layers[il].attn_comp_wgate, - model.layers[il].attn_comp_ape, model.layers[il].attn_comp_norm, - lctx.dsv4.cache.hca_state_kv[il], lctx.dsv4.cache.hca_state_score[il], lctx.dsv4.cache.hca_k[il], - nullptr, nullptr, - n_embd_head, il, false, "hca", gf, true); - - } - - ggml_tensor * raw_k_write = nullptr; - if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_write_idxs != nullptr) { - raw_k_write = dsv4_raw_cpy_k(&lctx, ctx0, kv_self.k_l[il], kv, lctx.dsv4.inputs.raw_k_write_src_idxs, lctx.dsv4.inputs.raw_k_write_idxs, gf, n_embd_head, cb, il); - if (raw_k_write != nullptr) { - cb(raw_k_write, "dsv4_raw_k_write", il); - } - } - if (raw_k_write == nullptr) { - llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, kv, nullptr, n_tokens, kv_head, cb, il); - } - if (il < (int64_t) kv_self.v_l.size() && kv_self.v_l[il] != nullptr) { - llm_build_kv_store(lctx, ctx0, hparams, cparams, kv_self, gf, nullptr, kv, n_tokens, kv_head, cb, il); - } - - ggml_tensor * raw_k = nullptr; - if (hparams.n_head_kv(il) == 1 && lctx.dsv4.inputs.raw_k_read_idxs != nullptr) { - raw_k = dsv4_raw_get_k(&lctx, ctx0, kv_self.k_l[il], lctx.dsv4.inputs.raw_k_read_idxs, n_embd_head, cb, il); - } - if (raw_k == nullptr) { - raw_k = ggml_view_3d(ctx0, kv_self.k_l[il], - n_embd_head, hparams.n_head_kv(il), n_kv, - ggml_row_size(kv_self.k_l[il]->type, n_embd_head), - ggml_row_size(kv_self.k_l[il]->type, n_embd_head) * hparams.n_head_kv(il), - 0); - } - cb(raw_k, "raw_k", il); - - const int64_t raw_kq_n_kv = raw_k != nullptr && lctx.dsv4.raw.n_kv > 0 - ? lctx.dsv4.raw.n_kv - : (raw_k != nullptr ? raw_k->ne[2] * raw_k->ne[3] : n_kv); - const int64_t raw_attn_n_kv = raw_kq_n_kv > 0 ? std::max(256, GGML_PAD(raw_kq_n_kv, 256)) : raw_kq_n_kv; - if (raw_k != nullptr && raw_k->ne[3] == 1) { - raw_k = dsv4_pad_raw_k_to(ctx0, raw_k, raw_attn_n_kv); - } - ggml_tensor * raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, - lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, raw_k->ne[3], cb, il); - cb(raw_mask, "raw_mask_view", il); - raw_mask = dsv4_pad_mask_tokens(ctx0, raw_mask, n_tokens); - raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); - cb(raw_mask, "dsv4_raw_mask_padded", il); - ggml_tensor * attn = nullptr; - - if (hparams.n_swa > 0) { - constexpr int k_fa_chunk = 256; - int n_swa = hparams.n_swa; - int ntokens = std::max(k_fa_chunk, int(q->ne[2])); - int nton = k_fa_chunk*((ntokens + n_swa + k_fa_chunk - 1)/k_fa_chunk); - int first = raw_k->ne[2] - nton; - if (first > 0) { - raw_k = ggml_view_4d(ctx0, raw_k, raw_k->ne[0], raw_k->ne[1], nton, raw_k->ne[3], - raw_k->nb[1], raw_k->nb[2], raw_k->nb[3], raw_k->nb[2]*first); - raw_mask = ggml_view_4d(ctx0, raw_mask, nton, raw_mask->ne[1], raw_mask->ne[2], raw_mask->ne[3], - raw_mask->nb[1], raw_mask->nb[2], raw_mask->nb[3], raw_mask->nb[0]*first); - } - } - - if (ratio == llama_context::dsv4_runtime::CSA_RATIO && - lctx.dsv4.inputs.csa.kq_mask != nullptr && - lctx.dsv4.csa_plan.n_kv > 0 && - lctx.dsv4.lid_plan.n_kv > 0 && - !cparams.k_cache_hadamard) { - ggml_tensor * csa_k = dsv4_comp_get_k(ctx0, - lctx.dsv4.cache.csa_k[il], - lctx.dsv4.csa_ctx, - n_embd_head, - lctx.dsv4.cache.csa_k[il]->ne[1]/std::max(1, lctx.dsv4.cache.n_stream)); - auto csa_mask = lctx.dsv4.inputs.csa.kq_mask; - if (hparams.indexer_top_k < lctx.dsv4.inputs.csa.kq_mask->ne[0]) { - auto top_k = dsv4_build_lid_top_k(ctx0, *this, qr, cur, inp_pos, il, gf, cb); - csa_mask = build_top_k_mask(ctx0, - dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.csa.kq_mask, nullptr, - lctx.dsv4.csa_plan.n_kv, n_tokens, csa_k->ne[3], cb, il), - top_k); - cb(csa_mask, "csa_mask", il); - } - const bool use_fattn = cparams.flash_attn; - if (use_fattn) { - csa_mask = dsv4_pad_mask_tokens(ctx0, csa_mask, n_tokens); - } - raw_k = dsv4_repeat_streams(ctx0, raw_k, csa_k->ne[3]); - if (!use_fattn) { - raw_mask = dsv4_build_raw_mask_view(ctx0, KQ_mask, - lctx.dsv4.inputs.raw_k_read_idxs, raw_kq_n_kv, n_tokens, csa_k->ne[3], cb, il); - raw_mask = dsv4_pad_raw_mask_to(ctx0, raw_mask, raw_attn_n_kv, n_tokens); - } - if (use_fattn && csa_mask->type != GGML_TYPE_F16) { - csa_mask = ggml_cast(ctx0, csa_mask, GGML_TYPE_F16); - } - if (raw_mask->type != csa_mask->type) { - raw_mask = ggml_cast(ctx0, raw_mask, csa_mask->type); - } - if (raw_k->type != csa_k->type) { - csa_k = ggml_cast(ctx0, csa_k, raw_k->type); - } - ggml_tensor * k_all = ggml_concat(ctx0, raw_k, csa_k, 2); - ggml_tensor * kq_mask = ggml_concat(ctx0, raw_mask, csa_mask, 0); - cb(csa_k, "csa_k", il); - cb(k_all, "csa_k_all", il); - cb(kq_mask, "csa_kq_mask", il); - int n_csa = hparams.n_swa + hparams.indexer_top_k; - attn = dsv4_build_attn(ctx0, hparams, cparams, q, k_all, k_all, kq_mask, model.layers[il].attn_sinks, kq_scale, cb, il, n_csa, gf); - cb(attn, "attn_csa", il); - } else if (ratio == llama_context::dsv4_runtime::HCA_RATIO && - lctx.dsv4.inputs.hca.kq_mask != nullptr && - lctx.dsv4.hca_plan.n_kv > 0 && - std::any_of(lctx.dsv4.hca_plan.n_visible.begin(), lctx.dsv4.hca_plan.n_visible.end(), - [](int32_t n_visible) { return n_visible > 0; }) && - !cparams.k_cache_hadamard) { - ggml_tensor * hca_k = dsv4_comp_get_k(ctx0, - lctx.dsv4.cache.hca_k[il], - lctx.dsv4.hca_ctx, - n_embd_head, - lctx.dsv4.cache.hca_k[il]->ne[1]/std::max(1, lctx.dsv4.cache.n_stream)); - const bool use_fattn = cparams.flash_attn; - ggml_tensor * hca_mask = dsv4_build_raw_mask_view(ctx0, lctx.dsv4.inputs.hca.kq_mask, nullptr, - lctx.dsv4.hca_plan.n_kv, n_tokens, hca_k->ne[3], cb, il); - hca_mask = dsv4_pad_mask_tokens(ctx0, hca_mask, n_tokens); - if (use_fattn && hca_mask->type != GGML_TYPE_F16) { - hca_mask = ggml_cast(ctx0, hca_mask, GGML_TYPE_F16); - } - raw_k = dsv4_repeat_streams(ctx0, raw_k, hca_k->ne[3]); - if (raw_mask->type != hca_mask->type) { - raw_mask = ggml_cast(ctx0, raw_mask, hca_mask->type); - } - if (hca_k->type != raw_k->type) { - hca_k = ggml_cast(ctx0, hca_k, raw_k->type); - } - ggml_tensor * k_all = ggml_concat(ctx0, raw_k, hca_k, 2); - ggml_tensor * kq_mask = ggml_concat(ctx0, raw_mask, hca_mask, 0); - cb(hca_k, "hca_k", il); - cb(k_all, "hca_k_all", il); - cb(kq_mask, "hca_kq_mask", il); - int n_hca = (n_kv + llama_context::dsv4_runtime::HCA_RATIO - 1)/llama_context::dsv4_runtime::HCA_RATIO; - n_hca += hparams.n_swa; - attn = dsv4_build_attn(ctx0, hparams, cparams, q, k_all, k_all, kq_mask, model.layers[il].attn_sinks, kq_scale, cb, il, n_hca, gf); - cb(attn, "attn_hca", il); - } else { - //printf("Regular attention for layer %d\n", il); - attn = dsv4_build_attn(ctx0, hparams, cparams, q, raw_k, raw_k, raw_mask, model.layers[il].attn_sinks, kq_scale, cb, il, -1, gf); - cb(attn, "attn_raw", il); - } - - attn = ggml_reshape_3d(ctx0, attn, n_embd_head, n_head, n_tokens); - ggml_tensor * attn_nope = ggml_view_3d(ctx0, attn, n_embd_head_nope, n_head, n_tokens, - ggml_row_size(attn->type, n_embd_head), - ggml_row_size(attn->type, n_embd_head) * n_head, - 0); - ggml_tensor * attn_pe = ggml_view_3d(ctx0, attn, n_embd_head_rope, n_head, n_tokens, - ggml_row_size(attn->type, n_embd_head), - ggml_row_size(attn->type, n_embd_head) * n_head, - ggml_row_size(attn->type, n_embd_head_nope)); - attn_pe = ggml_rope_back(ctx0, attn_pe, inp_pos, nullptr, n_embd_head_rope, rope_type, n_ctx_orig_l, - freq_base_l, freq_scale_l, ext_factor_l, attn_factor_l, beta_fast_l, beta_slow_l); - cb(attn_pe, "attn_derope", il); - attn = ggml_concat(ctx0, attn_nope, attn_pe, 0); - cb(attn, "attn", il); - - const int64_t o_group_dim = model.layers[il].wo_a->ne[0]; - const int64_t n_groups = (n_head * n_embd_head) / o_group_dim; - const int64_t o_lora_rank = model.layers[il].wo_b->ne[0] / n_groups; - - GGML_ASSERT((n_head * n_embd_head) % o_group_dim == 0); - GGML_ASSERT(model.layers[il].wo_b->ne[0] % n_groups == 0); - - attn = ggml_reshape_3d(ctx0, attn, o_group_dim, n_groups, n_tokens); - attn = ggml_permute(ctx0, attn, 0, 2, 1, 3); - - ggml_tensor * oa = ggml_mul_mat(ctx0, - ggml_reshape_3d(ctx0, model.layers[il].wo_a, model.layers[il].wo_a->ne[0], o_lora_rank, n_groups), - attn); - cb(oa, "attn_wo_a", il); - oa = ggml_permute(ctx0, oa, 0, 2, 1, 3); - if (n_tokens == 1) { - oa = ggml_reshape_2d(ctx0, oa, o_lora_rank * n_groups, n_tokens); - } else { - oa = ggml_cont_2d(ctx0, oa, o_lora_rank * n_groups, n_tokens); - } - - cur = llm_build_lora_mm(lctx, ctx0, model.layers[il].wo_b, oa); - cb(cur, "attn_out", il); - - inpL = build_mhc_post(cur, post, residual, comb, n_embd, hc, true); - cb(inpL, "hc_attn_post", il); - - residual = inpL; + ggml_tensor *post, *comb; + auto residual = inpL; cur = build_hc_pre(ctx0, *this, hparams, n_embd, hparams.f_norm_rms_eps, inpL, model.layers[il].hc_ffn_fn, diff --git a/src/llama.cpp b/src/llama.cpp index 23b57cac..6d540da5 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -3707,12 +3707,16 @@ static std::pair, double> get_layer_sizes(const llama_model_ continue; } if (name == "output.weight") { - ow_size = size; + ow_size += size; continue; } if (name == "output_norm.weight") { continue; } + if (auto pos = name.find("output_hc_"); pos == 0) { + ow_size += size; + continue; + } if (model.arch == LLM_ARCH_GEMMA4) { if (name == "per_layer_token_embd.weight" || name == "per_layer_model_proj.weight" ||