mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-08-17 04:15:30 +00:00
support gemma assistant as a draft model
This commit is contained in:
parent
a860ec0b37
commit
10e4b6d5e3
1 changed files with 143 additions and 17 deletions
|
|
@ -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<llama_token> 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<llama_token> 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<gpt_vocab::id> 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<parts.size();++p)
|
||||
{
|
||||
std::vector<gpt_vocab::id> 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)
|
||||
{
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue