diff --git a/common/speculative.cpp b/common/speculative.cpp index 89e9b2782c..12909822ba 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -2385,6 +2385,9 @@ common_speculative_init_result::common_speculative_init_result( cparams.ctx_type = LLAMA_CONTEXT_TYPE_MTP; } + // the draft context holds as many tokens per sequence as the target context + cparams.n_ctx = llama_n_ctx(ctx_tgt); + // note: for small models maybe we can set this to the maximum possible draft from all speculative types // the extra memory for small models is likely negligible? cparams.n_rs_seq = 0; diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 1293c86402..f8976c6df3 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -1060,6 +1060,18 @@ private: uint32_t hp_nct = 0; uint32_t hp_nex = 0; try { + // the draft context follows the target context, measure it at the largest context the target can take + if (cparams_dft.n_ctx == 0) { + auto mparams_tgt = common_model_params_to_llama(params_base); + auto cparams_tgt = common_context_params_to_llama(params_base); + + common_get_device_memory_data( + params_base.model.path.c_str(), &mparams_tgt, &cparams_tgt, + devs, hp_ngl, hp_nct, hp_nex, GGML_LOG_LEVEL_ERROR); + + cparams_dft.n_ctx = hp_nct * (params_base.kv_unified ? 1 : params_base.n_parallel); + } + auto dmd = common_get_device_memory_data( params_dft.model.path.c_str(), &mparams_dft, &cparams_dft, devs, hp_ngl, hp_nct, hp_nex, GGML_LOG_LEVEL_ERROR);