diff --git a/common/arg.cpp b/common/arg.cpp index 3f01b02599..8c9bc2cdb9 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -4209,6 +4209,21 @@ common_params_context common_params_parser_init(common_params & params, llama_ex params.speculative.draft.backend_sampling = value; } ).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_DRAFT_BACKEND_SAMPLING")); + add_opt(common_arg( + {"--spec-draft-sampling"}, "{greedy,probabilistic}", + string_format("how the draft is sampled: greedy takes its argmax, probabilistic samples it and has " + "the target verify by rejection sampling (default: %s)", + params.speculative.draft.probabilistic ? "probabilistic" : "greedy"), + [](common_params & params, const std::string & value) { + if (value == "greedy") { + params.speculative.draft.probabilistic = false; + } else if (value == "probabilistic") { + params.speculative.draft.probabilistic = true; + } else { + throw std::invalid_argument("invalid value, must be one of: greedy, probabilistic"); + } + } + ).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_DRAFT_SAMPLING")); add_opt(common_arg( {"--spec-draft-device", "-devd", "--device-draft"}, "", "comma-separated list of devices to use for offloading the draft model (none = don't offload, default: follows --device)\n" diff --git a/common/common.h b/common/common.h index 5700d20c58..f4b72d90a3 100644 --- a/common/common.h +++ b/common/common.h @@ -333,6 +333,8 @@ struct common_params_speculative_draft { bool backend_sampling = true; // offload draft sampling to the backend (default: on) + bool probabilistic = false; // sample the draft and verify by rejection, instead of argmax and match + common_params_model mparams; llama_context * ctx_tgt = nullptr; diff --git a/common/sampling.cpp b/common/sampling.cpp index e9e1cb372e..d9c508049d 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -12,6 +12,7 @@ #include #include #include +#include #include #include @@ -121,6 +122,9 @@ struct common_sampler { llama_token_data_array cur_p; + // for rejection sampling; independent of the draft, or the target distribution is not preserved + std::mt19937 rng; + void reset() { prev.clear(); @@ -432,6 +436,8 @@ struct common_sampler * common_sampler_init( /* .prev = */ ring_buffer(std::max(32, params.n_prev)), /* .cur = */ {}, /* .cur_p = */ {}, + // mix it, the chain and the draft are seeded from this one too + /* .rng = */ std::mt19937(llama_sampler_get_seed(chain) ^ 0x9e3779b9u), }; return result; @@ -515,6 +521,7 @@ struct common_sampler * common_sampler_clone(common_sampler * gsmpl) { /* .prev = */ gsmpl->prev, /* .cur = */ gsmpl->cur, /* .cur_p = */ gsmpl->cur_p, + /* .rng = */ gsmpl->rng, }; } @@ -535,6 +542,7 @@ void common_sampler_copy(const common_sampler * src, common_sampler * dst) { dst->cur = src->cur; dst->cur_p = src->cur_p; dst->cur_p.data = src->cur_p.data ? dst->cur.data() : nullptr; // re-point to dst's buffer + dst->rng = src->rng; dst->t_total_us = src->t_total_us; } @@ -709,6 +717,124 @@ std::vector common_sampler_sample_and_accept_n(struct common_sample return result; } +static float prob_of(const llama_token_data * data, size_t n, llama_token id) { + for (size_t k = 0; k < n; ++k) { + if (data[k].id == id) { + return data[k].p; + } + } + return 0.0f; +} + +// Accept a drafted token with probability min(1, p/q), else draw from norm(max(0, p - q)). +// Preserves the target distribution exactly, and accepts more often than matching does when the +// draft samples instead of taking its argmax. +std::vector common_sampler_sample_and_accept_n_rejection(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector & idxs, const llama_tokens & draft, const std::vector> & draft_q, bool grammar_first) { + GGML_ASSERT(idxs.size() == draft.size() + 1 && "idxs.size() must be draft.size() + 1"); + GGML_ASSERT(draft_q.size() == draft.size() && "draft_q must have one entry per draft token"); + + std::vector result; + result.reserve(idxs.size()); + + // draws come from the sampler's own stream, so they stay independent of what was drafted + std::uniform_real_distribution uni(0.0f, 1.0f); + + std::vector residual; + + std::vector cand; // candidate array masked by the grammar, if there is one + + size_t i = 0; + for (; i < draft.size(); i++) { + // leaves the target distribution in the candidate array + const llama_token id_tgt = common_sampler_sample(gsmpl, ctx, idxs[i], grammar_first); + + const auto * cur_p = common_sampler_get_candidates(gsmpl, true); + const auto & q = draft_q[i]; + + const bool masked = !grammar_first && grammar_should_apply(gsmpl); + if (masked) { + cand.assign(cur_p->data, cur_p->data + cur_p->size); + llama_token_data_array arr = { cand.data(), cand.size(), -1, false }; + llama_sampler_apply(gsmpl->grmr, &arr); + } + + // a candidate the grammar rejects carries no probability, whatever the target thinks + auto p_raw = [&](size_t k) { + return masked && cand[k].logit == -INFINITY ? 0.0f : cur_p->data[k].p; + }; + + // masking drops probability mass, so rescale what is left or the residual is over-weighted + float p_sum = 0.0f; + if (masked) { + for (size_t k = 0; k < cur_p->size; ++k) { + p_sum += p_raw(k); + } + } + + const float p_norm = masked && p_sum > 0.0f ? 1.0f/p_sum : 1.0f; + + auto p_of = [&](size_t k) { + return p_raw(k)*p_norm; + }; + + // q_x is never 0 for a token the draft produced, but guard the divide + const float q_x = prob_of(q.data(), q.size(), draft[i]); + + float p_x = 0.0f; + for (size_t k = 0; k < cur_p->size; ++k) { + if (cur_p->data[k].id == draft[i]) { + p_x = p_of(k); + break; + } + } + + if (q_x > 0.0f && (p_x >= q_x || uni(gsmpl->rng) < p_x / q_x)) { + common_sampler_accept(gsmpl, draft[i], true); + result.push_back(draft[i]); + continue; + } + + // rejected: tokens outside q's support keep all of p + residual.clear(); + float sum = 0.0f; + for (size_t k = 0; k < cur_p->size; ++k) { + const float r = p_of(k) - prob_of(q.data(), q.size(), cur_p->data[k].id); + if (r > 0.0f) { + residual.push_back({ cur_p->data[k].id, 0.0f, r }); + sum += r; + } + } + + llama_token id = id_tgt; + if (sum > 0.0f) { + float u = uni(gsmpl->rng) * sum; + id = residual.back().id; + for (const auto & e : residual) { + u -= e.p; + if (u <= 0.0f) { + id = e.id; + break; + } + } + } + + common_sampler_accept(gsmpl, id, true); + result.push_back(id); + + break; + } + + if (i == draft.size()) { + const llama_token id = common_sampler_sample(gsmpl, ctx, idxs[i], grammar_first); + + common_sampler_accept(gsmpl, id, true); + + result.push_back(id); + } + + return result; +} + std::vector common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const llama_tokens & draft, bool grammar_first) { std::vector idxs(draft.size() + 1); for (size_t i = 0; i < idxs.size(); ++i) { diff --git a/common/sampling.h b/common/sampling.h index ced3c8364b..7ebae3df82 100644 --- a/common/sampling.h +++ b/common/sampling.h @@ -85,6 +85,9 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co // std::vector common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector & idxs, const llama_tokens & draft, bool grammar_first = false); +// as above, but verifies by rejection sampling; draft_q holds the draft's candidates per token +std::vector common_sampler_sample_and_accept_n_rejection(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector & idxs, const llama_tokens & draft, const std::vector> & draft_q, bool grammar_first = false); + // assume idxs == [ 0, 1, 2, ..., draft.size() ] std::vector common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const llama_tokens & draft, bool grammar_first = false); diff --git a/common/speculative.cpp b/common/speculative.cpp index 5c36c9ca50..328ed241a6 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -30,6 +30,45 @@ #define SPEC_VOCAB_MAX_SIZE_DIFFERENCE 128 #define SPEC_VOCAB_CHECK_START_TOKEN_ID 5 +// Rebuild seq_id's draft sampler at the target's temperature: rejection weighs q against p, so +// both have to sample alike. Only temp and seed carry over; the draft keeps its own top_k. +static void spec_retune( + std::vector & smpls, + std::vector & cfg, + const llama_model * model, + llama_seq_id seq_id, + float temp, + uint32_t seed) { + if (cfg.size() != smpls.size()) { + const size_t n_old = cfg.size(); + cfg.resize(smpls.size()); + + // the initial sampler has no temperature, so no request may match the cache and skip a rebuild + for (size_t i = n_old; i < cfg.size(); ++i) { + cfg[i].temp = NAN; + } + } + + auto & cur = cfg[seq_id]; + + if (cur.temp == temp && cur.seed == seed) { + return; + } + + cur.temp = temp; + cur.seed = seed; + + common_params_sampling sparams; + sparams.no_perf = false; + sparams.top_k = 10; + sparams.temp = cur.temp; + // must be explicit, the default reseeds at random; mixed so it differs from the target's + sparams.seed = cur.seed == LLAMA_DEFAULT_SEED ? cur.seed : cur.seed ^ 0x85ebca6bu; + sparams.samplers = { COMMON_SAMPLER_TYPE_TOP_K, COMMON_SAMPLER_TYPE_TEMPERATURE }; + + smpls[seq_id].reset(common_sampler_init(model, sparams)); +} + const std::map common_speculative_type_from_name_map = { {"none", COMMON_SPECULATIVE_TYPE_NONE}, {"draft-simple", COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE}, @@ -187,6 +226,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl { std::vector smpls; + std::vector smpls_cfg; + common_speculative_impl_draft_simple(const common_params_speculative & params, uint32_t n_seq) : common_speculative_impl(COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE, n_seq, params.draft.n_max) , params(params.draft) @@ -255,8 +296,9 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl { } } - void begin(llama_seq_id /*seq_id*/, const llama_tokens & /*prompt*/) override { - // noop + void begin(llama_seq_id seq_id, const llama_tokens & /*prompt*/) override { + // reset here rather than per round, or two identical requests differ + common_sampler_reset(smpls[seq_id].get()); } bool process(const common_batch & batch_in) override { @@ -323,7 +365,20 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl { n_drafting++; drafting[seq_id] = true; - common_sampler_reset(smpls[seq_id].get()); + // greedy drafting leaves no candidates behind, so the verifier falls back to sample-and-match + if (!params.probabilistic) { + dp.result_q = nullptr; + } + + // result_q is only set when the caller wants rejection, so it also gates the retune + if (dp.result_q) { + spec_retune(smpls, smpls_cfg, llama_get_model(ctx_dft), seq_id, dp.temp, dp.seed); + } + + // a reset reseeds the chain, which breaks probabilistic drafting + if (!dp.result_q) { + common_sampler_reset(smpls[seq_id].get()); + } batch.add(dp.id_last, dp.pos0, seq_id, true); } @@ -348,7 +403,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl { auto * smpl = smpls[seq_id].get(); - common_sampler_sample(smpl, ctx_dft, i_batch, true); + const llama_token id_sampled = common_sampler_sample(smpl, ctx_dft, i_batch, true); ++i_batch; const auto * cur_p = common_sampler_get_candidates(smpl, true); @@ -360,7 +415,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl { } // add drafted token for each sequence - const llama_token id = cur_p->data[0].id; + const llama_token id = dparams.at(seq_id).result_q ? id_sampled : cur_p->data[0].id; // only collect very high-confidence draft tokens if (cur_p->data[0].p < params.p_min) { @@ -377,6 +432,10 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl { result.push_back(id); + if (dp.result_q) { + dp.result_q->emplace_back(cur_p->data, cur_p->data + cur_p->size); + } + if ((params.n_max <= (int) result.size()) || (dp.n_max > 0 && dp.n_max <= (int) result.size())) { drafting[seq_id] = false; @@ -1335,6 +1394,8 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { std::vector smpls; + std::vector smpls_cfg; + // backend sampler chain per seq, attached to ctx_dft std::vector backend_chains; @@ -1455,6 +1516,9 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { } void begin(llama_seq_id seq_id, const llama_tokens & prompt) override { + // reset here rather than per round, or two identical requests differ + common_sampler_reset(smpls[seq_id].get()); + const int32_t N = (int32_t) prompt.size(); if (N <= 0) { return; @@ -1599,7 +1663,20 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { n_drafting++; drafting[seq_id] = true; - common_sampler_reset(smpls[seq_id].get()); + // greedy drafting leaves no candidates behind, so the verifier falls back to sample-and-match + if (!params.probabilistic) { + dp.result_q = nullptr; + } + + // result_q is only set when the caller wants rejection, so it also gates the retune + if (dp.result_q) { + spec_retune(smpls, smpls_cfg, llama_get_model(ctx_dft), seq_id, dp.temp, dp.seed); + } + + // a reset reseeds the chain, which breaks probabilistic drafting + if (!dp.result_q) { + common_sampler_reset(smpls[seq_id].get()); + } const int32_t idx = batch.add(dp.id_last, dp.pos0, seq_id, true); batch.set_embd(idx, { pending_h[seq_id].data(), 1, (size_t) n_embd }); @@ -1648,7 +1725,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { auto * smpl = smpls[seq_id].get(); - common_sampler_sample(smpl, ctx_dft, i_last[seq_id], true); + const llama_token id_sampled = common_sampler_sample(smpl, ctx_dft, i_last[seq_id], true); const float * h_row = llama_get_embeddings_nextn_ith(ctx_dft, i_last[seq_id]); const auto * cur_p = common_sampler_get_candidates(smpl, true); @@ -1660,7 +1737,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { } // add drafted token for each sequence - const llama_token id = cur_p->data[0].id; + const llama_token id = dparams.at(seq_id).result_q ? id_sampled : cur_p->data[0].id; // only collect very high-confidence draft tokens if (cur_p->data[0].p < params.p_min) { @@ -1677,6 +1754,10 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { result.push_back(id); + if (dp.result_q) { + dp.result_q->emplace_back(cur_p->data, cur_p->data + cur_p->size); + } + if (params.n_max <= (int) result.size()) { drafting[seq_id] = false; n_drafting--; @@ -2833,6 +2914,11 @@ void common_speculative_draft(common_speculative * spec) { if (!result.empty() && (int) result.size() > dp.n_max) { SPC_DBG("truncating draft to %d tokens\n", dp.n_max); result.resize(dp.n_max); + + // the candidates are one per drafted token and must be cut with them + if (dp.result_q) { + dp.result_q->resize(dp.n_max); + } } } diff --git a/common/speculative.h b/common/speculative.h index d46b21eb71..0c9e0cf37e 100644 --- a/common/speculative.h +++ b/common/speculative.h @@ -69,6 +69,13 @@ struct common_speculative_draft_params { // the generated draft from the last _draft() call llama_tokens * result; + + // candidate distribution per drafted token; set it to make draft-simple and draft-mtp sample + std::vector> * result_q = nullptr; + + // the target's temp and seed, read only when the drafter samples probabilistically + float temp = 1.0f; + uint32_t seed = LLAMA_DEFAULT_SEED; }; common_speculative_draft_params & common_speculative_get_draft_params(common_speculative * spec, llama_seq_id seq_id); diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 615b035757..4da504e2ec 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -53,6 +53,31 @@ static common_speculative_output_limits server_output_limits(const common_params return result; } +// a checkpoint restore dropped tokens the target had accepted - re-accept them rather than verify again +static std::vector server_accept_replay( + common_sampler * smpl, + llama_context * ctx, + const std::vector & idxs, + const llama_tokens & draft) { + GGML_ASSERT(idxs.size() == draft.size() + 1); + + std::vector result; + result.reserve(idxs.size()); + + for (size_t i = 0; i < draft.size(); ++i) { + // the token is discarded - the call is what advances the sampler over this position + common_sampler_sample(smpl, ctx, idxs[i]); + common_sampler_accept(smpl, draft[i], true); + result.push_back(draft[i]); + } + + const llama_token id = common_sampler_sample(smpl, ctx, idxs[draft.size()]); + common_sampler_accept(smpl, id, true); + result.push_back(id); + + return result; +} + // synthetic draft verification for benchmarking - accept draft tokens at random instead of by match with the target // on replay the draft was already accepted before a context checkpoint restore, so repeat the same decisions static std::vector server_sample_and_accept_synth( @@ -212,6 +237,9 @@ struct server_slot { common_speculative * spec; llama_tokens spec_draft; + + // draft candidates per token in spec_draft; only draft-simple and draft-mtp fill it + std::vector> spec_draft_q; llama_tokens spec_prompt; std::vector spec_i_batch; common_prompt_checkpoint spec_ckpt; @@ -446,6 +474,11 @@ struct server_slot { return !!spec; } + // at temp 0 both p and q are point masses, so rejection is the same as sample-and-match + bool use_spec_rejection() const { + return task && task->params.sampling.temp > 0.0f; + } + void add_token(const completion_token_output & token) { if (!is_processing()) { SLT_WRN(*this, "%s", "slot is not processing\n"); @@ -3088,6 +3121,9 @@ private: if (n_draft_max > 0) { GGML_ASSERT(slot.can_speculate()); + // stale candidates: a replay never reads them, a new draft refills them + slot.spec_draft_q.clear(); + if (!slot.spec_draft.empty()) { // we have a previous (partial) draft to reuse if (use_ckpt_tgt) { @@ -3107,6 +3143,8 @@ private: slot.spec_prompt = slot.prompt.tokens.get_text_tokens(); + const bool spec_reject = slot.use_spec_rejection(); + common_speculative_get_draft_params(spec.get(), slot.id) = { /* .drafting = */ true, /* .n_max = */ n_draft_max, @@ -3114,6 +3152,9 @@ private: /* .id_last = */ slot.sampled, /* .prompt = */ &slot.spec_prompt, /* .result = */ &slot.spec_draft, + /* .result_q = */ spec_reject ? &slot.spec_draft_q : nullptr, + /* .temp = */ slot.task->params.sampling.temp, + /* .seed = */ slot.task->params.sampling.seed, }; drafting.push_back(&slot); @@ -4051,12 +4092,25 @@ private: common_sampler_ptr smpl_save(common_sampler_clone(slot.smpl.get())); GGML_ASSERT(slot.spec_i_batch.size() == n_draft + 1); + GGML_ASSERT(slot.spec_draft_q.empty() || (slot.spec_draft_q.size() == slot.spec_draft.size())); const auto & synth_probs = common_speculative_get_synth_probs(spec.get()); - auto accepted = synth_probs.empty() - ? common_sampler_sample_and_accept_n(slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft) - : server_sample_and_accept_synth( + + // drafters that fill no distribution fall back here + const bool use_rejection = slot.use_spec_rejection() && !slot.spec_draft_q.empty(); + + std::vector accepted; + if (!synth_probs.empty()) { + // synthetic acceptance replaces verification entirely, so it comes first + accepted = server_sample_and_accept_synth( slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft, synth_probs, slot.spec_synth_rng, slot.spec_is_replay); + } else if (slot.spec_is_replay && slot.use_spec_rejection()) { + accepted = server_accept_replay(slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft); + } else if (use_rejection) { + accepted = common_sampler_sample_and_accept_n_rejection(slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft, slot.spec_draft_q); + } else { + accepted = common_sampler_sample_and_accept_n(slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft); + } slot.spec_i_batch.clear(); GGML_ASSERT(accepted.size() >= 1);