Graph parallel for Gemma4-31B (#1596)

* Use build_std_attention for Gemma4 when possible

It is possible for the 26b MoE and 31b dense models.
It is not possible for the E4B/E2B vaiants because they
don't have KV cache in each layer.

* Standardize Gemma4 dense ffn

* WIP: Gemma4 split mode graph

Runs but produces NaNs

* WIP: Gemma4 split mode graph

Runs but very high PPL. At least it is no longer NaN.

* WIP

* This works!

* Put attn_norm, attn_post_norm, ffn_norm, ffn_post_norm on all GPUs

* Fix crash when saving/loading KV cache
This commit is contained in:
Kawrakow 2026-04-09 08:00:22 +02:00 committed by GitHub
parent fac404509c
commit 5950d0259e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
6 changed files with 421 additions and 46 deletions

View File

@ -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<ggml_tensor *> sa_inp(n_device, nullptr);
std::vector<ggml_tensor *> sa_out(n_device, nullptr);
std::vector<ggml_tensor *> ffn_inp(n_device, nullptr);
std::vector<ggml_tensor *> 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);

View File

@ -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);

View File

@ -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);

View File

@ -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<int> mirror(model.splits.size(), 1);
std::vector<size_t> 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<int>(model.splits.size(), 1), mem_used);
}
}
if (!gpu_split_count.empty()) {

View File

@ -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<ggml_tensor> computed_wk_b;

View File

@ -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;