diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index a5aa9fa89..b93fb66b7 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -81,10 +81,14 @@ static const std::map 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; diff --git a/src/llama-arch.h b/src/llama-arch.h index cc82d63a7..02a4b996e 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -80,6 +80,7 @@ enum llm_arch { LLM_ARCH_MISTRAL4, LLM_ARCH_GEMMA4, LLM_ARCH_GEMMA4_MTP, + LLM_ARCH_GEMMA4_ASSISTANT, LLM_ARCH_UNKNOWN, }; diff --git a/src/llama-build-context.cpp b/src/llama-build-context.cpp index d2e0fd975..d0aba9dab 100644 --- a/src/llama-build-context.cpp +++ b/src/llama-build-context.cpp @@ -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; diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index f00951f03..8101dad7a 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -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); diff --git a/src/llama-load-tensors.cpp b/src/llama-load-tensors.cpp index 6d65f084d..2522d4cdb 100644 --- a/src/llama-load-tensors.cpp +++ b/src/llama-load-tensors.cpp @@ -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"); diff --git a/src/llama-model.cpp b/src/llama-model.cpp index c5a0ac039..2be2074f5 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -845,6 +845,29 @@ static const std::map> 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) { diff --git a/src/llama-spec-features.cpp b/src/llama-spec-features.cpp index 5a32b8482..3bdc76f32 100644 --- a/src/llama-spec-features.cpp +++ b/src/llama-spec-features.cpp @@ -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; -} \ No newline at end of file +} diff --git a/src/llama.cpp b/src/llama.cpp index 4aced5326..d4f807dc9 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -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{ (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: