This repository has no description
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206diff --git a/common/speculative.cpp b/common/speculative.cppindex 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<uint8_t> & /*data*/) const { return false; } virtual void set_state(llama_seq_id /*seq_id*/, const std::vector<uint8_t> & /*data*/) {}@@ -2133,6 +2136,14 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { std::vector<int> i_last; std::vector<std::vector<float>> 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<std::vector<llama_token>> deferred_tok;+ std::vector<std::vector<llama_pos>> deferred_pos;+ std::vector<std::vector<float>> 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<llama_pos>::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<uint8_t> & data) { if (spec == nullptr) {diff --git a/common/speculative.h b/common/speculative.hindex 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<uint8_t> & data); void common_speculative_set_state(common_speculative * spec, llama_seq_id seq_id, const std::vector<uint8_t> & data);diff --git a/ggml/src/ggml-cuda/concat.cu b/ggml/src/ggml-cuda/concat.cuindex 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"); } };