diff --git a/src/llama-build-context.cpp b/src/llama-build-context.cpp index 2b02f6b5..2d206cf9 100644 --- a/src/llama-build-context.cpp +++ b/src/llama-build-context.cpp @@ -749,7 +749,7 @@ ggml_tensor * llm_build_context::llm_build_ffn( llm_ffn_op_type type_op, llm_ffn_gate_type type_gate, const llm_build_cb & cb, int il, ggml_cgraph * graph, bool add_input, - bool is_norm, ggml_tensor * add_extra) { + bool is_norm, ggml_tensor * add_extra, ggml_tensor * post_norm) { if (!up_b && !up_s && !gate_b && !gate_s && !down_b && !down_s && up->extra && gate->extra && down->extra && type_gate == LLM_FFN_PAR && @@ -858,6 +858,10 @@ ggml_tensor * llm_build_context::llm_build_ffn( cur = ggml_mul(ctx, cur, down_s); cb(cur, "ffn_down_s", il); } + if (post_norm) { + cur = llm_build_norm(ctx, cur, lctx.model.hparams, post_norm, NULL, LLM_NORM_RMS, cb, il); + cb(cur, "ffn_post_normed", il); + } if (add_input) { cur = ggml_add(ctx, cur, input); cb(cur, "ffn_out_with_inp", il); @@ -990,6 +994,11 @@ ggml_tensor * llm_build_context::llm_build_ffn( cb(cur, "ffn_down_s", il); } + if (post_norm) { + cur = llm_build_norm(ctx, cur, lctx.model.hparams, post_norm, NULL, LLM_NORM_RMS, cb, il); + cb(cur, "ffn_post_normed", il); + } + if (add_input) { cur = ggml_add(ctx, cur, input); cb(cur, "ffn_out_with_inp", il); @@ -6021,6 +6030,308 @@ ggml_cgraph * llm_build_context::build_gemma3() { return gf; } +static ggml_cgraph * build_gemma4_graph_paralle(llm_build_context & llm, llama_context & lctx, ggml_context * ctx0, + ggml_tensor * inpL, ggml_tensor * inp_pos, ggml_tensor * inp_out_ids, + ggml_tensor * KQ_mask, ggml_tensor * KQ_mask_swa, int n_tokens, const llm_build_cb & cb) { + auto & model = lctx.model; + auto & hparams = model.hparams; + auto & cparams = lctx.cparams; + auto & kv_self = lctx.kv_self; + int n_device = model.splits.size(); + GGML_ASSERT(n_device > 1); + GGML_ASSERT(cparams.flash_attn); + auto gf = ggml_new_graph_custom(ctx0, llama_model_max_nodes(model, n_tokens), false); + + std::vector sa_inp(n_device, nullptr); + std::vector sa_out(n_device, nullptr); + std::vector ffn_inp(n_device, nullptr); + std::vector ffn_out(n_device, nullptr); + + //ggml_tensor * last_ffn_inp = nullptr; + //ggml_tensor * last_sa_inp = nullptr; + + for (int il = 0; il < hparams.n_layer; ++il) { + auto & l = model.layers[il]; + const bool is_sliding = hparams.swa_layers[il] ? true : false; + const float freq_base_l = is_sliding ? hparams.rope_freq_base_train_swa : cparams.rope_freq_base; + const float freq_scale_l = is_sliding ? hparams.rope_freq_scale_train_swa : cparams.rope_freq_scale; + const int n_rot_l = is_sliding ? hparams.n_rot_swa : hparams.n_rot; + const int n_swa = is_sliding ? hparams.n_swa : 0; + //const int n_embd_head = hparams.n_embd_head_k(il); + //const int n_head = hparams.n_head(il); + //const int n_head_kv = hparams.n_head_kv(il); + + struct ggml_tensor * KQ_mask_l = is_sliding ? KQ_mask_swa : KQ_mask; + + auto freq_factors = !is_sliding ? model.layers[il].rope_freqs : nullptr; + if (freq_factors) { + GGML_ASSERT(freq_factors->extra); + } + + auto wq = (const ggml_split_tensor_t *)l.wq->extra; + auto wk = (const ggml_split_tensor_t *)l.wk->extra; + auto wv = l.wv ? (const ggml_split_tensor_t *)l.wv->extra : nullptr; + auto wo = (const ggml_split_tensor_t *)l.wo->extra; + GGML_ASSERT(wq && wk && wo); + + auto q_norm = (const ggml_split_tensor_t *)l.attn_q_norm->extra; + auto k_norm = (const ggml_split_tensor_t *)l.attn_k_norm->extra; + GGML_ASSERT(q_norm && k_norm); + + auto kl = (ggml_split_tensor_t *)kv_self.k_l[il]->extra; + auto vl = (ggml_split_tensor_t *)kv_self.v_l[il]->extra; + GGML_ASSERT(kl && vl); + + for (int id = 0; id < n_device; ++id) { + GGML_ASSERT((wq->splits[id] && wk->splits[id] && (!wv || wv->splits[id]) && wo->splits[id]) || + (!wq->splits[id] && !wk->splits[id] && (!wv || !wv->splits[id]) && !wo->splits[id])); + if (!wq->splits[id]) { + sa_inp[id] = sa_out[id] = nullptr; + continue; + } + GGML_ASSERT(kl->splits[id] && vl->splits[id]); + + int il_cb = 1000*(il + 1) + id; + + if (il == 0) { + sa_inp[id] = inpL; + //sa_inp[id] = do_split_norm(ctx0, inpL, l.attn_norm, hparams, cb, id, il_cb, false); + } else { + GGML_ASSERT(inpL->op == GGML_OP_REDUCE); + auto cur = get_input_tensor_sm_graph(ctx0, inpL, id); + cur = do_split_norm(ctx0, cur, model.layers[il-1].ffn_post_norm, hparams, cb, id, il_cb, false); + cb(cur, "ffn_post_norm", il_cb); + auto add = ffn_inp[id]; + if (!add) { + for (int j = 0; j < n_device; ++j) { + if (ffn_inp[j]) { + add = ffn_inp[j]; break; + } + } + GGML_ASSERT(add); + } + sa_inp[id] = ggml_add(ctx0, cur, add); + cb(sa_inp[id], "sa_inp", il_cb); + if (model.layers[il-1].out_scale) { + auto scale = (const ggml_split_tensor_t *)model.layers[il-1].out_scale->extra; + sa_inp[id] = ggml_mul(ctx0, sa_inp[id], scale->splits[id]); + cb(sa_inp[id], "sa_inp_scaled", il_cb); + } + } + auto cur = do_split_norm(ctx0, sa_inp[id], model.layers[il].attn_norm, hparams, cb, id, il_cb, false); + cb(cur, "sa_inp_normed", il_cb); + auto Qcur = llm.llm_build_lora_mm(lctx, ctx0, wq->splits[id], cur); + cb(Qcur, "Qcur", il_cb); + auto Kcur = llm.llm_build_lora_mm(lctx, ctx0, wk->splits[id], cur); + cb(Kcur, "Kcur", il_cb); + ggml_tensor * Vcur = nullptr; + if (wv) { + Vcur = llm.llm_build_lora_mm(lctx, ctx0, wv->splits[id], cur); + cb(Vcur, "Vcur", il_cb); + } + ggml_build_forward_expand(gf, Qcur); + ggml_build_forward_expand(gf, Kcur); + if (Vcur) { + ggml_build_forward_expand(gf, Vcur); + Vcur = ggml_reshape_3d(ctx0, Vcur, hparams.n_embd_head_v(il), Vcur->ne[0]/hparams.n_embd_head_v(il), n_tokens); + cb(Vcur, "Vcur", il_cb); + } + Qcur = ggml_reshape_3d(ctx0, Qcur, hparams.n_embd_head_k(il), Qcur->ne[0]/hparams.n_embd_head_k(il), n_tokens); + cb(Qcur, "Qcur", il_cb); + Kcur = ggml_reshape_3d(ctx0, Kcur, hparams.n_embd_head_k(il), Kcur->ne[0]/hparams.n_embd_head_k(il), n_tokens); + cb(Kcur, "Kcur", il_cb); + if (!Vcur) { + Vcur = Kcur; + } + Qcur = llm.llm_build_norm(ctx0, Qcur, hparams, q_norm->splits[id], NULL, LLM_NORM_RMS, cb, il_cb); + cb(Qcur, "Qcur_n", il_cb); + Kcur = llm.llm_build_norm(ctx0, Kcur, hparams, k_norm->splits[id], NULL, LLM_NORM_RMS, cb, il_cb); + cb(Kcur, "Kcur_n", il_cb); + Vcur = ggml_rms_norm(ctx0, Vcur, hparams.f_norm_rms_eps); + cb(Vcur, "Vcur_n", il_cb); + + auto rope_factors = freq_factors ? ((const ggml_split_tensor_t *)freq_factors->extra)->splits[id] : nullptr; + Qcur = ggml_rope_ext(ctx0, Qcur, inp_pos, rope_factors, n_rot_l, llm.rope_type, llm.n_ctx_orig, freq_base_l, freq_scale_l, + llm.ext_factor, llm.attn_factor, llm.beta_fast, llm.beta_slow); + Kcur = ggml_rope_ext(ctx0, Kcur, inp_pos, rope_factors, n_rot_l, llm.rope_type, llm.n_ctx_orig, freq_base_l, freq_scale_l, + llm.ext_factor, llm.attn_factor, llm.beta_fast, llm.beta_slow); + cb(Qcur, "Qcur_rope", il_cb); + cb(Kcur, "Kcur_rope", il_cb); + + const int64_t n_embd_head_k = hparams.n_embd_head_k(il); + const int64_t n_embd_head_v = hparams.n_embd_head_v(il); + const int64_t n_head_kv = wk->splits[id]->ne[1] / n_embd_head_k; + + if (cparams.k_cache_hadamard) { + Qcur = ggml_hadamard(ctx0, Qcur, n_embd_head_k); + Kcur = ggml_hadamard(ctx0, Kcur, n_embd_head_k); + cb(Qcur, "Qcur_h", il_cb); + cb(Kcur, "Kcur_h", il_cb); + } + if (cparams.v_cache_hadamard) { + Vcur = ggml_hadamard(ctx0, Vcur, n_embd_head_v); + cb(Vcur, "Vcur_h", il_cb); + } + + GGML_ASSERT(kv_self.size == cparams.n_ctx); + + ggml_build_forward_expand(gf, Qcur); + ggml_build_forward_expand(gf, Kcur); + ggml_build_forward_expand(gf, Vcur); + + auto idx = 2*n_device*il + 2*id; + GGML_ASSERT(idx+1 < (int)lctx.cache_copies.size()); + auto k_row_size = ggml_row_size(kl->splits[id]->type, n_embd_head_k); + ggml_tensor * k_cache_view = ggml_view_2d(ctx0, kl->splits[id], n_embd_head_k, n_tokens*n_head_kv, + k_row_size, k_row_size*n_head_kv*llm.kv_head); + + lctx.cache_copies[idx+0].cpy = ggml_cpy(ctx0, Kcur, k_cache_view); + cb(lctx.cache_copies[idx+0].cpy, "k_cache", il_cb); + lctx.cache_copies[idx+0].step = k_row_size*n_head_kv; + ggml_build_forward_expand(gf, lctx.cache_copies[idx+0].cpy); + + if (!wv) { + wv = wk; + } + auto v_cache_view = ggml_view_1d(ctx0, vl->splits[id], n_tokens*wv->splits[id]->ne[1], + llm.kv_head*ggml_row_size(vl->splits[id]->type, wv->splits[id]->ne[1])); + lctx.cache_copies[idx+1].step = ggml_row_size(vl->splits[id]->type, wv->splits[id]->ne[1]); + lctx.cache_copies[idx+1].cpy = ggml_cpy(ctx0, Vcur, v_cache_view); + cb(lctx.cache_copies[idx+1].cpy, "v_cache", il_cb); + ggml_build_forward_expand(gf, lctx.cache_copies[idx+1].cpy); + + auto split_kl = kl->splits[id]; + auto split_vl = vl->splits[id]; + + auto q = ggml_permute(ctx0, Qcur, 0, 2, 1, 3); + cb(q, "q", il_cb); + auto k = ggml_view_3d(ctx0, split_kl, n_embd_head_k, llm.n_kv, n_head_kv, + ggml_row_size(split_kl->type, n_embd_head_k)*n_head_kv, + ggml_row_size(split_kl->type, n_embd_head_k), 0); + cb(k, "k", il_cb); + auto v = ggml_view_3d(ctx0, split_vl, n_embd_head_v, llm.n_kv, n_head_kv, + ggml_row_size(split_vl->type, wv->splits[id]->ne[1]), + ggml_row_size(split_vl->type, n_embd_head_v), 0); + cb(v, "v", il_cb); + + cur = ggml_flash_attn_ext(ctx0, q, k, v, KQ_mask_l, hparams.f_attention_scale, hparams.f_max_alibi_bias, + hparams.attn_soft_cap ? hparams.f_attn_logit_softcapping : 0.0f); + cb(cur, "fa", il_cb); + cur->op_params[4] = n_swa; + if (cparams.v_cache_hadamard) { + cur = ggml_hadamard(ctx0, cur, n_embd_head_v); + cb(cur, "fa_h", il_cb); + } + cur = ggml_reshape_2d(ctx0, cur, wo->splits[id]->ne[0], n_tokens); + if (il == hparams.n_layer-1 && inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + sa_inp[id] = ggml_get_rows(ctx0, sa_inp[id], inp_out_ids); + } + cur = llm.llm_build_lora_mm(lctx, ctx0, wo->splits[id], cur); + cb(cur, "qkv", il_cb); + if (cur->ne[1] > 32 && cparams.reduce_type != GGML_TYPE_F32) { + cur = ggml_cast(ctx0, cur, cparams.reduce_type); + cb(cur, "qkv_cast", il_cb); + } + ggml_build_forward_expand(gf, cur); + sa_out[id] = cur; + + } + + auto last_ffn_inp = ggml_reduce(ctx0, sa_out.data(), n_device, GGML_OP_ADD); + ggml_build_forward_expand(gf, last_ffn_inp); + cb(last_ffn_inp, "sa_reduce", il); + + auto ffn_up = (const ggml_split_tensor_t *)l.ffn_up->extra; + auto ffn_gate = (const ggml_split_tensor_t *)l.ffn_gate->extra; + auto ffn_down = (const ggml_split_tensor_t *)l.ffn_down->extra; + GGML_ASSERT(ffn_up && ffn_gate && ffn_down); + + for (int id = 0; id < n_device; ++id) { + GGML_ASSERT((ffn_up->splits[id] && ffn_gate->splits[id] && ffn_down->splits[id]) || + (!ffn_up->splits[id] && !ffn_gate->splits[id] && !ffn_down->splits[id])); + if (!ffn_up->splits[id]) { + ffn_inp[id] = ffn_out[id] = nullptr; + continue; + } + int il_cb = 1000*(il + 1) + id; + + GGML_ASSERT(last_ffn_inp && last_ffn_inp->op == GGML_OP_REDUCE); + auto cur = get_input_tensor_sm_graph(ctx0, last_ffn_inp, id); + cur = do_split_norm(ctx0, cur, model.layers[il].attn_post_norm, hparams, cb, id, il_cb, false); + cb(cur, "sa_post", il_cb); + auto add = sa_inp[id]; + if (!add) { + for (int j = 0; j < n_device; ++j) { + if (sa_inp[j]) { + add = sa_inp[j]; break; + } + } + } + ffn_inp[id] = ggml_add(ctx0, cur, add); + cb(ffn_inp[id], "ffn_inp", il_cb); + cur = do_split_norm(ctx0, ffn_inp[id], model.layers[il].ffn_norm, hparams, cb, id, il_cb, false); + cb(cur, "ffn_inp_normed", il_cb); + cur = llm.llm_build_ffn(ctx0, lctx, nullptr, cur, + ffn_up->splits[id], nullptr, nullptr, + ffn_gate->splits[id], nullptr, nullptr, + ffn_down->splits[id], nullptr, nullptr, + nullptr, + LLM_FFN_GELU, LLM_FFN_PAR, cb, il, gf, false, false, nullptr, nullptr); + cb(cur, "ffn", il_cb); + if (cur->ne[1] > 32 && cparams.reduce_type != GGML_TYPE_F32) { + cur = ggml_cast(ctx0, cur, cparams.reduce_type); + cb(cur, "ffn_cast", il_cb); + } + ggml_build_forward_expand(gf, cur); + ffn_out[id] = cur; + } + + inpL = ggml_reduce(ctx0, ffn_out.data(), n_device, GGML_OP_ADD); + cb(inpL, "ffn_reduce", il); + ggml_build_forward_expand(gf, inpL); + + } + + int idx = lctx.model.default_layer_device[lctx.model.hparams.n_layer]; + int idx_out = ggml_backend_sched_get_backend_idx(lctx.sched, lctx.model.output->buffer); + if (idx_out >= 0) idx = idx_out; + auto cur = inpL->src[idx]; + if (!cur) { + cur = inpL->view_src; + } + + auto post_norm = (const ggml_split_tensor_t *)model.layers[hparams.n_layer-1].ffn_post_norm->extra; + cur = llm.llm_build_norm(ctx0, cur, hparams, post_norm->splits[idx], NULL, LLM_NORM_RMS, cb, -1); + cb(cur, "ffn_post", hparams.n_layer-1); + auto add = ffn_inp[idx]; + if (!add) { + for (int j = 0; j < n_device; ++j) { + if (ffn_inp[j]) { + add = ffn_inp[j]; break; + } + } + } + cur = ggml_add(ctx0, cur, add); + cb(cur, "ffn_out", hparams.n_layer-1); + + if (model.layers[hparams.n_layer-1].out_scale) { + auto scale = (const ggml_split_tensor_t *)model.layers[hparams.n_layer-1].out_scale->extra; + cur = ggml_mul(ctx0, cur, scale->splits[idx]); + cb(cur, "ffn_out_scaled", hparams.n_layer-1); + } + + cur = build_output(lctx, ctx0, cur, model.output, model.output_norm, cb); + if (hparams.f_final_logit_softcapping > 0) { + cur = ggml_softcap(ctx0, cur, 1.0f / hparams.f_final_logit_softcapping, hparams.f_final_logit_softcapping); + } + cb(cur, "result_output", -1); + + ggml_build_forward_expand(gf, cur); + + return gf; +} + static ggml_tensor * gemma4_project_per_layer_inputs(ggml_context * ctx0, const llama_model & model, const llm_build_cb & cb, int n_embd, int n_embd_per_layer, int n_layer, int n_tokens, ggml_tensor * inputs_embeds, ggml_tensor * inp_per_layer) { @@ -6044,7 +6355,7 @@ static ggml_tensor * gemma4_project_per_layer_inputs(ggml_context * ctx0, const } ggml_cgraph * llm_build_context::build_gemma4() { - struct ggml_cgraph * gf = ggml_new_graph_custom(ctx0, llama_model_max_nodes(model, n_tokens), false); + //struct ggml_cgraph * gf = ggml_new_graph_custom(ctx0, llama_model_max_nodes(model, n_tokens), false); struct ggml_tensor * cur; struct ggml_tensor * inpL; @@ -6092,6 +6403,13 @@ ggml_cgraph * llm_build_context::build_gemma4() { } + if (model.split_mode == LLAMA_SPLIT_MODE_GRAPH) { + return build_gemma4_graph_paralle(*this, lctx, ctx0, inpL, inp_pos, inp_out_ids, + KQ_mask, KQ_mask_swa, n_tokens, cb); + } + + auto gf = ggml_new_graph_custom(ctx0, llama_model_max_nodes(model, n_tokens), false); + // "5-to-1 interleaved attention" // 5 layers of local attention followed by 1 layer of global attention // we do this via swa_layers now @@ -6109,13 +6427,16 @@ ggml_cgraph * llm_build_context::build_gemma4() { struct ggml_tensor * KQ_mask_l = is_sliding ? KQ_mask_swa : KQ_mask; - // norm - cur = llm_build_norm(ctx0, inpL, hparams, model.layers[il].attn_norm, NULL, LLM_NORM_RMS, cb, il); - cb(cur, "attn_norm", il); + auto freq_factors = !is_sliding ? model.layers[il].rope_freqs : nullptr; - // self-attention - { - auto freq_factors = !is_sliding ? model.layers[il].rope_freqs : nullptr; + ggml_tensor * attn_out; + + if (hparams.has_kv(il) && model.layers[il].wv) { + attn_out = build_std_attention(gf, model.layers[il].attn_norm, inpL, inp_pos, il == n_layer - 1 ? inp_out_ids : nullptr, freq_factors, + KQ_mask_l, nullptr, nullptr, hparams.f_attention_scale, 0.0f, n_swa, il, true, false, true, false, false, model.layers[il].attn_post_norm); + } else { + cur = llm_build_norm(ctx0, inpL, hparams, model.layers[il].attn_norm, NULL, LLM_NORM_RMS, cb, il); + cb(cur, "attn_norm", il); ggml_tensor *Qcur, *Kcur = nullptr, *Vcur = nullptr; Qcur = llm_build_lora_mm(lctx, ctx0, model.layers[il].wq, cur); @@ -6151,20 +6472,20 @@ ggml_cgraph * llm_build_context::build_gemma4() { cur = llm_build_kv(ctx0, lctx, kv_self, gf, model.layers[il].wo, model.layers[il].bo, Kcur, Vcur, Qcur, KQ_mask_l, n_tokens, kv_head, n_kv, hparams.f_attention_scale, cb, il, nullptr, n_swa); + + if (il == n_layer - 1 && inp_out_ids) { + // skip computing output for unused tokens + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inpL = ggml_get_rows(ctx0, inpL, inp_out_ids); + } + + cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].attn_post_norm, NULL, LLM_NORM_RMS, cb, il); + cb(cur, "attn_post_norm", il); + + attn_out = ggml_add(ctx0, cur, inpL); + cb(attn_out, "attn_out", il); } - if (il == n_layer - 1 && inp_out_ids) { - // skip computing output for unused tokens - cur = ggml_get_rows(ctx0, cur, inp_out_ids); - inpL = ggml_get_rows(ctx0, inpL, inp_out_ids); - } - - cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].attn_post_norm, NULL, LLM_NORM_RMS, cb, il); - cb(cur, "attn_post_norm", il); - - auto attn_out = ggml_add(ctx0, cur, inpL); - cb(attn_out, "attn_out", il); - if (model.layers[il].ffn_gate_inp) { auto cur_mlp = llm_build_norm(ctx0, attn_out, hparams, model.layers[il].ffn_norm, nullptr, LLM_NORM_RMS, cb, il); @@ -6212,25 +6533,23 @@ ggml_cgraph * llm_build_context::build_gemma4() { cur = ggml_add(ctx0, cur_mlp, cur_moe); cb(cur, "ffn_moe_combined", il); - } else { - cur = llm_build_norm(ctx0, attn_out, hparams, model.layers[il].ffn_norm, nullptr, LLM_NORM_RMS, cb, il); - cb(cur, "ffn_norm", il); + cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].ffn_post_norm, NULL, LLM_NORM_RMS, cb, -1); + cb(cur, "ffn_post_norm", -1); - cur = llm_build_ffn(ctx0, lctx, nullptr, cur, + cur = ggml_add(ctx0, cur, attn_out); + + } else { + + cur = llm_build_ffn(ctx0, lctx, model.layers[il].ffn_norm, attn_out, model.layers[il].ffn_up, nullptr, nullptr, model.layers[il].ffn_gate, nullptr, nullptr, model.layers[il].ffn_down, nullptr, nullptr, nullptr, - LLM_FFN_GELU, LLM_FFN_PAR, cb, il, gf); + LLM_FFN_GELU, LLM_FFN_PAR, cb, il, gf, true, false, nullptr, model.layers[il].ffn_post_norm); cb(cur, "ffn_out", il); } - cur = llm_build_norm(ctx0, cur, hparams, model.layers[il].ffn_post_norm, NULL, LLM_NORM_RMS, cb, -1); - cb(cur, "ffn_post_norm", -1); - - cur = ggml_add(ctx0, cur, attn_out); - if (inp_per_layer) { ggml_tensor * pe_in = cur; cb(cur, "pe_in", il); @@ -6257,6 +6576,10 @@ ggml_cgraph * llm_build_context::build_gemma4() { // layer_scalar if (model.layers[il].out_scale) { + //if (ggml_backend_buffer_is_host(model.layers[il].out_scale->buffer)) { + // auto val = (const float *)model.layers[il].out_scale->data; + // printf("Layer %d: out_scale = %g\n", il, val[0]); + //} cur = ggml_mul(ctx0, cur, model.layers[il].out_scale); cb(cur, "out_scaled", il); } @@ -7781,9 +8104,9 @@ ggml_cgraph * llm_build_context::build_glm4_moe() { } struct ggml_tensor * llm_build_context::build_mtp_tail( - const llama_layer & mtp_layer, - struct ggml_tensor * prev_embeddings, - int64_t n_embd_head, + const llama_layer & mtp_layer, + struct ggml_tensor * prev_embeddings, + int64_t n_embd_head, struct ggml_cgraph * gf, struct ggml_tensor * inp_pos, struct ggml_tensor * rope_cache @@ -7803,7 +8126,7 @@ struct ggml_tensor * llm_build_context::build_mtp_tail( ggml_tensor * token_emb_norm = llm_build_norm(ctx0, token_emb, hparams, mtp_layer.nextn.enorm, NULL, LLM_NORM_RMS, cb, il); ggml_tensor * hidden_state_norm = llm_build_norm(ctx0, prev_embeddings, hparams, mtp_layer.nextn.hnorm, NULL, LLM_NORM_RMS, cb, il); - + ggml_tensor * combined = ggml_concat(ctx0, token_emb_norm, hidden_state_norm, 0); cb(combined, "mtp_concat", il); ggml_tensor* cur = llm_build_lora_mm(lctx, ctx0, mtp_layer.nextn.eh_proj, combined); @@ -10248,7 +10571,8 @@ ggml_cgraph * llm_build_context::llama_build_graph( ggml_tensor * llm_build_context::build_std_attention(ggml_cgraph * gf, ggml_tensor * the_attn_norm, ggml_tensor * input, ggml_tensor * inp_pos, ggml_tensor * inp_out_ids, ggml_tensor * rope_factors_in, ggml_tensor * KQ_mask, ggml_tensor * sinks, ggml_tensor * inp_attn_scale, float KQ_scale, float f_attn_scale, - int n_swa, int il, bool do_rope, bool add_graph_split, bool add_input, bool is_norm, bool is_multi) { + int n_swa, int il, bool do_rope, bool add_graph_split, bool add_input, bool is_norm, bool is_multi, + ggml_tensor * post_norm) { float freq_base_l = n_swa > 0 ? hparams.rope_freq_base_train_swa : cparams.rope_freq_base; float freq_scale_l = n_swa > 0 ? hparams.rope_freq_scale_train_swa : hparams.rope_freq_scale_train; @@ -10541,6 +10865,10 @@ ggml_tensor * llm_build_context::build_std_attention(ggml_cgraph * gf, ggml_tens model.layers[il].wq, model.layers[il].bq, model.layers[il].wk, model.layers[il].bk, model.layers[il].wv, model.layers[il].bv, model.layers[il].attn_q_norm, model.layers[il].attn_k_norm, f_attn_scale, il); Qcur = Q; Kcur = K; Vcur = V; + if (model.arch == LLM_ARCH_GEMMA4) { + Vcur = ggml_reshape_3d(ctx0, Vcur, model.hparams.n_embd_head_v(il), model.hparams.n_head_kv(il), n_tokens); + Vcur = ggml_rms_norm(ctx0, Vcur, model.hparams.f_norm_rms_eps); + } } if (do_rope) { @@ -10620,6 +10948,11 @@ ggml_tensor * llm_build_context::build_std_attention(ggml_cgraph * gf, ggml_tens } } + if (post_norm) { + cur = llm_build_norm(ctx0, cur, hparams, post_norm, NULL, LLM_NORM_RMS, cb, il); + cb(cur, "sa_normed", il); + } + if (add_input) { cb(cur, "attn_out", il); cur = ggml_add(ctx0, cur, input); diff --git a/src/llama-build-context.h b/src/llama-build-context.h index 58d68f90..81ae8c5c 100644 --- a/src/llama-build-context.h +++ b/src/llama-build-context.h @@ -356,7 +356,7 @@ struct llm_build_context { llm_ffn_op_type type_op, llm_ffn_gate_type type_gate, const llm_build_cb & cb, int il, ggml_cgraph * graph = nullptr, bool add_input = false, - bool is_norm = false, ggml_tensor * add_extra = nullptr); + bool is_norm = false, ggml_tensor * add_extra = nullptr, ggml_tensor * post_norm = nullptr); static ggml_tensor * llm_build_moe_ffn(ggml_context * ctx, llama_context & lctx, ggml_tensor * cur, @@ -440,7 +440,7 @@ llm_expert_gating_func_type gating_op, ggml_tensor * inp_pos, ggml_tensor * inp_out_ids, ggml_tensor * rope_factors, ggml_tensor * KQ_mask, ggml_tensor * sinks, ggml_tensor * inp_attn_scale, float KQ_scale, float f_attn_scale, int n_swa, int il, bool do_rope = true, bool add_graph_split = false, bool add_input = false, bool is_norm = false, - bool is_multi = false); + bool is_multi = false, ggml_tensor * post_norm = nullptr); static uint32_t llama_kv_qnext_state_slots(const llama_kv_cache & kv_self); diff --git a/src/llama-hparams.h b/src/llama-hparams.h index b5f481c0..7202ea84 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -332,7 +332,7 @@ struct llama_hparams { uint32_t rope_n_rot(uint32_t il) const { const uint32_t v = rope_dim_per_layer[il]; - return v ? v : n_rot; + return v ? v : swa_layers[il] ? n_rot_swa : n_rot; } static const char * rope_scaling_type_name(llama_rope_scaling_type); diff --git a/src/llama-load-tensors.cpp b/src/llama-load-tensors.cpp index 2adbee0f..3fa7bb86 100644 --- a/src/llama-load-tensors.cpp +++ b/src/llama-load-tensors.cpp @@ -4052,10 +4052,31 @@ bool create_tensors_helper::create_tensors() { use_mmap_buffer &= !has_buft_overrides; } + if (model.arch == LLM_ARCH_GEMMA4 && (model.split_mode == LLAMA_SPLIT_MODE_GRAPH || model.split_mode == LLAMA_SPLIT_MODE_ATTN)) { + bool supported = true; + if (model.tok_embd_per_layer) { + supported = false; + } + for (auto & l : model.layers) { + if (l.ffn_gate_inp) { + supported = false; + break; + } + } + if (!supported) { + LLAMA_LOG_WARN("\n=========================================================\n"); + LLAMA_LOG_WARN("Split mode 'graph' is not supported for this Gemma4 variant\n"); + LLAMA_LOG_WARN(" => changing split mode to 'layer'\n"); + LLAMA_LOG_WARN("===========================================================\n\n"); + model.split_mode = LLAMA_SPLIT_MODE_LAYER; + } + } + if (model.split_mode == LLAMA_SPLIT_MODE_GRAPH || model.split_mode == LLAMA_SPLIT_MODE_ATTN) { const int n_layer = model.mtp ? model.layers.size() : model.layers.size() - model.hparams.nextn_predict_layers; LLAMA_LOG_INFO("================================ max_gpu = %d\n", model.max_gpu); + std::vector mirror(model.splits.size(), 1); std::vector mem_used(model.splits.size(), 0); const auto & hparams = model.hparams; auto cur_splits = model.splits; @@ -4110,8 +4131,10 @@ bool create_tensors_helper::create_tensors() { auto & layer = model.layers[il]; auto ctx_split = ctx_for_layer_split(il); if (layer.attn_norm) { - auto split = create_split(ggml_nrows(layer.attn_norm), -1, cur_splits, mem_used); - prepare_split_tensors(-1, ctx_split, layer.attn_norm, layer.split_attn_norm, split, mem_used); + prepare_split_tensors(-1, ctx_split, layer.attn_norm, layer.split_attn_norm, mirror, mem_used); + } + if (model.arch == LLM_ARCH_GEMMA4 && layer.attn_post_norm) { + prepare_split_tensors(-1, ctx_split, layer.attn_post_norm, layer.split_attn_post_norm, mirror, mem_used); } if (layer.rope_freqs) { auto split = create_split(ggml_nrows(layer.rope_freqs), -1, cur_splits, mem_used); @@ -4120,7 +4143,7 @@ bool create_tensors_helper::create_tensors() { if (hparams.is_recurrent(il)) { split_recurrent_tensors(hparams, layer, cur_splits, mem_used, ctx_split, il); //, model.arch == LLM_ARCH_QWEN3NEXT ? 0 : 1); } - else if (layer.wo && layer.wq && layer.wk && layer.wv) { + else if (layer.wo && layer.wq && layer.wk && (layer.wv || model.arch == LLM_ARCH_GEMMA4)) { auto granularity_kq = hparams.n_embd_head_k(il) * gqa_ratio; int wq_ne1 = layer.wq->ne[1]; if (model.arch == LLM_ARCH_QWEN3NEXT || model.arch == LLM_ARCH_QWEN35MOE || model.arch == LLM_ARCH_QWEN35) { @@ -4227,16 +4250,22 @@ bool create_tensors_helper::create_tensors() { } } } - prepare_split_tensors(1, ctx_split, layer.wv, layer.split_wv, split_vo, mem_used); - if (layer.bv) { - prepare_split_tensors(0, ctx_split, layer.bv, layer.split_bv, split_vo, mem_used); + if (layer.wv) { + prepare_split_tensors(1, ctx_split, layer.wv, layer.split_wv, split_vo, mem_used); + if (layer.bv) { + prepare_split_tensors(0, ctx_split, layer.bv, layer.split_bv, split_vo, mem_used); + } } } if (layer.ffn_norm) { if (auto it = split_tensors.find(layer.ffn_norm); it != split_tensors.end()) { - auto split = create_split(ggml_nrows(layer.ffn_norm), -1, cur_splits, mem_used); - prepare_split_tensors(-1, ctx_split, layer.ffn_norm, layer.split_ffn_norm, split, mem_used); + prepare_split_tensors(-1, ctx_split, layer.ffn_norm, layer.split_ffn_norm, mirror, mem_used); + } + } + if (layer.ffn_post_norm) { + if (auto it = split_tensors.find(layer.ffn_post_norm); it != split_tensors.end()) { + prepare_split_tensors(-1, ctx_split, layer.ffn_post_norm, layer.split_ffn_post_norm, mirror, mem_used); } } @@ -4367,6 +4396,10 @@ bool create_tensors_helper::create_tensors() { prepare_split_tensors(-1, ctx_split, layer.ffn_exp_probs_b, layer.split_ffn_exp_probs_b, shared_split, mem_used); } } + + if (layer.out_scale) { + prepare_split_tensors(-1, ctx_split, layer.out_scale, layer.split_out_scale, std::vector(model.splits.size(), 1), mem_used); + } } if (!gpu_split_count.empty()) { diff --git a/src/llama-model.h b/src/llama-model.h index b48e1f86..9f4ce181 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -198,6 +198,7 @@ struct llama_layer { struct ggml_tensor * bkv = nullptr; llama_split_tensor split_attn_norm; + llama_split_tensor split_attn_post_norm; llama_split_tensor split_attn_sinks; llama_split_tensor split_wq; llama_split_tensor split_wk; @@ -257,6 +258,7 @@ struct llama_layer { llama_split_tensor split_ffn_down; llama_split_tensor split_ffn_norm; llama_split_tensor split_ffn_up_gate; + llama_split_tensor split_ffn_post_norm; // ff MoE struct ggml_tensor * ffn_gate_inp = nullptr; @@ -360,6 +362,8 @@ struct llama_layer { struct ggml_tensor * ffn_down_scale = nullptr; struct ggml_tensor * out_scale = nullptr; // gemma4 layer output scale + llama_split_tensor split_out_scale; + struct llama_layer_nextn nextn; std::unique_ptr computed_wk_b; diff --git a/src/llama.cpp b/src/llama.cpp index 31be5479..61d62589 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -918,6 +918,9 @@ static bool llama_kv_cache_init( bool split_cache_i = split_cache; auto K = model.layers[i].wk; auto V = model.layers[i].wv; + if (!V && model.arch == LLM_ARCH_GEMMA4) { + V = K; + } if (split_cache && (!K || !V || !K->extra || !V->extra)) { ctx = offload ? ctx_map.at(model.buft_layer[i].buft) : cache.ctxs.front(); split_cache_i = false; @@ -2000,6 +2003,7 @@ static bool is_model_split_supported(const llama_model & model) { //LLM_ARCH_QWEN3NEXT, LLM_ARCH_QWEN35, LLM_ARCH_QWEN35MOE, + LLM_ARCH_GEMMA4, }; auto it = k_supported.find(model.arch); return it != k_supported.end(); @@ -6389,6 +6393,7 @@ bool llama_save_session_file(struct llama_context * ctx, const char * path_sessi } static inline ggml_tensor * get_kv_cache_split_tensor(const ggml_tensor * tensor, const llama_layer & l) { + if (!l.wv) return l.wk; bool use_V_for_K = l.attn_k_norm && l.attn_k_norm->ne[0] == l.wk->ne[1] ? true : false; auto kv = tensor->ne[1] > 1 && !use_V_for_K ? l.wk : l.wv; return kv;