mirror of
https://github.com/ikawrakow/ik_llama.cpp.git
synced 2026-07-21 10:15:35 +00:00
Support for alternative Gemma4 assistant (#1937)
This commit is contained in:
@@ -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;
|
||||
|
||||
@@ -80,6 +80,7 @@ enum llm_arch {
|
||||
LLM_ARCH_MISTRAL4,
|
||||
LLM_ARCH_GEMMA4,
|
||||
LLM_ARCH_GEMMA4_MTP,
|
||||
LLM_ARCH_GEMMA4_ASSISTANT,
|
||||
LLM_ARCH_UNKNOWN,
|
||||
};
|
||||
|
||||
|
||||
@@ -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
@@ -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);
|
||||
|
||||
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user