mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-03 03:17:32 -05:00
spec : add probabilistic sampling for simple draft and MTP (#27694)
* Make the drafter probabilistic and the target verify by rejection sampling * Drop stale spec_draft_q before drafting * Fallback to argmax sampling for grammar-constrained requests and adding flag for enabling probabilistic draft sampling. Default flag value is greedy. * Support grammar-constrained requests in rejection sampling * Fix - renormalize distribution after masking * copy rng on sampler copy and re-accept drafted tokens on replay * Fix draft sampler sharing the target's rng stream * Simplify the rejection sampler's inputs and move replay to the server * Truncate the draft candidates along with the draft --------- Co-authored-by: praneshgo <227579474+praneshgo@users.noreply.github.com> Co-authored-by: Pranesh Gonegandla <pgonegandla@nvidia.com>
This commit is contained in:
co-authored by
praneshgo
Pranesh Gonegandla
parent
134b2bb756
commit
1fb7ef3e33
@@ -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"}, "<dev1,dev2,..>",
|
||||
"comma-separated list of devices to use for offloading the draft model (none = don't offload, default: follows --device)\n"
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -12,6 +12,7 @@
|
||||
#include <climits>
|
||||
#include <cmath>
|
||||
#include <cstring>
|
||||
#include <random>
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
|
||||
@@ -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<llama_token>(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<llama_token> 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<llama_token> common_sampler_sample_and_accept_n_rejection(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector<int> & idxs, const llama_tokens & draft, const std::vector<std::vector<llama_token_data>> & 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<llama_token> 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<float> uni(0.0f, 1.0f);
|
||||
|
||||
std::vector<llama_token_data> residual;
|
||||
|
||||
std::vector<llama_token_data> 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<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const llama_tokens & draft, bool grammar_first) {
|
||||
std::vector<int> idxs(draft.size() + 1);
|
||||
for (size_t i = 0; i < idxs.size(); ++i) {
|
||||
|
||||
@@ -85,6 +85,9 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co
|
||||
//
|
||||
std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector<int> & 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<llama_token> common_sampler_sample_and_accept_n_rejection(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector<int> & idxs, const llama_tokens & draft, const std::vector<std::vector<llama_token_data>> & draft_q, bool grammar_first = false);
|
||||
|
||||
// assume idxs == [ 0, 1, 2, ..., draft.size() ]
|
||||
std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const llama_tokens & draft, bool grammar_first = false);
|
||||
|
||||
|
||||
+94
-8
@@ -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<common_sampler_ptr> & smpls,
|
||||
std::vector<common_params_sampling> & 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<std::string, common_speculative_type> 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<common_sampler_ptr> smpls;
|
||||
|
||||
std::vector<common_params_sampling> 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<common_sampler_ptr> smpls;
|
||||
|
||||
std::vector<common_params_sampling> smpls_cfg;
|
||||
|
||||
// backend sampler chain per seq, attached to ctx_dft
|
||||
std::vector<llama_sampler *> 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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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<std::vector<llama_token_data>> * 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);
|
||||
|
||||
@@ -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<llama_token> server_accept_replay(
|
||||
common_sampler * smpl,
|
||||
llama_context * ctx,
|
||||
const std::vector<int32_t> & idxs,
|
||||
const llama_tokens & draft) {
|
||||
GGML_ASSERT(idxs.size() == draft.size() + 1);
|
||||
|
||||
std::vector<llama_token> 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<llama_token> 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<std::vector<llama_token_data>> spec_draft_q;
|
||||
llama_tokens spec_prompt;
|
||||
std::vector<int32_t> 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<llama_token> 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);
|
||||
|
||||
Reference in New Issue
Block a user