server: make the draft context follow the target context

With a non-unified KV cache the target context now holds n_ctx_train
tokens per sequence, while the draft context was still created with
n_ctx = 0 and fell back to n_ctx_train / n_streams per sequence. A slot
filled beyond that point makes the draft batch fail to decode, and the
server answers 500 on the request.

The draft context now takes its size from the target context, so both
hold the same number of tokens per sequence. Contexts that share their
cells with the target no longer need the kv_size override.

The memory reserved for the draft model before fitting is measured at
the largest context the target can take, since the draft context grows
with the target and a fixed byte margin cannot express that.
This commit is contained in:
Pascal
2026-08-21 22:55:36 +02:00
parent 2913a33e9a
commit aed7d0ba52
2 changed files with 15 additions and 0 deletions
+3
View File
@@ -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;
+12
View File
@@ -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);