diff --git a/common/speculative-dflash-impl.h b/common/speculative-dflash-impl.h index cf573904..e0de29d9 100644 --- a/common/speculative-dflash-impl.h +++ b/common/speculative-dflash-impl.h @@ -85,10 +85,12 @@ struct common_speculative_state_dflash : public common_speculative_state { llama_batch batch = {}; int32_t block_size = 0; + int32_t query_capacity = 0; int32_t mask_token_id = -1; int32_t n_target_features = 0; int32_t cross_ctx = 0; bool is_dspark = false; + bool is_dsv4_dspark = false; bool ready = false; std::vector target_layer_ids; @@ -112,7 +114,9 @@ struct common_speculative_state_dflash : public common_speculative_state { enum common_speculative_type type, llama_context * ctx_tgt, llama_context * ctx_dft, - int32_t cross_ctx) + int32_t cross_ctx, + int32_t configured_query_capacity, + int32_t active_width) : common_speculative_state(type) , ctx_tgt(ctx_tgt) , ctx_dft(ctx_dft) @@ -122,6 +126,7 @@ struct common_speculative_state_dflash : public common_speculative_state { const llama_model * model_dft = llama_get_model(ctx_dft); is_dspark = type == COMMON_SPECULATIVE_TYPE_DSPARK; + is_dsv4_dspark = is_dspark && llama_model_is_deepseek4(model_tgt); const bool has_dspark_head = llama_model_dflash_has_dspark_head(model_dft); if (is_dspark != has_dspark_head) { LOG_ERR("%s: %s stage requires %s DSpark Markov tensors\n", __func__, @@ -146,6 +151,15 @@ struct common_speculative_state_dflash : public common_speculative_state { return; } + query_capacity = is_dsv4_dspark + ? std::max(block_size, configured_query_capacity) + : block_size; + if (is_dsv4_dspark && active_width > query_capacity) { + LOG_ERR("%s: DSV4 DSpark active width %d exceeds allocated query capacity %d\n", + __func__, active_width, query_capacity); + return; + } + target_layer_ids.resize((size_t) n_target_layers); if (llama_model_dflash_target_layer_ids(model_dft, target_layer_ids.data(), n_target_layers) != n_target_layers) { LOG_ERR("%s: failed to read DFlash target layer ids\n", __func__); @@ -220,7 +234,7 @@ struct common_speculative_state_dflash : public common_speculative_state { return; } - batch = llama_batch_init(std::max(1, block_size), 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); @@ -231,8 +245,9 @@ struct common_speculative_state_dflash : public common_speculative_state { llama_set_dflash_visible_cross_ctx(ctx_dft, this->cross_ctx); llama_set_dflash_dspark(ctx_dft, is_dspark); - LOG_INF("%s: DFlash context ready (n_ctx=%d, block_size=%d, cross_ctx=%d, n_target_features=%d, n_target_layers=%d)\n", - __func__, llama_n_ctx(ctx_dft), block_size, this->cross_ctx, n_target_features, n_target_layers); + 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, + this->cross_ctx, n_target_features, n_target_layers); } ~common_speculative_state_dflash() override { @@ -263,8 +278,15 @@ struct common_speculative_state_dflash : public common_speculative_state { return; } - const int32_t max_draft_tokens = is_dspark ? block_size : block_size - 1; - const int32_t n_keep = std::min(params.n_max, max_draft_tokens); + const int32_t max_draft_tokens = is_dsv4_dspark ? query_capacity : (is_dspark ? block_size : block_size - 1); + if (is_dsv4_dspark && params.n_max > query_capacity) { + LOG_ERR("%s: DSV4 DSpark runtime width %d exceeds allocated query capacity %d\n", + __func__, params.n_max, query_capacity); + return; + } + const int32_t n_keep = is_dsv4_dspark + ? params.n_max + : std::min(params.n_max, max_draft_tokens); if (n_keep <= 0) { return; } @@ -453,6 +475,12 @@ static bool dflash_append_target_features( } const int32_t n_rows = (int32_t) new_positions.size(); + for (int32_t i = 1; i < n_rows; ++i) { + if (new_positions[i] <= new_positions[i - 1]) { + return false; + } + } + if (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()); @@ -468,6 +496,50 @@ static bool dflash_append_target_features( return true; } + // 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()); + if (overlap != state.target_window_pos.end()) { + dflash_materialize_target_window_features(state); + + const int32_t prefix_rows = (int32_t) (overlap - state.target_window_pos.begin()); + const int32_t keep_old_rows = std::min(prefix_rows, state.cross_ctx - n_rows); + const int32_t old_start = prefix_rows - keep_old_rows; + const int32_t total_rows = keep_old_rows + n_rows; + const size_t total_floats = (size_t) total_rows * row_width; + + state.target_window_stage.resize(total_floats); + if (keep_old_rows > 0) { + std::copy( + state.target_window.begin() + (size_t) old_start * row_width, + state.target_window.begin() + (size_t) prefix_rows * row_width, + state.target_window_stage.begin()); + } + std::copy( + new_rows.begin(), new_rows.end(), + state.target_window_stage.begin() + (size_t) keep_old_rows * row_width); + + std::vector & next_window_pos = state.target_window_pos_stage; + next_window_pos.resize((size_t) total_rows); + if (keep_old_rows > 0) { + std::copy( + state.target_window_pos.begin() + old_start, + state.target_window_pos.begin() + prefix_rows, + next_window_pos.begin()); + } + 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_pos.swap(next_window_pos); + state.target_window_stage.clear(); + next_window_pos.clear(); + state.target_window_rows = total_rows; + state.target_window_ring_filled = total_rows; + dflash_ring_reset_rows(state, state.target_window.data(), total_rows); + state.last_target_pos = state.target_window_pos.back(); + dflash_record_window_update(state, 0, total_rows, true); + return true; + } + const int32_t keep_old_rows = std::min(state.target_window_rows, state.cross_ctx - n_rows); std::vector & next_window_pos = state.target_window_pos_stage; next_window_pos.resize((size_t) (keep_old_rows + n_rows)); diff --git a/common/speculative.cpp b/common/speculative.cpp index b641f3ef..c81a70e8 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -1278,6 +1278,17 @@ common_speculative * common_speculative_init( } const auto stages = params.get_resolved_stages(); + const llama_model * target_model = llama_get_model(ctx_tgt); + const bool is_dsv4_target = llama_model_is_deepseek4(target_model); + int32_t dsv4_dspark_query_capacity = 0; + for (const auto & stage : stages) { + if (stage.type != COMMON_SPECULATIVE_TYPE_DSPARK || !is_dsv4_target) { + continue; + } + + const int32_t stage_n_max = params.with_stage_overrides(stage).n_max; + dsv4_dspark_query_capacity = std::max(dsv4_dspark_query_capacity, std::max(1, stage_n_max)); + } if (params.model_dft && llama_model_is_gemma4_mtp_assistant(params.model_dft)) { const bool has_draft_stage = std::any_of(stages.begin(), stages.end(), [](const common_speculative_stage_params & stage) { return stage.type == COMMON_SPECULATIVE_TYPE_DRAFT; @@ -1329,10 +1340,11 @@ common_speculative * common_speculative_init( return nullptr; } - const int64_t required_n_ctx = (int64_t) max_cross_ctx + (int64_t) block_size; + 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; if (required_n_ctx > std::numeric_limits::max()) { - LOG_ERR("%s: invalid DFlash draft context size cross_ctx=%d block_size=%d required_n_ctx=%lld\n", - __func__, max_cross_ctx, block_size, (long long) required_n_ctx); + LOG_ERR("%s: invalid DFlash draft context size cross_ctx=%d query_capacity=%d required_n_ctx=%lld\n", + __func__, max_cross_ctx, query_capacity, (long long) required_n_ctx); return nullptr; } @@ -1368,7 +1380,6 @@ common_speculative * common_speculative_init( configs.push_back(common_speculative_config(stage, stage_params)); } - const llama_model * target_model = llama_get_model(ctx_tgt); if (!configs.empty() && common_speculative_needs_checkpoint(target_model)) { const int ckpt_tokens = std::max(1, params.get_max_stage_n_max() + 1); const int actual_mode = llama_spec_ckpt_init(ctx_tgt, params.spec_ckpt_mode, ckpt_tokens); @@ -1405,11 +1416,16 @@ common_speculative * common_speculative_init( } case COMMON_SPECULATIVE_TYPE_DFLASH: case COMMON_SPECULATIVE_TYPE_DSPARK: { + const int32_t query_capacity = config.type == COMMON_SPECULATIVE_TYPE_DSPARK && is_dsv4_target + ? std::max(1, dsv4_dspark_query_capacity) + : 0; auto state = std::make_unique( config.type, ctx_tgt, ctx_dft, - config.params.dflash_cross_ctx); + config.params.dflash_cross_ctx, + query_capacity, + config.params.n_max); if (!state->ready) { LOG_ERR("%s: failed to initialize %s speculative state\n", __func__, common_speculative_type_to_str(config.type).c_str());