dflash: use draft context as capacity contract (#2341)

dflash: account selector drafts in telemetry
This commit is contained in:
Samuel Oliveira Alves 2026-08-26 13:26:50 -03:00 committed by GitHub
parent 850320be8b
commit 2f068b5d87
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 122 additions and 79 deletions

View File

@ -168,9 +168,6 @@ common_params_speculative common_params_speculative::with_stage_overrides(const
if (stage.has_mtp_heads_override()) { if (stage.has_mtp_heads_override()) {
result.mtp_heads = stage.mtp_heads; result.mtp_heads = stage.mtp_heads;
} }
if (stage.has_dflash_cross_ctx_override()) {
result.dflash_cross_ctx = stage.dflash_cross_ctx;
}
if (stage.has_ngram_size_n_override()) { if (stage.has_ngram_size_n_override()) {
result.ngram_size_n = stage.ngram_size_n; result.ngram_size_n = stage.ngram_size_n;
result.ngram_mod.reset(); result.ngram_mod.reset();
@ -321,9 +318,6 @@ bool common_speculative_validate_chain(const common_params_speculative & params,
return fail(common_speculative_type_to_str(stage.type) + " speculative stage requires a draft model or draft params"); return fail(common_speculative_type_to_str(stage.type) + " speculative stage requires a draft model or draft params");
} }
if (common_speculative_type_is_dflash_family(stage.type) && stage_params.dflash_cross_ctx < 1) {
return fail(common_speculative_type_to_str(stage.type) + " speculative stage requires cross_ctx >= 1");
}
} }
if (resolved.size() == 2) { if (resolved.size() == 2) {
@ -959,13 +953,6 @@ static void common_speculative_stage_apply_kv(
} }
return; return;
} }
if (key == "cross_ctx" || key == "dflash_cross_ctx") {
stage.dflash_cross_ctx = std::stoi(value_raw);
if (stage.dflash_cross_ctx < 1) {
throw std::invalid_argument("speculative stage dflash cross_ctx must be at least 1");
}
return;
}
if (key == "ngram_size_n") { if (key == "ngram_size_n") {
stage.ngram_size_n = std::stoi(value_raw); stage.ngram_size_n = std::stoi(value_raw);
if (stage.ngram_size_n < 1 || stage.ngram_size_n > 1024) { if (stage.ngram_size_n < 1 || stage.ngram_size_n > 1024) {
@ -3071,7 +3058,7 @@ void gpt_params_print_usage(int /*argc*/, char ** argv, const gpt_params & param
"path to dynamic lookup cache to use for lookup decoding (updated by generation)" }); "path to dynamic lookup cache to use for lookup decoding (updated by generation)" });
options.push_back({ "*", "-c, --ctx-size N", "size of the prompt context (default: %d, 0 = loaded from model)", params.n_ctx }); options.push_back({ "*", "-c, --ctx-size N", "size of the prompt context (default: %d, 0 = loaded from model)", params.n_ctx });
options.push_back({ "*", "-cd, --ctx-size-draft N", "size of the prompt context for the draft model (default: %d, 0 = loaded from model)", params.speculative.n_ctx }); options.push_back({ "*", "-cd, --ctx-size-draft N", "size of the prompt context for the draft model (default: %d, 0 = inherits target context for DFlash/DSpark, otherwise loaded from model)", params.speculative.n_ctx });
options.push_back({ "*", "-ctx-ckpt N, --ctx-checkpoints N", "max number of context checkpoints to create per slot (default: %d)",params.ctx_checkpoints_n}); options.push_back({ "*", "-ctx-ckpt N, --ctx-checkpoints N", "max number of context checkpoints to create per slot (default: %d)",params.ctx_checkpoints_n});
options.push_back({ "*", "-ctx-ckpt-i N, --ctx-checkpoints-interval N", "minimum number of tokens between each context checkpoint. (default: %d, <=0 disable)",params.ctx_checkpoints_interval}); options.push_back({ "*", "-ctx-ckpt-i N, --ctx-checkpoints-interval N", "minimum number of tokens between each context checkpoint. (default: %d, <=0 disable)",params.ctx_checkpoints_interval});
@ -3389,12 +3376,12 @@ void gpt_params_print_usage(int /*argc*/, char ** argv, const gpt_params & param
" --recurrent-ckpt-mode remains as a deprecated alias" }); " --recurrent-ckpt-mode remains as a deprecated alias" });
options.push_back({ "*", "--spec-type SPEC[:k=v,...]", "canonical speculative stage entry; repeat for a supported two-stage chain.\n" options.push_back({ "*", "--spec-type SPEC[:k=v,...]", "canonical speculative stage entry; repeat for a supported two-stage chain.\n"
"types: none, draft, dflash, dspark, mtp, ngram-cache, ngram-simple, ngram-map-k, ngram-map-k4v, ngram-mod, suffix\n" "types: none, draft, dflash, dspark, mtp, ngram-cache, ngram-simple, ngram-map-k, ngram-map-k4v, ngram-mod, suffix\n"
"canonical keys: n_max,n_min,p_min,heads,cross_ctx,ngram_size_n,ngram_size_m,ngram_min_hits,suffix_min_match_len,suffix_max_depth,suffix_corpus\n" "canonical keys: n_max,n_min,p_min,heads,ngram_size_n,ngram_size_m,ngram_min_hits,suffix_min_match_len,suffix_max_depth,suffix_corpus\n"
"MTP heads: heads=1 is the default; heads>1 and heads=0 (all model heads) are experimental\n" "MTP heads: heads=1 is the default; heads>1 and heads=0 (all model heads) are experimental\n"
"for comma-bearing string values, quote the value inside the stage payload for normal shell use\n" "for comma-bearing string values, quote the value inside the stage payload for normal shell use\n"
"if argv is passed directly without shell unescaping, the parser also accepts escaped commas as \\,\n" "if argv is passed directly without shell unescaping, the parser also accepts escaped commas as \\,\n"
"examples: --spec-type mtp:n_max=1,p_min=0.0\n" "examples: --spec-type mtp:n_max=1,p_min=0.0\n"
" --model-draft draft.gguf --spec-type dflash:n_max=4,cross_ctx=512\n" " --model-draft draft.gguf --spec-type dflash:n_max=4\n"
" --spec-type ngram-mod:n_max=64,n_min=2,ngram_size_n=8 --spec-type mtp:n_max=1,p_min=0.0\n" " --spec-type ngram-mod:n_max=64,n_min=2,ngram_size_n=8 --spec-type mtp:n_max=1,p_min=0.0\n"
" --spec-type \"suffix:n_max=16,n_min=2,suffix_min_match_len=5,suffix_max_depth=64,suffix_corpus='/tmp/spec,type-corpus.json'\"\n" " --spec-type \"suffix:n_max=16,n_min=2,suffix_min_match_len=5,suffix_max_depth=64,suffix_corpus='/tmp/spec,type-corpus.json'\"\n"
"legacy --spec-stage, --draft-*, --spec-ngram-*, --suffix-* and -mtp flags are rejected" }); "legacy --spec-stage, --draft-*, --spec-ngram-*, --suffix-* and -mtp flags are rejected" });

