Initial dflash implementation

This commit is contained in:
SamuelOliveirads
2026-05-28 18:57:58 -03:00
parent 6eff055a0c
commit 82cff238fe
23 changed files with 1618 additions and 205 deletions
+67 -53
View File
@@ -134,6 +134,14 @@ static bool server_speculative_has_mtp(const common_params_speculative & spec) {
return spec.has_stage_type(COMMON_SPECULATIVE_TYPE_MTP);
}
static bool server_speculative_has_dflash(const common_params_speculative & spec) {
return spec.has_stage_type(COMMON_SPECULATIVE_TYPE_DFLASH);
}
static bool server_speculative_has_target_features(const common_params_speculative & spec) {
return server_speculative_has_mtp(spec) || server_speculative_has_dflash(spec);
}
static bool server_speculative_same_stage_types(
const common_params_speculative & lhs,
const common_params_speculative & rhs) {
@@ -217,6 +225,18 @@ static common_speculative_stage_params server_parse_speculative_stage_json(const
}
server_context::~server_context() {
// Speculative state may reference the live target context during teardown.
for (server_slot& slot : slots) {
if (slot.ctx_sampling != nullptr) {
common_sampler_free(slot.ctx_sampling);
}
slot.spec_ckpt.clear();
common_speculative_free(slot.spec);
slot.spec = nullptr;
slot.ctx_dft = nullptr;
llama_batch_free(slot.batch_spec);
}
if (ctx) {
llama_free(ctx);
ctx = nullptr;
@@ -238,19 +258,6 @@ server_context::~server_context() {
model_draft = nullptr;
}
// Clear any sampling context
for (server_slot& slot : slots) {
if (slot.ctx_sampling != nullptr) {
common_sampler_free(slot.ctx_sampling);
}
slot.spec_ckpt.clear();
if (slot.ctx_dft) {
llama_free(slot.ctx_dft);
}
common_speculative_free(slot.spec);
llama_batch_free(slot.batch_spec);
}
llama_batch_free(batch);
}
@@ -286,6 +293,13 @@ bool server_context::load_model(const gpt_params& params_) {
params_base.speculative.model_dft = nullptr;
}
if (server_speculative_has_dflash(params_base.speculative) && params_base.n_parallel > 1) {
LOG_ERROR("DFlash is currently limited to a single server slot (-np 1).\n", {
{"n_parallel", params_base.n_parallel},
});
return false;
}
bool has_draft_model = !params_base.speculative.model.empty() || !params_base.speculative.params.empty();
std::string& mmproj_path = params_base.mmproj.path;
if (!mmproj_path.empty()) {
@@ -470,7 +484,7 @@ void server_context::init() {
bool can_spec = true;
if (!params_base.dry_run) {
can_spec = common_speculative_is_compat(ctx);
}
}
if (!can_spec) {
SRV_WRN("%s", "speculative decoding not supported by this context\n");
}
@@ -1656,7 +1670,7 @@ bool server_context::launch_slot_with_task(server_slot& slot, server_task& task)
int32_t banbuffer_size = json_value(data, "banbuffer_size", 0);
slot.n_buffer = 0; // Ensure buffer calculation starts fresh for this slot
slot.rewind_count_max = json_value(data, "rewind_count_max", -1);
const auto& banned_strings = data.find("banned_strings");
if (banned_strings != data.end() && banned_strings->is_array()) {
slot.ban_phrases.clear();
@@ -2805,7 +2819,7 @@ static size_t load_server_tokens_from_file(const std::string & filename, server
size_t pos = 0;
json token_json;
if (file.is_open()) {
file >> token_json;
file >> token_json;
pos = file.tellg();
file.close();
}
@@ -3727,7 +3741,7 @@ bool server_context::create_checkpoint(server_slot & slot) {
slot.server_cached_prompt.checkpoints.erase(slot.server_cached_prompt.checkpoints.begin());
}
auto & cur = slot.server_cached_prompt.checkpoints.emplace_back();
server_prompt_checkpoint_update(cur, ctx, slot.id, slot.cache_tokens.n_tokens(), pos_min, pos_max, slot.n_past_offset);
@@ -4060,7 +4074,7 @@ void server_context::batch_pending_prompt(const int32_t n_ubatch, const int32_t
slot.do_checkpoint = true;
break;
}
}
LOG_VERBOSE("prompt processing progress", {
{"id_slot", slot.id},
@@ -4143,7 +4157,7 @@ static void restore_speculative_checkpoint(
common_speculative_type spec_type_used,
llama_token sampled_before,
const std::vector<llama_token> & ids, int n_draft,
const std::vector<float> & mtp_hidden_state_pre, int32_t mtp_n_past_base) {
const std::vector<float> & spec_feature_rows_pre, int32_t spec_n_past_base) {
if (slot.spec_ckpt.per_step_enabled) {
const int step = (int)ids.size() - 1;
llama_spec_ckpt_restore(ctx, slot.id, slot.spec_ckpt.n_past, step);
@@ -4155,16 +4169,16 @@ static void restore_speculative_checkpoint(
common_sampler_accept(slot.ctx_sampling, ctx, id, true);
}
// Update MTP KV cache and hidden state using embeddings collected before checkpoint restore.
if (slot.has_mtp && !mtp_hidden_state_pre.empty()) {
// Update speculative target features using rows collected before checkpoint restore.
if (server_speculative_has_target_features(slot.params.speculative) && !spec_feature_rows_pre.empty()) {
if (!common_speculative_commit_accepted_hidden_rows(
slot.spec,
spec_type_used,
slot.id,
mtp_n_past_base,
spec_n_past_base,
sampled_before,
ids,
mtp_hidden_state_pre)) {
spec_feature_rows_pre)) {
common_speculative_clear_sequence_hidden(slot.spec, slot.id);
} else if (spec_type_used != COMMON_SPECULATIVE_TYPE_MTP) {
SLT_DBG(slot, "%s", "synced MTP target hidden state from accepted-prefix rows after per-step restore");
@@ -4201,7 +4215,7 @@ static void restore_speculative_checkpoint(
if (ret != 0) {
SLT_ERR(slot, "failed to re-decode accepted tokens after checkpoint restore: %d\n", ret);
}
if (slot.has_mtp) {
if (server_speculative_has_target_features(slot.params.speculative)) {
const int n_accepted = (int)ids.size();
std::vector<int32_t> redecoded_indices(n_accepted);
for (int j = 0; j < n_accepted; ++j) {
@@ -4272,20 +4286,20 @@ void server_context::speculative_decoding_accept() {
}
const bool any_rejected = (ids.size() - 1) < n_draft;
int32_t mtp_n_past_base = 0;
std::vector<float> mtp_hidden_state_pre;
int32_t spec_n_past_base = 0;
std::vector<float> spec_feature_rows_pre;
std::vector<int32_t> accepted_output_indices;
if (slot.has_mtp) {
if (server_speculative_has_target_features(slot.params.speculative)) {
const int32_t n_pre_spec_tokens = slot.cache_tokens.n_tokens() - (int32_t)(slot.drafted.size() + 1);
mtp_n_past_base = slot.cache_tokens.pos_next(n_pre_spec_tokens);
spec_n_past_base = slot.cache_tokens.pos_next(n_pre_spec_tokens);
if (!ids.empty()) {
accepted_output_indices.assign(slot.i_batch_dft.begin(), slot.i_batch_dft.begin() + ids.size());
}
if (any_rejected && slot.spec_ckpt.valid && !accepted_output_indices.empty()) {
if (!common_speculative_copy_output_hidden_rows(slot.spec, ctx, accepted_output_indices, mtp_hidden_state_pre)) {
mtp_hidden_state_pre.clear();
if (!common_speculative_copy_output_hidden_rows(slot.spec, ctx, accepted_output_indices, spec_feature_rows_pre)) {
spec_feature_rows_pre.clear();
}
}
}
@@ -4317,15 +4331,15 @@ void server_context::speculative_decoding_accept() {
// for recurrent/hybrid models: if any drafts were rejected, restore recurrent state
if (any_rejected && slot.spec_ckpt.valid) {
restore_speculative_checkpoint(slot, ctx, model, spec_type_used, sampled_before, ids, n_draft, mtp_hidden_state_pre, mtp_n_past_base);
restore_speculative_checkpoint(slot, ctx, model, spec_type_used, sampled_before, ids, n_draft, spec_feature_rows_pre, spec_n_past_base);
} else {
if (slot.has_mtp && !accepted_output_indices.empty()) {
if (server_speculative_has_target_features(slot.params.speculative) && !accepted_output_indices.empty()) {
if (!common_speculative_commit_accepted_output(
slot.spec,
ctx,
spec_type_used,
slot.id,
mtp_n_past_base,
spec_n_past_base,
sampled_before,
ids,
accepted_output_indices)) {
@@ -4395,15 +4409,15 @@ void server_context::release_slot_after_final_response(server_slot & slot) {
void server_context::send_token_results(completion_token_outputs& results, server_slot& slot, int32_t n) {
int count = 0;
bool released = false;
int32_t start_pos = slot.n_past - (int32_t)slot.token_buffer.size() + 1;
for (auto& it : results) {
bool has_next = process_token(it, slot);
// Clean up positional bans for the token we just confirmed/sent
slot.positional_bans.erase(start_pos + count);
count++;
if (!has_next) {
if (slot.stopped_limit && !slot.stopped_eos && !slot.stopped_word) {
@@ -4436,7 +4450,7 @@ inline int32_t check_ban_phrase(server_slot& slot) {
std::string string_buffer;
std::vector<size_t> token_offsets;
for (const auto& it : slot.token_buffer) {
token_offsets.push_back(string_buffer.size());
string_buffer += it.text_to_send;
@@ -4488,10 +4502,10 @@ inline int32_t check_ban_phrase(server_slot& slot) {
if (found) {
int32_t token_idx = -1;
for (size_t i = 0; i < token_offsets.size(); ++i) {
size_t len = (i == token_offsets.size() - 1)
? string_buffer.size() - token_offsets[i]
size_t len = (i == token_offsets.size() - 1)
? string_buffer.size() - token_offsets[i]
: token_offsets[i+1] - token_offsets[i];
if (best_start >= token_offsets[i] && best_start < token_offsets[i] + len) {
token_idx = (int32_t)i;
break;
@@ -4509,7 +4523,7 @@ inline int32_t check_ban_phrase(server_slot& slot) {
inline void rewind_context(server_slot& slot, int32_t ban_pos) {
slot.rewind_count++;
int32_t buffer_start_pos = slot.n_past - (int32_t)slot.token_buffer.size() + 1;
int32_t n_keep_buffer = ban_pos - buffer_start_pos;
if (n_keep_buffer < 0) n_keep_buffer = 0;
@@ -4518,9 +4532,9 @@ inline void rewind_context(server_slot& slot, int32_t ban_pos) {
int32_t n = 0;
for (auto result = slot.token_buffer.begin() + n_keep_buffer; result != slot.token_buffer.end(); result++) {
llama_token banned_tok = result->tok;
if (n == 0) {
LLAMA_LOG_DEBUG("Banned pattern detected at pos %d. Banning token %d ('%s') and rewinding.\n",
LLAMA_LOG_DEBUG("Banned pattern detected at pos %d. Banning token %d ('%s') and rewinding.\n",
ban_pos, banned_tok, result->text_to_send.c_str());
}
@@ -4533,7 +4547,7 @@ inline void rewind_context(server_slot& slot, int32_t ban_pos) {
}
int32_t n_rewind_total = (slot.n_past + 1) - ban_pos;
size_t n_keep_cache = 0;
if (ban_pos > 0) {
n_keep_cache = (size_t)(ban_pos - 1);
@@ -4546,13 +4560,13 @@ inline void rewind_context(server_slot& slot, int32_t ban_pos) {
if (n_keep_cache < slot.cache_tokens.size()) {
slot.sampled = slot.cache_tokens[n_keep_cache];
} else {
slot.sampled = 0;
slot.sampled = 0;
}
// Truncate cache
slot.cache_tokens.keep_first(n_keep_cache);
slot.n_past = slot.cache_tokens.n_tokens();
// Remove from KV cache
llama_kv_cache_seq_rm(slot.ctx, slot.id, slot.cache_tokens.pos_next(slot.n_past), -1);
@@ -4590,13 +4604,13 @@ void server_context::buffer_and_check_string_ban(server_slot & slot, completion_
// Automatic / Heuristic logic
// Account for strings + regex + regex_ci
size_t total_bans = slot.ban_phrases.size() + slot.ban_regex.size() + slot.ban_regex_ci.size();
// Heuristic: Allow if under 20 OR under 2 * total_bans
// Conversely: Stop if >= 20 AND > 2 * total_bans
if (slot.rewind_count >= 20 && slot.rewind_count > 2 * total_bans) {
allow_rewind = false;
}
}
}
else if (slot.rewind_count_max > 0) {
// Strict limit logic
if (slot.rewind_count >= slot.rewind_count_max) {
@@ -4613,7 +4627,7 @@ void server_context::buffer_and_check_string_ban(server_slot & slot, completion_
else if (buffer_full || !next_token) {
slot.rewind_status = false;
slot.rewind_count = 0;
if (!next_token) {
// send all remaining tokens
send_token_results(slot.token_buffer, slot);
@@ -4625,7 +4639,7 @@ void server_context::buffer_and_check_string_ban(server_slot & slot, completion_
}
else {
// buffer the result, wait for more tokens to validate string
slot.sampled = result.tok;
slot.sampled = result.tok;
}
}
@@ -4710,9 +4724,9 @@ void server_context::process_batch_tokens(int32_t & n_batch) {
continue; // continue loop of n_batch
}
if (server_speculative_has_mtp(params_base.speculative)) {
if (server_speculative_has_target_features(params_base.speculative)) {
for (auto & slot : slots) {
if (!slot.spec || !slot.has_mtp) {
if (!slot.spec || !server_speculative_has_target_features(slot.params.speculative)) {
continue;
}
@@ -4722,7 +4736,7 @@ void server_context::process_batch_tokens(int32_t & n_batch) {
}
if (common_speculative_on_target_seq_batch(slot.spec, ctx, batch_view, slot.id, true) != 0) {
LOG_ERROR("failed to warm up MTP state from prompt batch for slot %d\n", slot.id);
LOG_ERROR("failed to warm up speculative target-feature state from prompt batch for slot %d\n", slot.id);
}
}
}