From 10e4b6d5e310b6fdd02e4e0c126eb5a042fdfdcb Mon Sep 17 00:00:00 2001 From: Concedo <39025047+LostRuins@users.noreply.github.com> Date: Thu, 11 Jun 2026 12:16:43 +0800 Subject: [PATCH] support gemma assistant as a draft model --- gpttype_adapter.cpp | 160 +++++++++++++++++++++++++++++++++++++++----- 1 file changed, 143 insertions(+), 17 deletions(-) diff --git a/gpttype_adapter.cpp b/gpttype_adapter.cpp index ab15d7ac1..824823eb5 100644 --- a/gpttype_adapter.cpp +++ b/gpttype_adapter.cpp @@ -49,6 +49,7 @@ #include "mpt_v3.cpp" #include "tools/mtmd/mtmd.h" #include "tools/mtmd/mtmd-helper.h" +#include "common/speculative.h" #include "vendor/stb/stb_image.h" #include "otherarch/sdcpp/thirdparty/stb_image_resize.h" #include "common/common.h" @@ -116,6 +117,8 @@ static llama_v2_context * llama_ctx_v2 = nullptr; static llama_v3_context * llama_ctx_v3 = nullptr; static llama_context * llama_ctx_v4 = nullptr; static llama_context * draft_ctx = nullptr; //will remain null if speculative is unused +static common_speculative * draft_spec = nullptr; // llama.cpp speculative state, used for MTP draft heads +static bool draft_is_mtp = false; static llama_context * guidance_ctx = nullptr; //for classifier free guidance, will be null if unused static mtmd_context * mtmd_ctx = nullptr; //for multimodal media @@ -655,7 +658,7 @@ const char * kcpp_print_system_info(void) { } //loads a model for speculative decoding. -static void speculative_decoding_setup(std::string spec_model_filename, const llama_model_params & base_model_params, const llama_context_params & base_ctx_params, int base_n_vocab, const float * draft_gpusplit, int draft_gpulayers) +static void speculative_decoding_setup(std::string spec_model_filename, llama_context * main_ctx, const llama_model_params & base_model_params, const llama_context_params & base_ctx_params, int base_n_vocab, const float * draft_gpusplit, int draft_gpulayers) { llama_model_params draft_model_params = llama_model_default_params(); llama_context_params draft_ctx_params = llama_context_default_params(); @@ -694,17 +697,33 @@ static void speculative_decoding_setup(std::string spec_model_filename, const ll draft_ctx_params.swa_full = base_ctx_params.swa_full; llama_model * draftmodel = llama_model_load_from_file(spec_model_filename.c_str(), draft_model_params); + if(draftmodel == nullptr) + { + printf("Error: failed to load speculative decoding draft model '%s'\n", spec_model_filename.c_str()); + printf("Speculative Decoding will not be used!\n"); + draft_is_mtp = false; + return; + } + draft_is_mtp = draftmodel && draftmodel->hparams.n_layer_nextn > 0; + if(draft_is_mtp) + { + printf("Detected MTP draft head, using llama.cpp MTP speculative decoding.\n"); + draft_ctx_params.ctx_type = LLAMA_CONTEXT_TYPE_MTP; + draft_ctx_params.ctx_other = main_ctx; + draft_ctx_params.n_rs_seq = speculative_chunk_amt; + } draft_ctx = llama_init_from_model(draftmodel, draft_ctx_params); if(draft_ctx == NULL) { printf("Error: failed to load speculative decoding draft model '%s'\n", spec_model_filename.c_str()); printf("Speculative Decoding will not be used!\n"); + draft_is_mtp = false; } else { const llama_vocab * tmpvocab = llama_model_get_vocab(draftmodel); int draftvocab = llama_vocab_n_tokens(tmpvocab); - if(llama_model_is_recurrent(draftmodel) || llama_model_is_hybrid(draftmodel)) + if(!draft_is_mtp && (llama_model_is_recurrent(draftmodel) || llama_model_is_hybrid(draftmodel))) { printf("Error: Speculative decoding cannot be used with Recurrent draft models!\n"); llama_free(draft_ctx); @@ -728,13 +747,51 @@ static void speculative_decoding_setup(std::string spec_model_filename, const ll printf("If you REALLY want to override this, run in --debugmode and this restriction will be disabled. However, you might encounter unwanted results!\n"); llama_free(draft_ctx); draft_ctx = nullptr; + draft_is_mtp = false; } } } + + if(draft_ctx && draft_is_mtp) + { + common_params_speculative spec_params; + spec_params.types = { COMMON_SPECULATIVE_TYPE_DRAFT_MTP }; + spec_params.draft.ctx_tgt = main_ctx; + spec_params.draft.ctx_dft = draft_ctx; + spec_params.draft.n_max = speculative_chunk_amt; + spec_params.draft.n_min = 0; + spec_params.draft.p_min = 0.0f; + spec_params.draft.backend_sampling = false; + spec_params.draft.n_gpu_layers = draft_gpulayers; + spec_params.draft.cache_type_k = draft_ctx_params.type_k; + spec_params.draft.cache_type_v = draft_ctx_params.type_v; + draft_spec = common_speculative_init(spec_params, 1); + if(draft_spec == nullptr) + { + printf("Error: failed to initialize MTP speculative decoding state.\n"); + llama_free(draft_ctx); + draft_ctx = nullptr; + draft_is_mtp = false; + } + } } } +static int32_t kcpp_decode_main_and_mtp_spec(llama_context * main_ctx, llama_batch batch) +{ + const int32_t decode_status = llama_decode(main_ctx, batch); + if(decode_status == 0 && draft_spec && draft_is_mtp) + { + if(!common_speculative_process(draft_spec, batch)) + { + printf("\nERROR: MTP speculative state update failed!\n"); + return -1; + } + } + return decode_status; +} + static speculative_draft_result speculative_decoding_eval_chunk(llama_context * draft_ctx, llama_context * main_ctx, const llama_tokens & embd, const int n_vocab, const int & n_past) { speculative_draft_result results; @@ -796,6 +853,59 @@ static speculative_draft_result speculative_decoding_eval_chunk(llama_context * return results; } +static speculative_draft_result speculative_decoding_eval_mtp_chunk(llama_context * main_ctx, const llama_tokens & embd, const int & n_past) +{ + speculative_draft_result results; + results.draft_success = false; + if(embd.size()!=1 || draft_spec==nullptr) + { + printf("\nERROR: MTP speculative decoding applied to invalid batch!\n"); + return results; + } + + std::vector drafted_ids; + llama_tokens prompt_tokens(current_context_tokens.begin(), current_context_tokens.end()); + auto & dp = common_speculative_get_draft_params(draft_spec, 0); + dp.drafting = true; + dp.n_max = speculative_chunk_amt; + dp.n_past = n_past; + dp.id_last = embd[0]; + dp.prompt = &prompt_tokens; + dp.result = &drafted_ids; + + common_speculative_draft(draft_spec); + if(drafted_ids.empty()) + { + printf("\nERROR: MTP draft model produced no draft tokens!\n"); + return results; + } + + std::vector real_embd; + real_embd.reserve(drafted_ids.size()); + real_embd.push_back(embd[0]); + for(size_t i = 0; i + 1 < drafted_ids.size(); ++i) + { + real_embd.push_back(drafted_ids[i]); + } + + kcpp_embd_batch batch = kcpp_embd_batch(real_embd, n_past, use_mrope, true); + const int32_t decode_status = kcpp_decode_main_and_mtp_spec(main_ctx, batch.batch); + if(decode_status != 0) + { + printf("\nERROR: MTP speculative verification failed! (code:%d)\n", decode_status); + return results; + } + + results.drafted_amount = drafted_ids.size(); + for(size_t i = 0; i < drafted_ids.size(); ++i) + { + results.draftids.push_back(drafted_ids[i]); + results.actual_logits.push_back(llama_get_logits_ith(main_ctx, (int32_t)i)); + } + results.draft_success = true; + return results; +} + // KCPP SAMPLING FUNCTIONS void sample_softmax(llama_token_data_array * cur_p, bool do_sort=true) { if(!(cur_p->size > 0)) @@ -2341,7 +2451,13 @@ ModelLoadResult gpttype_load_model(const load_model_inputs inputs, FileFormat in kcpp_data->swa_full = inputs.prevent_swa; debugmode = inputs.debugmode; + if(draft_spec) + { + common_speculative_free(draft_spec); + draft_spec = nullptr; + } draft_ctx = nullptr; + draft_is_mtp = false; guidance_ctx = nullptr; if(mtmd_ctx) { @@ -2961,11 +3077,7 @@ ModelLoadResult gpttype_load_model(const load_model_inputs inputs, FileFormat in if(draftmodel_filename !="" && file_format==FileFormat::GGUF_GENERIC) { - if(llama_model_is_recurrent(llamamodel) || llama_model_is_hybrid(llamamodel)) - { - printf("Error: Speculative decoding cannot be used with Recurrent models!\n"); - } - else if(mtmd_ctx!=nullptr) + if(mtmd_ctx!=nullptr) { printf("Error: Speculative decoding cannot be used with multimodal projectors!\n"); } @@ -2973,7 +3085,7 @@ ModelLoadResult gpttype_load_model(const load_model_inputs inputs, FileFormat in { printf("\nAttempting to load draft model for speculative decoding. It will be fully offloaded if possible. Vocab must match the main model.\n"); speculative_chunk_amt = inputs.draft_amount; - speculative_decoding_setup(draftmodel_filename, model_params, llama_ctx_params, n_vocab, inputs.draft_gpusplit, inputs.draft_gpulayers); + speculative_decoding_setup(draftmodel_filename, llama_ctx_v4, model_params, llama_ctx_params, n_vocab, inputs.draft_gpusplit, inputs.draft_gpulayers); } } @@ -5549,7 +5661,7 @@ generation_outputs gpttype_generate(const generation_inputs inputs) if(embd.size()!=1 || draft_ctx==nullptr || remaining_tokens<=speculative_chunk_amt || grammar!=nullptr || startedsampling==false) //for large batch, or if no draft model, PP/TG as usual { draft_used = false; - kcpp_embd_batch batch = kcpp_embd_batch(embd, n_past, use_mrope, false); + kcpp_embd_batch batch = kcpp_embd_batch(embd, n_past, use_mrope, draft_is_mtp); int32_t decode_status = -1; bool skipdecodelater = false; @@ -5580,8 +5692,8 @@ generation_outputs gpttype_generate(const generation_inputs inputs) rnn_snapshot_taken = true; } std::vector chunk = parts[p]; - kcpp_embd_batch smallbatch = kcpp_embd_batch(chunk, temp_past, use_mrope, false); - decode_status = llama_decode(llama_ctx_v4, smallbatch.batch); + kcpp_embd_batch smallbatch = kcpp_embd_batch(chunk, temp_past, use_mrope, draft_is_mtp); + decode_status = kcpp_decode_main_and_mtp_spec(llama_ctx_v4, smallbatch.batch); if(p==0 && decode_status==1) { skipdecodelater = false; @@ -5596,7 +5708,7 @@ generation_outputs gpttype_generate(const generation_inputs inputs) if(!skipdecodelater) { - decode_status = llama_decode(llama_ctx_v4, batch.batch); + decode_status = kcpp_decode_main_and_mtp_spec(llama_ctx_v4, batch.batch); if(decode_status==1 && embd.size()>128) { printf("Couldn't find a big KV slot. Retry with smaller batch size of 128...\n"); @@ -5606,8 +5718,8 @@ generation_outputs gpttype_generate(const generation_inputs inputs) for(int p=0;p chunk = parts[p]; - kcpp_embd_batch smallbatch = kcpp_embd_batch(chunk, temp_past, use_mrope, false); - int32_t decode_status2 = llama_decode(llama_ctx_v4, smallbatch.batch); + kcpp_embd_batch smallbatch = kcpp_embd_batch(chunk, temp_past, use_mrope, draft_is_mtp); + int32_t decode_status2 = kcpp_decode_main_and_mtp_spec(llama_ctx_v4, smallbatch.batch); if(debugmode==1 && !is_quiet) { printf("Retry chunk: %zu at %d... status: %s\n",chunk.size(),temp_past,(decode_status2==0?"ok":"fail")); @@ -5622,13 +5734,20 @@ generation_outputs gpttype_generate(const generation_inputs inputs) } } - if(draft_ctx) + if(draft_ctx && !draft_is_mtp) { evalres = (evalres && (llama_decode(draft_ctx, batch.batch)==0)); } } else { //individual tokens AND speculative is used (generation) draft_used = true; - draft_results = speculative_decoding_eval_chunk(draft_ctx, llama_ctx_v4, embd, n_vocab, n_past); + if(draft_is_mtp) + { + draft_results = speculative_decoding_eval_mtp_chunk(llama_ctx_v4, embd, n_past); + } + else + { + draft_results = speculative_decoding_eval_chunk(draft_ctx, llama_ctx_v4, embd, n_vocab, n_past); + } evalres = draft_results.draft_success; if(debugmode==1 && !is_quiet) { @@ -5782,6 +5901,7 @@ generation_outputs gpttype_generate(const generation_inputs inputs) int logits_to_sample = 1; int logits_sampled = 0; bool abort_draft = false; + int draft_accepted_this_round = 0; if(draft_used) { logits_to_sample = draft_results.drafted_amount; @@ -5803,7 +5923,7 @@ generation_outputs gpttype_generate(const generation_inputs inputs) } else { - logitsPtr = llama_get_logits(llama_ctx_v4); + logitsPtr = draft_is_mtp ? llama_get_logits_ith(llama_ctx_v4, -1) : llama_get_logits(llama_ctx_v4); } } else if(file_format == FileFormat::GGJT_3) @@ -5923,6 +6043,7 @@ generation_outputs gpttype_generate(const generation_inputs inputs) abort_draft = true; } else { draft_successes += 1; + draft_accepted_this_round += 1; } } @@ -6105,6 +6226,11 @@ generation_outputs gpttype_generate(const generation_inputs inputs) logits_sampled += 1; } + if(draft_used && draft_is_mtp && draft_spec) + { + common_speculative_accept(draft_spec, 0, draft_accepted_this_round); + } + //if we have somehow skipped ahead (e.g drafting), ensure that all tokens after npast are purged if (file_format == FileFormat::GGUF_GENERIC && draft_used) {