View File

@ -176,7 +176,6 @@ struct common_speculative_stage_params {
int32_t n_min = -1; int32_t n_min = -1;
float p_min = -1.0f; float p_min = -1.0f;
int32_t mtp_heads = -1; int32_t mtp_heads = -1;
int32_t dflash_cross_ctx = -1;
uint16_t ngram_size_n = 0; uint16_t ngram_size_n = 0;
uint16_t ngram_size_m = 0; uint16_t ngram_size_m = 0;
@ -190,7 +189,6 @@ struct common_speculative_stage_params {
bool has_n_min_override() const { return n_min >= 0; } bool has_n_min_override() const { return n_min >= 0; }
bool has_p_min_override() const { return p_min >= 0.0f; } bool has_p_min_override() const { return p_min >= 0.0f; }
bool has_mtp_heads_override() const { return mtp_heads >= 0; } bool has_mtp_heads_override() const { return mtp_heads >= 0; }
bool has_dflash_cross_ctx_override() const { return dflash_cross_ctx >= 0; }
bool has_ngram_size_n_override() const { return ngram_size_n > 0; } bool has_ngram_size_n_override() const { return ngram_size_n > 0; }
bool has_ngram_size_m_override() const { return ngram_size_m > 0; } bool has_ngram_size_m_override() const { return ngram_size_m > 0; }
bool has_ngram_min_hits_override() const { return ngram_min_hits > 0; } bool has_ngram_min_hits_override() const { return ngram_min_hits > 0; }
@ -224,7 +222,6 @@ struct common_params_speculative {
int32_t n_min = 0; // minimum number of tokens to draft during speculative decoding int32_t n_min = 0; // minimum number of tokens to draft during speculative decoding
std::vector<common_speculative_stage_params> stages; // explicit stage chain for single-spec or self-spec + model fallback std::vector<common_speculative_stage_params> stages; // explicit stage chain for single-spec or self-spec + model fallback
int32_t mtp_heads = 1; // MTP heads to use; 1 is the default, while >1 and 0 (all model heads) are experimental int32_t mtp_heads = 1; // MTP heads to use; 1 is the default, while >1 and 0 (all model heads) are experimental
int32_t dflash_cross_ctx = 512; // target-feature context window for DFlash
// Samplers for DFlash2 // Samplers for DFlash2
float draft_temperature = 0.0f; float draft_temperature = 0.0f;

View File

@ -9,6 +9,8 @@
#include <random> #include <random>
#include <vector> #include <vector>
#include "ggml.h"
static bool common_speculative_are_dflash_compatible( static bool common_speculative_are_dflash_compatible(
const llama_model * model_tgt, const llama_model * model_tgt,
const llama_model * model_dft) { const llama_model * model_dft) {
@ -80,6 +82,8 @@ static bool common_speculative_are_dflash_compatible(
struct common_speculative_state_dflash; struct common_speculative_state_dflash;
static void dflash_materialize_target_window_features(common_speculative_state_dflash & state); static void dflash_materialize_target_window_features(common_speculative_state_dflash & state);
static void dflash_record_window_update(common_speculative_state_dflash & state,
int32_t keep_rows, int32_t append_rows, bool replace);
// DFlash runtime state and draft path. // DFlash runtime state and draft path.
struct common_speculative_state_dflash : public common_speculative_state { struct common_speculative_state_dflash : public common_speculative_state {
@ -108,7 +112,6 @@ struct common_speculative_state_dflash : public common_speculative_state {
std::vector<int32_t> target_layer_ids; std::vector<int32_t> target_layer_ids;
std::vector<float> target_window; std::vector<float> target_window;
std::vector<llama_pos> target_window_pos; std::vector<llama_pos> target_window_pos;
std::vector<float> target_window_stage;
std::vector<llama_pos> target_window_pos_stage; std::vector<llama_pos> target_window_pos_stage;
std::vector<float> target_window_ring; std::vector<float> target_window_ring;
std::vector<float> target_window_append_features; std::vector<float> target_window_append_features;
@ -121,6 +124,10 @@ struct common_speculative_state_dflash : public common_speculative_state {
bool target_window_replace = false; bool target_window_replace = false;
bool target_window_materialized = false; bool target_window_materialized = false;
llama_pos last_target_pos = -1; llama_pos last_target_pos = -1;
uint64_t rebuild_count = 0;
uint64_t rebuild_rows = 0;
uint64_t rebuild_decode_time_us = 0;
uint64_t generated_tokens = 0;
common_speculative_state_dflash( common_speculative_state_dflash(
enum common_speculative_type type, enum common_speculative_type type,
@ -251,10 +258,7 @@ struct common_speculative_state_dflash : public common_speculative_state {
} }
batch = llama_batch_init(std::max(1, query_capacity), 0, 1); batch = llama_batch_init(std::max(1, query_capacity), 0, 1);
target_window.reserve((size_t) this->cross_ctx * (size_t) n_target_features);
target_window_stage.reserve((size_t) this->cross_ctx * (size_t) n_target_features);
target_window_ring.resize((size_t) this->cross_ctx * (size_t) n_target_features); target_window_ring.resize((size_t) this->cross_ctx * (size_t) n_target_features);
target_window_append_features.reserve((size_t) this->cross_ctx * (size_t) n_target_features);
target_window_pos.reserve((size_t) this->cross_ctx); target_window_pos.reserve((size_t) this->cross_ctx);
target_window_pos_stage.reserve((size_t) this->cross_ctx); target_window_pos_stage.reserve((size_t) this->cross_ctx);
ready = true; ready = true;
@ -262,11 +266,18 @@ struct common_speculative_state_dflash : public common_speculative_state {
llama_set_dflash_visible_cross_ctx(ctx_dft, this->cross_ctx); llama_set_dflash_visible_cross_ctx(ctx_dft, this->cross_ctx);
llama_set_dflash_dspark(ctx_dft, is_dspark); llama_set_dflash_dspark(ctx_dft, is_dspark);
LOG_INF("%s: DFlash context ready (n_ctx=%d, block_size=%d, query_capacity=%d, active_width=%d, cross_ctx=%d, n_target_features=%d, n_target_layers=%d)\n", LOG_INF("%s: DFlash context ready (n_ctx=%d, block_size=%d, query_capacity=%d, active_width=%d, cross_ctx=%d, n_target_features=%d, n_target_layers=%d)\n",
__func__, llama_n_ctx(ctx_dft), block_size, query_capacity, active_width, __func__, llama_n_ctx(ctx_dft), block_size, query_capacity, active_width, this->cross_ctx,
this->cross_ctx, n_target_features, n_target_layers); n_target_features, n_target_layers);
} }
~common_speculative_state_dflash() override { ~common_speculative_state_dflash() override {
if (rebuild_count > 0) {
LOG_INF("%s: DFlash cache rebuilds=%llu rows=%llu rebuild+decode=%.3f ms rows/token=%.3f\n",
__func__,
(unsigned long long) rebuild_count,
(unsigned long long) rebuild_rows,
(double) rebuild_decode_time_us / 1000.0,
generated_tokens > 0 ? (double) rebuild_rows / (double) generated_tokens : 0.0);
}
llama_clear_dflash_capture(ctx_tgt); llama_clear_dflash_capture(ctx_tgt);
if (ctx_dft) { if (ctx_dft) {
llama_free(ctx_dft); llama_free(ctx_dft);
@ -282,6 +293,10 @@ struct common_speculative_state_dflash : public common_speculative_state {
llama_reset_dflash_kv_cache_state(ctx_dft); llama_reset_dflash_kv_cache_state(ctx_dft);
proposal_dists.clear(); proposal_dists.clear();
selector_rng_initialized = false; selector_rng_initialized = false;
rebuild_count = 0;
rebuild_rows = 0;
rebuild_decode_time_us = 0;
generated_tokens = 0;
} }
void draft( void draft(
@ -312,6 +327,7 @@ struct common_speculative_state_dflash : public common_speculative_state {
const float * target_features = nullptr; const float * target_features = nullptr;
size_t target_feature_floats = 0; size_t target_feature_floats = 0;
const int32_t batch_len = is_dspark ? n_keep : n_keep + 1;
llama_dflash_window_update window_update = { llama_dflash_window_update window_update = {
target_window_version, target_window_version,
target_window_keep_rows, target_window_keep_rows,
@ -323,6 +339,7 @@ struct common_speculative_state_dflash : public common_speculative_state {
const llama_dflash_kv_cache_transition cache_plan = const llama_dflash_kv_cache_transition cache_plan =
llama_plan_dflash_kv_cache_transition_for_ctx(ctx_dft, window_update, target_window_rows); llama_plan_dflash_kv_cache_transition_for_ctx(ctx_dft, window_update, target_window_rows);
const int64_t rebuild_start_us = cache_plan.rebuild_cache ? ggml_time_us() : 0;
if (cache_plan.rebuild_cache) { if (cache_plan.rebuild_cache) {
dflash_materialize_target_window_features(*this); dflash_materialize_target_window_features(*this);
target_features = target_window.data(); target_features = target_window.data();
@ -339,20 +356,31 @@ struct common_speculative_state_dflash : public common_speculative_state {
llama_kv_cache_clear(ctx_dft); llama_kv_cache_clear(ctx_dft);
batch.n_tokens = 0; batch.n_tokens = 0;
const int32_t batch_len = is_dspark ? n_keep : n_keep + 1;
const bool output_seed_row = is_dspark; const bool output_seed_row = is_dspark;
const bool output_mask_rows = !is_dflash2; const bool output_mask_rows = !is_dflash2;
// id_last's true position is one past the newest committed feature row // id_last's true position is one past the newest committed feature row
// (last_target_pos): seed there, masks follow. Mirrors mainline's // (last_target_pos): seed there, masks follow. Mirrors mainline's
// [id_last @ n_past, mask @ n_past+1, ...] block geometry. // [id_last @ n_past, mask @ n_past+1, ...] block geometry.
const llama_pos draft_pos_base = last_target_pos >= 0 ? last_target_pos + 1 : (llama_pos) target_window_rows; const llama_pos draft_pos_base = last_target_pos >= 0
? last_target_pos + 1
: (llama_pos) target_window_rows;
common_batch_add(batch, id_last, draft_pos_base, { 0 }, output_seed_row); common_batch_add(batch, id_last, draft_pos_base, { 0 }, output_seed_row);
for (int32_t i = 1; i < batch_len; ++i) { for (int32_t i = 1; i < batch_len; ++i) {
common_batch_add(batch, mask_token_id, draft_pos_base + i, { 0 }, output_mask_rows); common_batch_add(batch, mask_token_id, draft_pos_base + i, { 0 }, output_mask_rows);
} }
if (llama_decode(ctx_dft, batch) != 0) { const int decode_status = llama_decode(ctx_dft, batch);
LOG_ERR("%s: llama_decode() failed for DFlash draft batch\n", __func__); if (cache_plan.rebuild_cache) {
rebuild_count++;
rebuild_rows += (uint64_t) target_window_rows;
rebuild_decode_time_us += (uint64_t) std::max<int64_t>(0, ggml_time_us() - rebuild_start_us);
std::vector<float>().swap(target_window);
target_window_materialized = false;
}
if (decode_status != 0) {
LOG_ERR("%s: DFlash draft decode failed (status=%d local_pos=%d history_rows=%d capacity=%d)\n",
__func__, decode_status, (int) draft_pos_base,
target_window_rows, cross_ctx + block_size);
batch.n_tokens = 0; batch.n_tokens = 0;
return; return;
} }
@ -417,6 +445,7 @@ struct common_speculative_state_dflash : public common_speculative_state {
} }
} }
generated_tokens += (uint64_t) result.size();
batch.n_tokens = 0; batch.n_tokens = 0;
return; return;
} }
@ -431,6 +460,8 @@ struct common_speculative_state_dflash : public common_speculative_state {
result.push_back(id); result.push_back(id);
} }
generated_tokens += (uint64_t) result.size();
batch.n_tokens = 0; batch.n_tokens = 0;
} }
@ -509,7 +540,8 @@ static void dflash_materialize_target_window_features(common_speculative_state_d
} }
const size_t row_width = (size_t) state.n_target_features; const size_t row_width = (size_t) state.n_target_features;
state.target_window.resize((size_t) state.target_window_rows * row_width); const size_t window_elements = (size_t) state.target_window_rows * row_width;
state.target_window.resize(window_elements);
const int32_t read_start = (state.target_window_ring_write_pos - state.target_window_rows + state.cross_ctx) % state.cross_ctx; const int32_t read_start = (state.target_window_ring_write_pos - state.target_window_rows + state.cross_ctx) % state.cross_ctx;
const int32_t first_rows = std::min<int32_t>(state.target_window_rows, state.cross_ctx - read_start); const int32_t first_rows = std::min<int32_t>(state.target_window_rows, state.cross_ctx - read_start);
@ -569,10 +601,11 @@ static bool dflash_append_target_features(
if (n_rows >= state.cross_ctx) { if (n_rows >= state.cross_ctx) {
const int32_t keep_from = n_rows - state.cross_ctx; const int32_t keep_from = n_rows - state.cross_ctx;
state.target_window_pos.assign(new_positions.begin() + keep_from, new_positions.end()); state.target_window_pos.assign(new_positions.begin() + keep_from, new_positions.end());
state.target_window_append_features.assign( std::vector<float>().swap(state.target_window_append_features);
new_rows.begin() + (ptrdiff_t) keep_from * (ptrdiff_t) row_width, dflash_ring_reset_rows(
new_rows.end()); state,
dflash_ring_reset_rows(state, state.target_window_append_features.data(), state.cross_ctx); new_rows.data() + (size_t) keep_from * row_width,
state.cross_ctx);
state.target_window_rows = state.cross_ctx; state.target_window_rows = state.cross_ctx;
state.target_window_ring_filled = state.target_window_rows; state.target_window_ring_filled = state.target_window_rows;
@ -582,7 +615,8 @@ static bool dflash_append_target_features(
} }
// In case we are re-decoding we should replace the suffix // In case we are re-decoding we should replace the suffix
const auto overlap = std::lower_bound(state.target_window_pos.begin(), state.target_window_pos.end(), new_positions.front()); const auto overlap = std::lower_bound(
state.target_window_pos.begin(), state.target_window_pos.end(), new_positions.front());
if (overlap != state.target_window_pos.end()) { if (overlap != state.target_window_pos.end()) {
dflash_materialize_target_window_features(state); dflash_materialize_target_window_features(state);
@ -592,19 +626,19 @@ static bool dflash_append_target_features(
const int32_t total_rows = keep_old_rows + n_rows; const int32_t total_rows = keep_old_rows + n_rows;
const size_t total_floats = (size_t) total_rows * row_width; const size_t total_floats = (size_t) total_rows * row_width;
state.target_window_stage.resize(total_floats); state.target_window.resize(total_floats);
std::vector<float> next_window(total_floats);
if (keep_old_rows > 0) { if (keep_old_rows > 0) {
std::copy( std::copy(
state.target_window.begin() + (size_t) old_start * row_width, state.target_window.begin() + (size_t) old_start * row_width,
state.target_window.begin() + (size_t) prefix_rows * row_width, state.target_window.begin() + (size_t) prefix_rows * row_width,
state.target_window_stage.begin()); next_window.begin());
} }
std::copy( std::copy(
new_rows.begin(), new_rows.end(), new_rows.begin(), new_rows.end(),
state.target_window_stage.begin() + (size_t) keep_old_rows * row_width); next_window.begin() + (size_t) keep_old_rows * row_width);
std::vector<llama_pos> & next_window_pos = state.target_window_pos_stage; std::vector<llama_pos> next_window_pos((size_t) total_rows);
next_window_pos.resize((size_t) total_rows);
if (keep_old_rows > 0) { if (keep_old_rows > 0) {
std::copy( std::copy(
state.target_window_pos.begin() + old_start, state.target_window_pos.begin() + old_start,
@ -613,10 +647,8 @@ static bool dflash_append_target_features(
} }
std::copy(new_positions.begin(), new_positions.end(), next_window_pos.begin() + keep_old_rows); std::copy(new_positions.begin(), new_positions.end(), next_window_pos.begin() + keep_old_rows);
state.target_window.swap(state.target_window_stage); state.target_window.swap(next_window);
state.target_window_pos.swap(next_window_pos); state.target_window_pos.swap(next_window_pos);
state.target_window_stage.clear();
next_window_pos.clear();
state.target_window_rows = total_rows; state.target_window_rows = total_rows;
state.target_window_ring_filled = total_rows; state.target_window_ring_filled = total_rows;
dflash_ring_reset_rows(state, state.target_window.data(), total_rows); dflash_ring_reset_rows(state, state.target_window.data(), total_rows);
@ -649,7 +681,6 @@ static bool dflash_append_target_features(
static void dflash_clear_target_features(common_speculative_state_dflash & state) { static void dflash_clear_target_features(common_speculative_state_dflash & state) {
state.target_window.clear(); state.target_window.clear();
state.target_window_pos.clear(); state.target_window_pos.clear();
state.target_window_stage.clear();
state.target_window_pos_stage.clear(); state.target_window_pos_stage.clear();
state.target_window_append_features.clear(); state.target_window_append_features.clear();
state.target_window_rows = 0; state.target_window_rows = 0;
@ -660,6 +691,10 @@ static void dflash_clear_target_features(common_speculative_state_dflash & state
state.target_window_replace = false; state.target_window_replace = false;
state.target_window_materialized = false; state.target_window_materialized = false;
state.last_target_pos = -1; state.last_target_pos = -1;
state.rebuild_count = 0;
state.rebuild_rows = 0;
state.rebuild_decode_time_us = 0;
state.generated_tokens = 0;
llama_reset_dflash_kv_cache_state(state.ctx_dft); llama_reset_dflash_kv_cache_state(state.ctx_dft);
} }
@ -678,11 +713,7 @@ static void dflash_context_shift(
const llama_pos discard_begin = kv_keep; const llama_pos discard_begin = kv_keep;
const llama_pos discard_end = kv_keep + kv_discard; const llama_pos discard_end = kv_keep + kv_discard;
std::vector<float> shifted_rows; int32_t write_row = 0;
std::vector<llama_pos> shifted_positions;
shifted_rows.reserve(state.target_window.size());
shifted_positions.reserve(state.target_window_pos.size());
for (int32_t row = 0; row < state.target_window_rows; ++row) { for (int32_t row = 0; row < state.target_window_rows; ++row) {
llama_pos pos = state.target_window_pos[(size_t) row]; llama_pos pos = state.target_window_pos[(size_t) row];
if (pos >= discard_begin && pos < discard_end) { if (pos >= discard_begin && pos < discard_end) {
@ -694,15 +725,21 @@ static void dflash_context_shift(
} }
const float * row_src = state.target_window.data() + (size_t) row * row_width; const float * row_src = state.target_window.data() + (size_t) row * row_width;
shifted_rows.insert(shifted_rows.end(), row_src, row_src + row_width); if (write_row != row) {
shifted_positions.push_back(pos); std::memmove(
state.target_window.data() + (size_t) write_row * row_width,
row_src,
row_width * sizeof(float));
}
state.target_window_pos[(size_t) write_row++] = pos;
} }
state.target_window = std::move(shifted_rows); state.target_window_rows = write_row;
state.target_window_pos = std::move(shifted_positions); state.target_window_pos.resize((size_t) write_row);
state.target_window_rows = (int32_t) state.target_window_pos.size();
dflash_ring_reset_rows(state, state.target_window.data(), state.target_window_rows); dflash_ring_reset_rows(state, state.target_window.data(), state.target_window_rows);
state.last_target_pos = state.target_window_pos.empty() ? -1 : state.target_window_pos.back(); state.last_target_pos = state.target_window_pos.empty() ? -1 : state.target_window_pos.back();
dflash_record_window_update(state, 0, state.target_window_rows, true); dflash_record_window_update(state, 0, state.target_window_rows, true);
std::vector<float>().swap(state.target_window);
state.target_window_materialized = false;
llama_reset_dflash_kv_cache_state(state.ctx_dft); llama_reset_dflash_kv_cache_state(state.ctx_dft);
} }

View File

@ -1319,6 +1319,8 @@ common_speculative * common_speculative_init(
}); });
llama_context * ctx_dft = nullptr; llama_context * ctx_dft = nullptr;
int32_t dflash_cross_ctx = 0;
if (needs_draft_ctx) { if (needs_draft_ctx) {
if (!params.model_dft) { if (!params.model_dft) {
LOG_ERR("%s: draft speculative stage requires a loaded draft model\n", __func__); LOG_ERR("%s: draft speculative stage requires a loaded draft model\n", __func__);
@ -1333,15 +1335,6 @@ common_speculative * common_speculative_init(
return nullptr; return nullptr;
} }
int32_t max_cross_ctx = 0;
for (const auto & stage : stages) {
if (!common_speculative_type_is_dflash_family(stage.type)) {
continue;
}
max_cross_ctx = std::max(max_cross_ctx, params.with_stage_overrides(stage).dflash_cross_ctx);
}
const int32_t block_size = llama_model_dflash_block_size(params.model_dft); const int32_t block_size = llama_model_dflash_block_size(params.model_dft);
if (block_size <= 0) { if (block_size <= 0) {
LOG_ERR("%s: invalid DFlash draft block size\n", __func__); LOG_ERR("%s: invalid DFlash draft block size\n", __func__);
@ -1349,14 +1342,25 @@ common_speculative * common_speculative_init(
} }
const int32_t query_capacity = std::max(block_size, dsv4_dspark_query_capacity); const int32_t query_capacity = std::max(block_size, dsv4_dspark_query_capacity);
const int64_t required_n_ctx = (int64_t) max_cross_ctx + (int64_t) query_capacity; const int32_t target_ctx = std::max<int32_t>(1, llama_n_ctx(ctx_tgt));
if (required_n_ctx > std::numeric_limits<int32_t>::max()) { const int32_t requested_draft_ctx = cparams_dft.n_ctx > 0
LOG_ERR("%s: invalid DFlash draft context size cross_ctx=%d query_capacity=%d required_n_ctx=%lld\n", ? (int32_t) cparams_dft.n_ctx : target_ctx;
__func__, max_cross_ctx, query_capacity, (long long) required_n_ctx); const int32_t effective_draft_ctx = std::min(target_ctx, requested_draft_ctx);
if (requested_draft_ctx > target_ctx) {
LOG_INF("%s: DFlash draft context %d exceeds target context %d, clamping to target capacity\n",
__func__, requested_draft_ctx, target_ctx);
}
const int32_t cross_ctx = effective_draft_ctx - query_capacity;
if (cross_ctx <= 0) {
LOG_ERR("%s: invalid DFlash draft context size draft=%d target=%d query_capacity=%d, draft context must exceed the query block\n",
__func__, requested_draft_ctx, target_ctx, query_capacity);
return nullptr; return nullptr;
} }
cparams_dft.n_ctx = (uint32_t) required_n_ctx; cparams_dft.n_ctx = (uint32_t) effective_draft_ctx;
dflash_cross_ctx = cross_ctx;
LOG_INF("%s: DFlash context target/slot=%d logical=%d cross_ctx=%d query_block=%d\n",
__func__, target_ctx, effective_draft_ctx, cross_ctx, query_capacity);
} }
ctx_dft = llama_init_from_model(params.model_dft, cparams_dft); ctx_dft = llama_init_from_model(params.model_dft, cparams_dft);
@ -1431,7 +1435,7 @@ common_speculative * common_speculative_init(
config.type, config.type,
ctx_tgt, ctx_tgt,
ctx_dft, ctx_dft,
config.params.dflash_cross_ctx, dflash_cross_ctx,
query_capacity, query_capacity,
config.params.n_max); config.params.n_max);
if (!state->ready) { if (!state->ready) {
@ -2014,13 +2018,15 @@ bool common_speculative_load_draft_model(
LOG_INF("%s: loading draft model '%s'\n", __func__, params_dft.model.c_str()); LOG_INF("%s: loading draft model '%s'\n", __func__, params_dft.model.c_str());
if (params_dft.n_ctx == 0) { if (params.has_dflash_family_stage() && params.n_ctx > 0) {
params_dft.n_ctx = params.n_ctx; params_dft.n_ctx = params.n_ctx;
} }
if (params.has_dflash_family_stage() && params_dft.n_gpu_layers < 0) { if (params.has_dflash_family_stage() && params_dft.n_gpu_layers < 0) {
params_dft.n_gpu_layers = params_base.n_gpu_layers; params_dft.n_gpu_layers = params_base.n_gpu_layers;
} }
params_dft.n_ctx = params_dft.n_ctx == 0 ? params_base.n_ctx / params_base.n_parallel : params_dft.n_ctx; if (params_dft.n_ctx == 0) {
params_dft.n_ctx = params.n_ctx > 0 ? params.n_ctx : params_base.n_ctx / params_base.n_parallel;
}
params_dft.n_parallel = 1; params_dft.n_parallel = 1;
params.mparams_dft.path = params_dft.model; params.mparams_dft.path = params_dft.model;
@ -3178,12 +3184,17 @@ void common_speculative_context_shift(
spec->last_step_target_only = false; spec->last_step_target_only = false;
spec->t_step_start_us = 0; spec->t_step_start_us = 0;
} }
if (auto * ctx_mtp = common_speculative_get_companion_ctx(spec); ctx_mtp != nullptr) { auto * dflash_state = common_speculative_get_dflash_state(spec);
if (dflash_state == nullptr) {
auto * ctx_mtp = common_speculative_get_companion_ctx(spec);
if (ctx_mtp == nullptr) {
return;
}
llama_kv_cache_seq_rm (ctx_mtp, seq_id, kv_keep, kv_keep + kv_discard); llama_kv_cache_seq_rm (ctx_mtp, seq_id, kv_keep, kv_keep + kv_discard);
llama_kv_cache_seq_add(ctx_mtp, seq_id, kv_keep + kv_discard, kv_past, -kv_discard); llama_kv_cache_seq_add(ctx_mtp, seq_id, kv_keep + kv_discard, kv_past, -kv_discard);
} }
if (auto * dflash_state = common_speculative_get_dflash_state(spec); dflash_state != nullptr) { if (dflash_state != nullptr) {
dflash_context_shift(*dflash_state, kv_keep, kv_discard, kv_past); dflash_context_shift(*dflash_state, kv_keep, kv_discard, kv_past);
} }
} }

View File

@ -129,11 +129,11 @@ Check the details [here](./speculative.md).
| `-td, --threads-draft N` | Number of threads to use during generation | Same as `--threads` | | | `-td, --threads-draft N` | Number of threads to use during generation | Same as `--threads` | |
| `-tbd, --threads-batch-draft N` | Number of threads to use during batch and prompt processing | Same as `--threads-draft` | | | `-tbd, --threads-batch-draft N` | Number of threads to use during batch and prompt processing | Same as `--threads-draft` | |
| `-ps, --p-split N` | Speculative decoding split probability | 0.1 | | | `-ps, --p-split N` | Speculative decoding split probability | 0.1 | |
| `-cd, --ctx-size-draft N` | Size of the prompt context for the draft model | 0 (loaded from model) | Similar to `--ctx-size` but applied to the draft model, if used. | | `-cd, --ctx-size-draft N` | Size of the prompt context for the draft model | 0 (inherits target context for DFlash/DSpark, otherwise loaded from model) | Similar to `--ctx-size` but applied to the draft model, if used. For DFlash/DSpark, the retained history is this value minus the query block. |
| `-ctkd, --cache-type-k-draft TYPE` | KV cache data type for K for the draft model | - | For draft model, see: `-ctk` | | `-ctkd, --cache-type-k-draft TYPE` | KV cache data type for K for the draft model | - | For draft model, see: `-ctk` |
| `-ctvd, --cache-type-v-draft TYPE` | KV cache data type for V for the draft model | - | For draft model, see: `-ctk` | | `-ctvd, --cache-type-v-draft TYPE` | KV cache data type for V for the draft model | - | For draft model, see: `-ctk` |
| `-draft, --draft-params` | Comma-separated list of draft model parameters | - | | | `-draft, --draft-params` | Comma-separated list of draft model parameters | - | |
| `--spec-type SPEC[:k=v,...]` | Canonical speculative stage entry; repeat to configure the supported two-stage chain | - | Types: `none`, `draft`, `dflash`, `dspark`, `mtp`, `ngram-cache`, `ngram-simple`, `ngram-map-k`, `ngram-map-k4v`, `ngram-mod`, `suffix`. Canonical keys include `n_max`, `n_min`, `p_min`, `heads`, `cross_ctx`, `ngram_size_n`, `ngram_size_m`, `ngram_min_hits`, `suffix_min_match_len`, `suffix_max_depth`, `suffix_corpus`. For MTP, `heads=1` is the default; values above `1` and `heads=0` (all model heads) are experimental. String values may escape commas as `\,` or quote the value inside the stage payload. Examples: `--spec-type ngram-mod:n_max=64,n_min=2,ngram_size_n=8 --spec-type mtp:n_max=1,p_min=0.0`, `--model-draft draft.gguf --spec-type dflash:n_max=4,cross_ctx=512` | | `--spec-type SPEC[:k=v,...]` | Canonical speculative stage entry; repeat to configure the supported two-stage chain | - | Types: `none`, `draft`, `dflash`, `dspark`, `mtp`, `ngram-cache`, `ngram-simple`, `ngram-map-k`, `ngram-map-k4v`, `ngram-mod`, `suffix`. Canonical keys include `n_max`, `n_min`, `p_min`, `heads`, `ngram_size_n`, `ngram_size_m`, `ngram_min_hits`, `suffix_min_match_len`, `suffix_max_depth`, `suffix_corpus`. For MTP, `heads=1` is the default; values above `1` and `heads=0` (all model heads) are experimental. String values may escape commas as `\,` or quote the value inside the stage payload. Examples: `--spec-type ngram-mod:n_max=64,n_min=2,ngram_size_n=8 --spec-type mtp:n_max=1,p_min=0.0`, `--model-draft draft.gguf --spec-type dflash:n_max=4` |
| `--spec-autotune` | Automatically tune speculative params to maximize tokens/sec | - | Automatically determines the near-optimal arguments for the type of speculation being performed [PR 1595](https://github.com/ikawrakow/ik_llama.cpp/pull/1595) | | `--spec-autotune` | Automatically tune speculative params to maximize tokens/sec | - | Automatically determines the near-optimal arguments for the type of speculation being performed [PR 1595](https://github.com/ikawrakow/ik_llama.cpp/pull/1595) |
| `--spec-ckpt-mode MODE` (deprecated alias: `--recurrent-ckpt-mode MODE`) | Checkpoint strategy for speculative decoding | auto | One of: - `auto` choose direct per-step, then device fallback, then host fallback - `per-step` save architecture state per draft step; no re-decode on rejection - `gpu-fallback` copy architecture state to a device buffer; re-decode on rejection - `cpu` serialize architecture state in host storage; re-decode on rejection [PR 1669](https://github.com/ikawrakow/ik_llama.cpp/pull/1669) [PR 1774](https://github.com/ikawrakow/ik_llama.cpp/pull/1774) | | `--spec-ckpt-mode MODE` (deprecated alias: `--recurrent-ckpt-mode MODE`) | Checkpoint strategy for speculative decoding | auto | One of: - `auto` choose direct per-step, then device fallback, then host fallback - `per-step` save architecture state per draft step; no re-decode on rejection - `gpu-fallback` copy architecture state to a device buffer; re-decode on rejection - `cpu` serialize architecture state in host storage; re-decode on rejection [PR 1669](https://github.com/ikawrakow/ik_llama.cpp/pull/1669) [PR 1774](https://github.com/ikawrakow/ik_llama.cpp/pull/1774) |

View File

@ -59,9 +59,7 @@ static ggml_backend_t llama_backend_for_tensor(const llama_context & lctx, const
bool llama_context::ensure_dflash_kv_cache_tensors(int32_t cross_ctx) { bool llama_context::ensure_dflash_kv_cache_tensors(int32_t cross_ctx) {
const int32_t target_cross_ctx = std::max<int32_t>(1, cross_ctx); const int32_t target_cross_ctx = std::max<int32_t>(1, cross_ctx);
const int32_t target_token_capacity = std::max<int32_t>( const int32_t target_token_capacity = std::max<int32_t>(1, (int32_t) model.hparams.dflash_block_size);
1,
std::max<int32_t>((int32_t) model.hparams.dflash_block_size, (int32_t) cparams.n_ubatch));
const int32_t target_cache_n_kv_total = GGML_PAD(target_cross_ctx + target_token_capacity, (int32_t) llama_kv_cache::get_padding(cparams.flash_attn)); const int32_t target_cache_n_kv_total = GGML_PAD(target_cross_ctx + target_token_capacity, (int32_t) llama_kv_cache::get_padding(cparams.flash_attn));
const ggml_type target_cache_type = cparams.flash_attn ? GGML_TYPE_F16 : GGML_TYPE_F32; const ggml_type target_cache_type = cparams.flash_attn ? GGML_TYPE_F16 : GGML_TYPE_F32;
const int32_t n_layer = model.hparams.n_layer; const int32_t n_layer = model.hparams.n_layer;

View File

@ -7329,6 +7329,9 @@ static int32_t llama_kv_cache_update_internal(struct llama_context & lctx) {
// TODO: extract to a function // TODO: extract to a function
// build worst-case graph // build worst-case graph
int n_tokens = (int)std::min(lctx.cparams.n_ctx, lctx.cparams.n_ubatch); int n_tokens = (int)std::min(lctx.cparams.n_ctx, lctx.cparams.n_ubatch);
if (llm_arch_is_dflash_family(lctx.model.arch)) {
n_tokens = std::min(n_tokens, (int) lctx.model.hparams.dflash_block_size);
}
// MTP draft generation consumes one token and one hidden-state vector per decode step. // MTP draft generation consumes one token and one hidden-state vector per decode step.
if (lctx.cparams.mtp_op_type == MTP_OP_DRAFT_GEN) { if (lctx.cparams.mtp_op_type == MTP_OP_DRAFT_GEN) {
n_tokens = 1; n_tokens = 1;
@ -8366,6 +8369,13 @@ struct llama_context * llama_init_from_model(
ggml_type type_k = params.type_k; ggml_type type_k = params.type_k;
ggml_type type_v = params.type_v; ggml_type type_v = params.type_v;
if (llm_arch_is_dflash_family(model->arch)) {
const uint32_t pad = llama_kv_cache::get_padding(cparams.flash_attn);
kv_size = GGML_PAD(std::max<uint32_t>(1u, hparams.dflash_block_size), pad);
LLAMA_LOG_INFO("%s: DFlash ordinary KV bookkeeping capacity = %u (history owned by custom cache)\n",
__func__, kv_size);
}
// Mamba only needs a constant number of KV cache cells per sequence // Mamba only needs a constant number of KV cache cells per sequence
if (model->arch == LLM_ARCH_MAMBA) { if (model->arch == LLM_ARCH_MAMBA) {
// Mamba needs at least as many KV cells as there are sequences kept at any time // Mamba needs at least as many KV cells as there are sequences kept at any time
@ -8659,6 +8669,9 @@ struct llama_context * llama_init_from_model(
} }
int n_tokens = (int)std::min(cparams.n_ctx, cparams.n_ubatch); int n_tokens = (int)std::min(cparams.n_ctx, cparams.n_ubatch);
if (llm_arch_is_dflash_family(model->arch)) {
n_tokens = std::min(n_tokens, (int) model->hparams.dflash_block_size);
}
const size_t max_nodes = ctx->max_nodes(n_tokens, cparams.n_ctx); const size_t max_nodes = ctx->max_nodes(n_tokens, cparams.n_ctx);
// buffer used to store the computation graph and the tensor meta data // buffer used to store the computation graph and the tensor meta data