Support for alternative Gemma4 assistant (#1937)

This commit is contained in:
Kawrakow
2026-06-09 09:30:12 +02:00
committed by GitHub
parent a38d29232d
commit 11c3546235
8 changed files with 74 additions and 21 deletions
+4
View File
@@ -81,10 +81,14 @@ static const std::map<llm_arch, const char *> LLM_ARCH_NAMES = {
{ LLM_ARCH_MISTRAL4, "mistral4" },
{ LLM_ARCH_GEMMA4, "gemma4" },
{ LLM_ARCH_GEMMA4_MTP, "gemma4_mtp" },
{ LLM_ARCH_GEMMA4_ASSISTANT,"gemma4_assistant" },
{ LLM_ARCH_UNKNOWN, "(unknown)" },
};
llm_arch llm_arch_from_string(const std::string & name) {
//if (name == "gemma4_assistant") {
// return llm_arch_from_string("gemma4_mtp");
//}
for (const auto & kv : LLM_ARCH_NAMES) { // NOLINT
if (kv.second == name) {
return kv.first;
+1
View File
@@ -80,6 +80,7 @@ enum llm_arch {
LLM_ARCH_MISTRAL4,
LLM_ARCH_GEMMA4,
LLM_ARCH_GEMMA4_MTP,
LLM_ARCH_GEMMA4_ASSISTANT,
LLM_ARCH_UNKNOWN,
};
+1
View File
@@ -2377,6 +2377,7 @@ ggml_cgraph * llm_build_context::llama_build_graph(
result = llm.build_gemma4();
} break;
case LLM_ARCH_GEMMA4_MTP:
case LLM_ARCH_GEMMA4_ASSISTANT:
{
result = llm.build_gemma4_mtp();
} break;
+10 -3
View File
@@ -780,11 +780,18 @@ void llm_load_hparams(
}
} break;
case LLM_ARCH_GEMMA4_MTP:
case LLM_ARCH_GEMMA4_ASSISTANT:
{
ml.get_key(LLM_KV_MTP_BACKBONE_EMBEDDING_LENGTH, hparams.mtp_backbone_n_embd);
if (model.arch == LLM_ARCH_GEMMA4_MTP) {
ml.get_key(LLM_KV_MTP_BACKBONE_EMBEDDING_LENGTH, hparams.mtp_backbone_n_embd);
ml.get_key(LLM_KV_MTP_CENTROID_COUNT, hparams.mtp_num_centroids, false);
ml.get_key(LLM_KV_MTP_CENTROID_TOP_K, hparams.mtp_centroid_top_k, false);
} else {
ml.get_key("gemma4_assistant.n_embd_backbone", hparams.mtp_backbone_n_embd);
ml.get_key("gemma4_assistant.n_centroids", hparams.mtp_num_centroids, false);
ml.get_key("gemma4_assistant.centroid_top_k", hparams.mtp_centroid_top_k, false);
}
ml.get_key(LLM_KV_MTP_USE_ORDERED_EMBEDDINGS, hparams.mtp_use_ordered_embeddings, false);
ml.get_key(LLM_KV_MTP_CENTROID_COUNT, hparams.mtp_num_centroids, false);
ml.get_key(LLM_KV_MTP_CENTROID_TOP_K, hparams.mtp_centroid_top_k, false);
ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa);
ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps);
+16 -5
View File
@@ -2204,11 +2204,19 @@ bool create_tensors_helper::create_gemma4_mtp_tensors(const LLM_TN & tn) {
if (model.output == NULL) {
model.output = create_tensor(ctx_output, tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, llama_model_loader::TENSOR_DUPLICATED);
}
model.mtp_pre_proj = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_PRE_PROJ, "weight"), {2*n_backbone, n_embd}, 0);
model.mtp_post_proj = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_POST_PROJ, "weight"), {n_embd, n_backbone}, 0);
if (model.arch == LLM_ARCH_GEMMA4_MTP) {
model.mtp_pre_proj = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_PRE_PROJ, "weight"), {2*n_backbone, n_embd}, 0);
model.mtp_post_proj = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_POST_PROJ, "weight"), {n_embd, n_backbone}, 0);
model.mtp_token_ordering = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_TOKEN_ORDERING, "weight"), {n_vocab}, llama_model_loader::TENSOR_NOT_REQUIRED);
model.mtp_centroids = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_CENTROIDS, "weight"), {n_embd, hparams.mtp_num_centroids}, llama_model_loader::TENSOR_NOT_REQUIRED);
} else {
model.mtp_pre_proj = create_tensor(ctx_output, "mtp.pre_projection.weight", {2*n_backbone, n_embd}, 0);
model.mtp_post_proj = create_tensor(ctx_output, "mtp.post_projection.weight", {n_embd, n_backbone}, 0);
model.mtp_token_ordering = create_tensor(ctx_output, "mtp.token_ordering.weight", {n_vocab}, llama_model_loader::TENSOR_NOT_REQUIRED);
printf("========================== hparams.mtp_num_centroids = %d\n", hparams.mtp_num_centroids);
model.mtp_centroids = create_tensor(ctx_output, "mtp.centroids.weight", {n_embd, hparams.mtp_num_centroids}, llama_model_loader::TENSOR_NOT_REQUIRED);
}
model.mtp_token_ordering = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_TOKEN_ORDERING, "weight"), {n_vocab}, llama_model_loader::TENSOR_NOT_REQUIRED);
model.mtp_centroids = create_tensor(ctx_output, tn(LLM_TENSOR_MTP_CENTROIDS, "weight"), {n_embd, hparams.mtp_num_centroids}, llama_model_loader::TENSOR_NOT_REQUIRED);
for (int i = 0; i < n_layer; ++i) {
ggml_context * ctx_layer = ctx_for_layer(i);
@@ -2218,6 +2226,8 @@ bool create_tensors_helper::create_gemma4_mtp_tensors(const LLM_TN & tn) {
const int64_t n_embd_head = hparams.n_embd_head_k(i);
const int64_t n_ff_cur = hparams.n_ff(i);
layer.rope_freqs = create_tensor(ctx_layer, tn(LLM_TENSOR_ROPE_FREQS, "weight"), {n_rot/2}, llama_model_loader::TENSOR_NOT_REQUIRED | (i != 0 ? llama_model_loader::TENSOR_DUPLICATED : 0));
layer.attn_norm = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0);
layer.wq = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd_head*n_head}, 0);
layer.wo = create_tensor(ctx_split, tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd_head*n_head, n_embd}, 0);
@@ -4308,6 +4318,7 @@ bool create_tensors_helper::create_tensors() {
case LLM_ARCH_GEMMA4:
use_mmap_buffer = create_gemma4_tensors(tn); break;
case LLM_ARCH_GEMMA4_MTP:
case LLM_ARCH_GEMMA4_ASSISTANT:
use_mmap_buffer = create_gemma4_mtp_tensors(tn); break;
case LLM_ARCH_STARCODER2:
use_mmap_buffer = create_starcoder2_tensors(tn); break;
@@ -4382,7 +4393,7 @@ bool create_tensors_helper::create_tensors() {
{
const bool unsupported =
(model.arch == LLM_ARCH_GEMMA4_MTP) ||
(model.arch == LLM_ARCH_GEMMA4_MTP || model.arch == LLM_ARCH_GEMMA4_ASSISTANT) ||
(model.arch == LLM_ARCH_GEMMA4 && model.tok_embd_per_layer);
if (unsupported && (model.split_mode == LLAMA_SPLIT_MODE_GRAPH || model.split_mode == LLAMA_SPLIT_MODE_ATTN)) {
LLAMA_LOG_WARN("\n=========================================================\n");
+24 -1
View File
@@ -845,6 +845,29 @@ static const std::map<llm_arch, std::map<llm_tensor, std::string>> LLM_TENSOR_NA
{ LLM_TENSOR_MTP_CENTROIDS, "mtp_centroids" },
},
},
{
LLM_ARCH_GEMMA4_ASSISTANT,
{
{ LLM_TENSOR_TOKEN_EMBD, "token_embd" },
{ LLM_TENSOR_OUTPUT_NORM, "output_norm" },
{ LLM_TENSOR_ATTN_NORM, "blk.%d.attn_norm" },
{ LLM_TENSOR_ATTN_Q, "blk.%d.attn_q" },
{ LLM_TENSOR_ATTN_Q_NORM, "blk.%d.attn_q_norm" },
{ LLM_TENSOR_ATTN_OUT, "blk.%d.attn_output" },
{ LLM_TENSOR_ATTN_POST_NORM, "blk.%d.post_attention_norm" },
{ LLM_TENSOR_FFN_NORM, "blk.%d.ffn_norm" },
{ LLM_TENSOR_FFN_GATE, "blk.%d.ffn_gate" },
{ LLM_TENSOR_FFN_DOWN, "blk.%d.ffn_down" },
{ LLM_TENSOR_FFN_UP, "blk.%d.ffn_up" },
{ LLM_TENSOR_FFN_POST_NORM, "blk.%d.post_ffw_norm" },
{ LLM_TENSOR_LAYER_OUT_SCALE, "blk.%d.layer_output_scale" },
{ LLM_TENSOR_MTP_PRE_PROJ, "mtp_pre_proj" },
{ LLM_TENSOR_MTP_POST_PROJ, "mtp_post_proj" },
{ LLM_TENSOR_MTP_TOKEN_ORDERING, "mtp_token_ordering" },
{ LLM_TENSOR_MTP_CENTROIDS, "mtp_centroids" },
{ LLM_TENSOR_ROPE_FREQS, "rope_freqs" },
},
},
{
LLM_ARCH_STARCODER2,
{
@@ -1958,7 +1981,7 @@ bool llama_model_has_recurrent(const llama_model * model) {
}
bool llama_model_is_gemma4_mtp_assistant(const llama_model * model) {
return model && model->arch == LLM_ARCH_GEMMA4_MTP;
return model && (model->arch == LLM_ARCH_GEMMA4_MTP || model->arch == LLM_ARCH_GEMMA4_ASSISTANT);
}
bool llama_is_gemma4_mtp_file(const char * path) {
+2 -2
View File
@@ -11,7 +11,7 @@ uint32_t llama_mtp_state_n_embd(const struct llama_context * ctx) {
}
const auto & hparams = ctx->model.hparams;
if (ctx->cparams.mtp && ctx->model.arch == LLM_ARCH_GEMMA4_MTP && hparams.mtp_backbone_n_embd > 0) {
if (ctx->cparams.mtp && (ctx->model.arch == LLM_ARCH_GEMMA4_MTP || ctx->model.arch == LLM_ARCH_GEMMA4_ASSISTANT) && hparams.mtp_backbone_n_embd > 0) {
return hparams.mtp_backbone_n_embd;
}
@@ -179,4 +179,4 @@ bool llama_spec_copy_hidden_rows_from_output_indices(
}
return hidden_rows.size() == (size_t) output_indices.size() * view.width;
}
}
+16 -10
View File
@@ -568,7 +568,7 @@ void llama_context::reset_scheduler() {
bool llama_context::can_reuse_graph(const llama_batch & u_batch) {
if (!cparams.graph_reuse) return false;
//if (kv_self.save_per_step_ssm) return false;
if (model.arch == LLM_ARCH_GEMMA4_MTP && mtp_target_ctx != nullptr) return false;
if ((model.arch == LLM_ARCH_GEMMA4_MTP || model.arch == LLM_ARCH_GEMMA4_ASSISTANT) && mtp_target_ctx != nullptr) return false;
auto the_prev = cparams.mtp_op_type == MTP_OP_NONE ? prev.get() : prev_mtp.get();
if (!the_prev || !the_prev->graph) return false;
//if (u_batch.n_tokens > 1) return false;
@@ -3349,12 +3349,13 @@ static bool llm_load_tensors(
if (split_mode == LLAMA_SPLIT_MODE_GRAPH || split_mode == LLAMA_SPLIT_MODE_ATTN) {
const bool unsupported_gemma_split =
model.arch == LLM_ARCH_GEMMA4_MTP ||
model.arch == LLM_ARCH_GEMMA4_ASSISTANT ||
(model.arch == LLM_ARCH_GEMMA4 && hparams.n_embd_per_layer > 0);
if (unsupported_gemma_split) {
LLAMA_LOG_WARN("\n=========================================================\n");
LLAMA_LOG_WARN("Split mode 'graph' is not supported for %s\n",
model.arch == LLM_ARCH_GEMMA4_MTP ? "Gemma 4 MTP assistant"
(model.arch == LLM_ARCH_GEMMA4_MTP || model.arch == LLM_ARCH_GEMMA4_ASSISTANT) ? "Gemma 4 MTP assistant"
: "this Gemma4 variant");
LLAMA_LOG_WARN(" => changing split mode to 'layer'\n");
LLAMA_LOG_WARN("===========================================================\n\n");
@@ -3710,7 +3711,7 @@ static bool llm_load_tensors(
}
}
}
if (model.arch == LLM_ARCH_GEMMA4_MTP && split_mode == LLAMA_SPLIT_MODE_LAYER && device_count > 0 && n_gpu_layers > 0) {
if ((model.arch == LLM_ARCH_GEMMA4_MTP || model.arch == LLM_ARCH_GEMMA4_ASSISTANT) && split_mode == LLAMA_SPLIT_MODE_LAYER && device_count > 0 && n_gpu_layers > 0) {
const int mtp_device = std::clamp(main_gpu, 0, device_count - 1);
LLAMA_LOG_INFO("%s: Gemma 4 MTP assistant forcing layer placement to GPU %d under layer split\n",
@@ -4325,7 +4326,7 @@ static void llama_set_inputs(llama_context & lctx, const llama_batch & batch) {
// NOTE: hparams.causal_attn indicates the model is capable of generation and uses the kv cache.
if (cparams.causal_attn && !lctx.is_encoding) {
const llama_kv_cache & mask_kv_self =
(lctx.model.arch == LLM_ARCH_GEMMA4_MTP && lctx.mtp_target_ctx != nullptr)
((lctx.model.arch == LLM_ARCH_GEMMA4_MTP || lctx.model.arch == LLM_ARCH_GEMMA4_ASSISTANT) && lctx.mtp_target_ctx != nullptr)
? lctx.mtp_target_ctx->kv_self
: kv_self;
const int64_t n_kv = mask_kv_self.n;
@@ -4852,7 +4853,8 @@ static bool llama_context_has_mtp_outputs(const llama_context & lctx) {
return lctx.cparams.mtp && (
lctx.model.hparams.nextn_predict_layers > 0 ||
lctx.model.arch == LLM_ARCH_GEMMA4 ||
lctx.model.arch == LLM_ARCH_GEMMA4_MTP);
lctx.model.arch == LLM_ARCH_GEMMA4_MTP ||
lctx.model.arch == LLM_ARCH_GEMMA4_ASSISTANT);
}
static size_t llama_output_reserve(llama_context & lctx, size_t n_outputs) {
@@ -5304,7 +5306,7 @@ static int llama_decode_internal(
#endif
//if (u_batch.n_tokens == 1 && u_batch.embd == nullptr && lctx.cparams.graph_reuse) {
if (u_batch.embd == nullptr && lctx.cparams.graph_reuse &&
!(lctx.model.arch == LLM_ARCH_GEMMA4_MTP && lctx.mtp_target_ctx != nullptr)) {
!((lctx.model.arch == LLM_ARCH_GEMMA4_MTP || lctx.model.arch == LLM_ARCH_GEMMA4_ASSISTANT) && lctx.mtp_target_ctx != nullptr)) {
prev = std::make_unique<llama_context::Prev>(llama_context::Prev{
(int)u_batch.all_seq_id, (int)lctx.n_outputs, (int)lctx.kv_self.n,
(int)u_batch.n_tokens,
@@ -5332,10 +5334,9 @@ static int llama_decode_internal(
}
else {
const bool has_mtp = llama_context_has_mtp_outputs(lctx);
const bool use_raw_mtp_embd = has_mtp && (lctx.model.arch == LLM_ARCH_QWEN35 ||
lctx.model.arch == LLM_ARCH_QWEN35MOE ||
lctx.model.arch == LLM_ARCH_GEMMA4 ||
lctx.model.arch == LLM_ARCH_GEMMA4_MTP);
const bool use_raw_mtp_embd = has_mtp && (lctx.model.arch == LLM_ARCH_GEMMA4 ||
lctx.model.arch == LLM_ARCH_GEMMA4_MTP||
lctx.model.arch == LLM_ARCH_GEMMA4_ASSISTANT);
if (cparams.embeddings || has_mtp) {
for (int i = gf->n_nodes - 1; i >= 0; --i) {
if (use_raw_mtp_embd && strcmp(gf->nodes[i]->name, "result_mtp_embd") == 0) {
@@ -5347,6 +5348,9 @@ static int llama_decode_internal(
embd = gf->nodes[i];
break;
}
// Strictly speaking we should use if (!use_raw_mtp_embd && strcmp(gf->nodes[i]->name, "result_norm") == 0)
// as Gemma4 MTP is supposed to be using embeddings before rms_norm.
// I don't see any significant difference between this and what we had before, so not making the change (yet).
if (strcmp(gf->nodes[i]->name, "result_norm") == 0) {
embd = gf->nodes[i];
break;
@@ -6886,6 +6890,7 @@ struct llama_context * llama_init_from_model(
if (model->arch != LLM_ARCH_GLM4_MOE && model->arch != LLM_ARCH_QWEN35 &&
model->arch != LLM_ARCH_QWEN35MOE && model->arch != LLM_ARCH_GEMMA4 &&
model->arch != LLM_ARCH_GEMMA4_MTP && model->arch != LLM_ARCH_GLM_DSA &&
model->arch != LLM_ARCH_GEMMA4_ASSISTANT &&
cparams.mtp != 0) {
cparams.mtp = 0;
}
@@ -7434,6 +7439,7 @@ enum llama_rope_type llama_rope_type(const struct llama_model * model) {
case LLM_ARCH_LAGUNA:
case LLM_ARCH_GEMMA4:
case LLM_ARCH_GEMMA4_MTP:
case LLM_ARCH_GEMMA4_ASSISTANT:
return LLAMA_ROPE_TYPE_NEOX;
case LLM_ARCH_QWEN2VL: