diff --git a/common/speculative.cpp b/common/speculative.cpp index ebde17e1..c48e21ff 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -190,6 +190,9 @@ struct common_speculative_impl { virtual void accept(llama_seq_id seq_id, uint16_t n_accepted, bool is_other) = 0; + // decodes work for seq_id the implementation has put off; called when the sequence stops generating + virtual bool flush(llama_seq_id /*seq_id*/) { return true; } + // (optional) serialize/restore per-seq internal state (e.g. eagle3's deferred boundary). virtual bool get_state(llama_seq_id /*seq_id*/, std::vector & /*data*/) const { return false; } virtual void set_state(llama_seq_id /*seq_id*/, const std::vector & /*data*/) {} @@ -2133,6 +2136,14 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { std::vector i_last; std::vector> chain_h; + // Catch-up rows of small (verification) batches, per sequence, not yet decoded: draft() decodes the + // accepted ones in the same batch as its first draft row, so a round reads the MTP weights once + // instead of twice. Rejected rows (pos >= n_past) are dropped instead of decoded and then removed. + // Each row is decoded exactly as the catch-up would have decoded it, so the drafts do not change. + std::vector> deferred_tok; + std::vector> deferred_pos; + std::vector> deferred_h; + common_speculative_impl_draft_mtp(const common_params_speculative & params, uint32_t n_seq) : common_speculative_impl(COMMON_SPECULATIVE_TYPE_DRAFT_MTP, n_seq) , params(params.draft) @@ -2210,6 +2221,54 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { verify_h.assign(n_seq, {}); verify_h_rows.assign(n_seq, 0); + + deferred_tok.assign(n_seq, {}); + deferred_pos.assign(n_seq, {}); + deferred_h.assign(n_seq, {}); + } + + // drops the deferred rows of seq_id at positions >= pos + void drop_deferred_from(llama_seq_id seq_id, llama_pos pos) { + auto & toks = deferred_tok[seq_id]; + auto & poss = deferred_pos[seq_id]; + size_t n = 0; + while (n < poss.size() && poss[n] < pos) { + ++n; + } + toks.resize(n); + poss.resize(n); + deferred_h[seq_id].resize(n * (size_t) n_embd); + } + + // appends the deferred rows of seq_id at positions < pos_end to the batch, without outputs, and forgets them + void add_deferred(llama_seq_id seq_id, llama_pos pos_end) { + const size_t row_bytes = (size_t) n_embd * sizeof(float); + for (size_t r = 0; r < deferred_pos[seq_id].size() && deferred_pos[seq_id][r] < pos_end; ++r) { + common_batch_add(batch, deferred_tok[seq_id][r], deferred_pos[seq_id][r], { seq_id }, false); + std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, deferred_h[seq_id].data() + r * n_embd, row_bytes); + } + deferred_tok[seq_id].clear(); + deferred_pos[seq_id].clear(); + deferred_h[seq_id].clear(); + } + + // decodes the deferred rows of seq_id, or of every sequence for seq_id < 0 + bool flush_deferred(llama_seq_id only = -1) { + common_batch_clear(batch); + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + if (only < 0 || seq_id == only) { + add_deferred(seq_id, std::numeric_limits::max()); + } + } + if (batch.n_tokens == 0) { + return true; + } + const int32_t rc = llama_decode(params.ctx_dft, batch); + if (rc != 0) { + SPC_ERR("llama_decode(ctx_dft) of deferred rows failed rc=%d\n", (int) rc); + return false; + } + return true; } ~common_speculative_impl_draft_mtp() override { @@ -2284,8 +2343,39 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { const size_t row_bytes = (size_t) n_embd * sizeof(float); - // if kv is shared with target (e.g Gemma4), then we can skip this catch-up decode + // verification batches (at most n_max + 1 rows per sequence) wait for draft(), see deferred_tok + bool defer = false; if (!is_mem_shared) { + // a batch rewrites its sequences from its first position on + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + if (i_batch_beg[seq_id] >= 0) { + drop_deferred_from(seq_id, batch_in.pos[i_batch_beg[seq_id]]); + } + } + + defer = !chain_heads; + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + if (i_batch_beg[seq_id] >= 0 && i_batch_end[seq_id] - i_batch_beg[seq_id] + 1 > params.n_max + 1) { + defer = false; + } + } + + if (defer) { + const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt); + for (int k = 0; k < n_tokens; ++k) { + const llama_seq_id seq_id = batch_in.seq_id[k][0]; + const float * h = k == i_batch_beg[seq_id] ? pending_h[seq_id].data() : h_tgt + (size_t) (k - 1) * n_embd; + deferred_tok[seq_id].push_back(batch_in.token[k]); + deferred_pos[seq_id].push_back(batch_in.pos[k]); + deferred_h[seq_id].insert(deferred_h[seq_id].end(), h, h + n_embd); + } + } else if (!flush_deferred()) { + return false; + } + } + + // if kv is shared with target (e.g Gemma4), then we can skip this catch-up decode + if (!is_mem_shared && !defer) { common_batch_clear(batch); for (int k = 0; k < n_tokens; ++k) { @@ -2390,6 +2480,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { drafting[seq_id] = true; common_sampler_reset(smpls[seq_id].get()); + add_deferred(seq_id, dp.n_past); common_batch_add(batch, dp.id_last, dp.n_past, { seq_id }, true); std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, pending_h[seq_id].data(), row_bytes); @@ -2519,6 +2610,10 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { } } + bool flush(llama_seq_id seq_id) override { + return flush_deferred(seq_id); + } + void accept(llama_seq_id seq_id, uint16_t n_accepted, bool /*is_other*/) override { if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) { return; @@ -3584,6 +3679,18 @@ void common_speculative_accept(common_speculative * spec, llama_seq_id seq_id, u } } +bool common_speculative_flush(common_speculative * spec, llama_seq_id seq_id) { + if (spec == nullptr) { + return true; + } + + bool ok = true; + for (auto & impl : spec->impls) { + ok = impl->flush(seq_id) && ok; + } + return ok; +} + // TODO: support the case of more than one speculative implementations having a state bool common_speculative_get_state(common_speculative * spec, llama_seq_id seq_id, std::vector & data) { if (spec == nullptr) { diff --git a/common/speculative.h b/common/speculative.h index a7a4775a..5a946379 100644 --- a/common/speculative.h +++ b/common/speculative.h @@ -82,6 +82,9 @@ void common_speculative_draft(common_speculative * spec); // informs the speculative context that n_accepted tokens were accepted by the target model void common_speculative_accept(common_speculative * spec, llama_seq_id, uint16_t n_accepted); +// decodes draft-side work put off for seq_id (MTP catch-up rows); call when the sequence stops generating +bool common_speculative_flush(common_speculative * spec, llama_seq_id seq_id); + // (optional) get/set internal state bool common_speculative_get_state(common_speculative * spec, llama_seq_id seq_id, std::vector & data); void common_speculative_set_state(common_speculative * spec, llama_seq_id seq_id, const std::vector & data); diff --git a/ggml/src/ggml-cuda/concat.cu b/ggml/src/ggml-cuda/concat.cu index f5970889..49dc9b69 100644 --- a/ggml/src/ggml-cuda/concat.cu +++ b/ggml/src/ggml-cuda/concat.cu @@ -200,8 +200,7 @@ static void concat_cuda(const ggml_tensor * src0, const ggml_tensor * src1, ggml dim3 grid_dim(dst->ne[1], dst->ne[2], dst->ne[3]); if constexpr (sizeof(T) == sizeof(uint32_t)) { - const bool transpose_dim0 = ggml_cuda_info().devices[ggml_cuda_get_device()].cc == GGML_CUDA_CC_DGX_SPARK && - dim == 0 && src0->ne[2] == 1 && src0->ne[3] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1 && + const bool transpose_dim0 = dim == 0 && src0->ne[2] == 1 && src0->ne[3] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1 && dst->ne[2] == 1 && dst->ne[3] == 1 && src0->ne[0] <= 8 && src0->nb[0] == sizeof(uint32_t) && src0->nb[1] == (uint64_t) src0->ne[0]*sizeof(uint32_t) && src1->nb[1] == sizeof(uint32_t) && dst->nb[0] == sizeof(uint32_t) && diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -1255,6 +1255,10 @@ // flush the generated token stats before reset() if (slot.stats.n_gen > 0) { metrics_on_prediction(slot); + } + // the draft context of a sequence that stops generating must hold every accepted position + if (slot.can_speculate() && !common_speculative_flush(spec.get(), slot.id)) { + SLT_ERR(slot, "%s", "failed to flush the draft context\n"); } };