WIP: split mode graph

This commit is contained in:
Kawrakow
2026-08-11 12:44:02 +00:00
parent b39b20f0b8
commit 21160cac03
4 changed files with 47 additions and 9 deletions
+35 -4
View File
@@ -1112,7 +1112,8 @@ ggml_tensor * llm_build_context::llm_build_ffn(
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,
ggml_tensor * post_norm, float post_norm_eps) {
ggml_tensor * post_norm, float post_norm_eps,
post_norm_data * pnd) {
if (!up_b && !up_s && !gate_b && !gate_s && !down_b && !down_s &&
up->extra && gate->extra && down->extra && type_gate == LLM_FFN_PAR &&
@@ -1134,6 +1135,17 @@ ggml_tensor * llm_build_context::llm_build_ffn(
GGML_ASSERT((!split_u && !split_g && !split_d) || (split_u && split_g && split_d));
if (!split_u) continue;
auto cur = get_input_tensor_sm_graph(ctx, input, id);
if (pnd) {
auto pn_extra = (ggml_split_tensor_t *)pnd->norm->extra;
GGML_ASSERT(pn_extra && pn_extra->splits[id]);
cur = ggml_fused_rms_norm(ctx, cur, pn_extra->splits[id], pnd->f_rms_eps);
cb(cur, "ffn_post_norm", il_cb);
if (pnd->add) {
auto add_id = get_input_tensor_sm_graph(ctx, pnd->add, id);
cur = ggml_add(ctx, cur, add_id);
cb(cur, "inp_added", il_cb);
}
}
cur = do_split_norm(ctx, cur, ffn_norm, lctx.model.hparams, cb, id, il_cb, is_norm);
if (input->op != GGML_OP_REDUCE) {
cur->op_params[GGML_MAX_OP_PARAMS / sizeof(int32_t) - 1] = 0xff;
@@ -3018,7 +3030,7 @@ ggml_tensor * llm_build_context::build_std_attention(ggml_cgraph * gf, ggml_tens
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,
ggml_tensor * post_norm, int kv_il, float post_norm_eps) {
ggml_tensor * post_norm, int kv_il, float post_norm_eps, post_norm_data * pnd) {
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;
@@ -3098,6 +3110,17 @@ ggml_tensor * llm_build_context::build_std_attention(ggml_cgraph * gf, ggml_tens
(split_wq && split_wk && split_wv && split_wo && split_kl && split_vl));
if (!split_wq) continue;
auto cur = get_input_tensor_sm_graph(ctx0, input, id);
if (pnd) {
auto pn_extra = (ggml_split_tensor_t *)pnd->norm->extra;
GGML_ASSERT(pn_extra && pn_extra->splits[id]);
cur = ggml_fused_rms_norm(ctx0, cur, pn_extra->splits[id], pnd->f_rms_eps);
cb(cur, "att_post_norm", il_cb);
if (pnd->add) {
auto add_id = get_input_tensor_sm_graph(ctx0, pnd->add, id);
cur = ggml_add(ctx0, cur, add_id);
cb(cur, "inp_added", il_cb);
}
}
cur = do_split_norm(ctx0, cur, the_attn_norm, lctx.model.hparams, cb, id, il_cb, is_norm);
auto input_normed = cur;
auto the_q_norm = model.layers[il].attn_q_norm ? model.layers[il].attn_q_norm->extra ?
@@ -3270,8 +3293,15 @@ ggml_tensor * llm_build_context::build_std_attention(ggml_cgraph * gf, ggml_tens
cur = ggml_mul(ctx0, cur, gate);
}
} else {
auto gate_3d = ggml_reshape_3d(ctx0, gate, 1, nh, n_tokens);
cur = ggml_fused_mul_unary(ctx0, gate_3d, attn_3d, GGML_UNARY_OP_SIGMOID);
if (gate->ne[0] == n_embd_head_v * nh) {
gate = ggml_sigmoid(ctx0, gate);
cb(gate, "gate", il_cb);
gate = ggml_reshape_3d(ctx0, gate, cur->ne[0], cur->ne[1], cur->ne[2]);
cur = ggml_mul(ctx0, cur, gate);
} else {
auto gate_3d = ggml_reshape_3d(ctx0, gate, 1, nh, n_tokens);
cur = ggml_fused_mul_unary(ctx0, gate_3d, attn_3d, GGML_UNARY_OP_SIGMOID);
}
}
cb(attn_3d, "attn_gated_3d", il_cb);
}
@@ -3305,6 +3335,7 @@ ggml_tensor * llm_build_context::build_std_attention(ggml_cgraph * gf, ggml_tens
cb(cur, "kqv_wo_biased", il_cb);
output_bias_added = true;
}
if (cur->ne[1] > 32 && lctx.cparams.reduce_type != GGML_TYPE_F32) {
cur = ggml_cast(ctx0, cur, lctx.cparams.reduce_type);
}
+10 -2
View File
@@ -37,6 +37,12 @@ enum llm_norm_type {
LLM_NORM_RMS,
};
struct post_norm_data {
ggml_tensor * norm;
ggml_tensor * add;
float f_rms_eps;
};
struct llm_build_context {
const llama_model & model;
llama_context & lctx;
@@ -508,7 +514,8 @@ struct llm_build_context {
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,
ggml_tensor * post_norm = nullptr, float post_norm_eps = 0.0f);
ggml_tensor * post_norm = nullptr, float post_norm_eps = 0.0f,
post_norm_data * pnd = nullptr);
static ggml_tensor * build_dspark_logits(llm_build_context & llm,
ggml_tensor * base_logits, ggml_tensor * input_tokens,
@@ -599,7 +606,8 @@ 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, ggml_tensor * post_norm = nullptr, int kv_il = -1, float post_norm_eps = 0.0f);
bool is_multi = false, ggml_tensor * post_norm = nullptr, int kv_il = -1, float post_norm_eps = 0.0f,
post_norm_data * pnd = nullptr);
static ggml_tensor * build_output(llama_context & lctx, ggml_context * ctx, ggml_tensor * cur, ggml_tensor * output, const llm_build_cb & cb);
+1 -3
View File
@@ -5376,9 +5376,7 @@ bool create_tensors_helper::create_tensors() {
}
if (layer.wqkv_gate) {
auto wqkv_gate_split = split_kq;
if (model.arch == LLM_ARCH_LAGUNA && layer.wqkv_gate->ne[1] == layer.wo->ne[0]) {
// Full-width Laguna M.1 gates follow the value/output partition.
// Head-wise gates still follow the K/Q partition collapsed by head size.
if (layer.wqkv_gate->ne[1] == layer.wo->ne[0]) {
wqkv_gate_split = split_vo;
} else {
for (auto & s : wqkv_gate_split) s /= hparams.n_embd_head_k(il);
+1
View File
@@ -3884,6 +3884,7 @@ static bool is_model_split_supported(const llama_model & model) {
LLM_ARCH_MISTRAL4,
LLM_ARCH_MELLUM,
LLM_ARCH_LAGUNA,
LLM_ARCH_MUSE_GLIMMER,
};
auto it = k_supported.find(model.arch);
return it != k_supported.end();