Skip to content
Merged
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 46 additions & 0 deletions common/speculative.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -922,6 +922,9 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {

std::vector<common_sampler_ptr> smpls;

// backend sampler chain per seq, attached to ctx_dft
std::vector<llama_sampler *> backend_chains;

int32_t n_embd_dec = 0; // draft hidden size
int32_t n_embd_enc = 0; // target_layer_ids_n * target_hidden_size
int32_t n_embd_tgt = 0; // target model hidden size
Expand Down Expand Up @@ -995,6 +998,22 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
s.reset(common_sampler_init(model_dft, sparams));
}

// offload draft sampling to the backend
backend_chains.assign(n_seq, nullptr);
if (this->params.backend_sampling) {
Comment thread
ruixiang63 marked this conversation as resolved.
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
llama_sampler * chain = llama_sampler_chain_init(llama_sampler_chain_default_params());
llama_sampler_chain_add(chain, llama_sampler_init_top_k(10));

if (!llama_set_sampler(ctx_dft, seq_id, chain)) {
SPC_WRN("backend offload failed for seq_id=%d; using CPU sampler\n", (int) seq_id);
llama_sampler_free(chain);
chain = nullptr;
}
backend_chains[seq_id] = chain;
}
}

// turn on extraction of the target layers' input embeddings
for (uint32_t k = 0; k < target_layer_ids_n; ++k) {
llama_set_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k], true);
Expand All @@ -1005,6 +1024,18 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
}

~common_speculative_impl_draft_dflash() override {
auto * ctx_dft = this->params.ctx_dft;
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) backend_chains.size(); ++seq_id) {
if (backend_chains[seq_id] == nullptr) {
continue;
}
if (ctx_dft) {
llama_set_sampler(ctx_dft, seq_id, nullptr);
}
llama_sampler_free(backend_chains[seq_id]);
}
backend_chains.clear();

llama_batch_free(batch);
llama_batch_free(batch_inject);
}
Expand Down Expand Up @@ -2301,6 +2332,21 @@ common_params common_base_params_to_speculative(const common_params & params) {
result.n_outputs_max = params.n_parallel;
result.n_outputs_max_per_seq = 1;

// dflash/dspark decode the whole noise block in a single pass and sample every block position on the backend
const bool has_block_draft = std::any_of(
params.speculative.types.begin(), params.speculative.types.end(),
[](common_speculative_type t) {
return t == COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH || t == COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK;
});
Comment thread
ggerganov marked this conversation as resolved.
if (has_block_draft) {
// per-seq output positions: DFlash decodes anchor + n_max masks (n_max + 1); DSpark n_max -> +1 covers both
const int32_t per_seq = std::max(1, params_spec.n_max + 1);
result.n_outputs_max = params.n_parallel * per_seq;
if (params_spec.backend_sampling) {
result.n_outputs_max_per_seq = per_seq;
}
}

return result;
}

Expand Down
Loading