mtp-relaxed-decoding/ holds the patch, build, and guide for a custom llama-server with relaxed-acceptance MTP speculative decoding (see its README for the full writeup). Bumps the open-webui submodule to the commit that adds it as a new service, connection, and the ollama-auth proxy swap. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
161 lines
8.4 KiB
Diff
161 lines
8.4 KiB
Diff
diff --git a/common/arg.cpp b/common/arg.cpp
|
|
index 86f8610..3fb6af4 100644
|
|
--- a/common/arg.cpp
|
|
+++ b/common/arg.cpp
|
|
@@ -4124,6 +4124,15 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
|
params.speculative.draft.p_min = std::stof(value);
|
|
}
|
|
).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_DRAFT_P_MIN"));
|
|
+ add_opt(common_arg(
|
|
+ {"--spec-draft-relaxed-top-n"}, "N",
|
|
+ string_format("accept a draft token if it ranks in the target's top-N candidates by "
|
|
+ "probability, instead of requiring an exact match (default: %d, disabled)",
|
|
+ params.speculative.draft.relaxed_top_n),
|
|
+ [](common_params & params, int value) {
|
|
+ params.speculative.draft.relaxed_top_n = value;
|
|
+ }
|
|
+ ).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_DRAFT_RELAXED_TOP_N"));
|
|
add_opt(common_arg(
|
|
{"--spec-draft-backend-sampling"},
|
|
{"--no-spec-draft-backend-sampling"},
|
|
diff --git a/common/common.h b/common/common.h
|
|
index de49dac..5c8a2c2 100644
|
|
--- a/common/common.h
|
|
+++ b/common/common.h
|
|
@@ -328,6 +328,11 @@ struct common_params_speculative_draft {
|
|
float p_split = 0.1f; // speculative decoding split probability
|
|
float p_min = 0.0f; // minimum speculative decoding probability (greedy)
|
|
|
|
+ // relaxed acceptance: accept a draft token if it ranks in the target's top-N
|
|
+ // candidates by probability, instead of requiring exact equality with the
|
|
+ // target's own sampled token. 0 = disabled (strict, distribution-exact behavior).
|
|
+ int32_t relaxed_top_n = 0;
|
|
+
|
|
bool backend_sampling = true; // offload draft sampling to the backend (default: on)
|
|
|
|
common_params_model mparams;
|
|
diff --git a/common/sampling.cpp b/common/sampling.cpp
|
|
index 06dea1e..0f2a26d 100644
|
|
--- a/common/sampling.cpp
|
|
+++ b/common/sampling.cpp
|
|
@@ -675,7 +675,34 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co
|
|
return id;
|
|
}
|
|
|
|
-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) {
|
|
+// Returns true if `token` ranks within the top `relaxed_top_n` candidates of `cur_p`
|
|
+// by probability. O(cur_p.size) and correct regardless of `cur_p.sorted`, since it
|
|
+// counts strictly-higher-probability entries rather than assuming sort order.
|
|
+static bool relaxed_top_n_contains(const llama_token_data_array & cur_p, llama_token token, int32_t relaxed_top_n) {
|
|
+ float token_p = -1.0f;
|
|
+ for (size_t k = 0; k < cur_p.size; k++) {
|
|
+ if (cur_p.data[k].id == token) {
|
|
+ token_p = cur_p.data[k].p;
|
|
+ break;
|
|
+ }
|
|
+ }
|
|
+ if (token_p < 0.0f) {
|
|
+ return false; // token isn't even in the candidate set (e.g. filtered out by top_k/top_p)
|
|
+ }
|
|
+
|
|
+ int32_t rank = 0;
|
|
+ for (size_t k = 0; k < cur_p.size; k++) {
|
|
+ if (cur_p.data[k].p > token_p) {
|
|
+ rank++;
|
|
+ if (rank >= relaxed_top_n) {
|
|
+ return false;
|
|
+ }
|
|
+ }
|
|
+ }
|
|
+ return true;
|
|
+}
|
|
+
|
|
+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, int32_t relaxed_top_n) {
|
|
GGML_ASSERT(idxs.size() == draft.size() + 1 && "idxs.size() must be draft.size() + 1");
|
|
|
|
std::vector<llama_token> result;
|
|
@@ -685,11 +712,30 @@ std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sample
|
|
for (; i < draft.size(); i++) {
|
|
const llama_token id = common_sampler_sample(gsmpl, ctx, idxs[i], grammar_first);
|
|
|
|
- common_sampler_accept(gsmpl, id, true);
|
|
+ const bool exact = (draft[i] == id);
|
|
+ bool accept = exact;
|
|
+ if (!accept && relaxed_top_n > 0) {
|
|
+ accept = relaxed_top_n_contains(gsmpl->cur_p, draft[i], relaxed_top_n);
|
|
+ if (getenv("RELAXED_VERIFY_LOG")) {
|
|
+ fprintf(stderr, "[RELAXED_VERIFY] pos=%zu relaxed_top_n=%d cur_p.size=%zu cur_p.sorted=%d draft='%s'(id=%d) target_sample='%s'(id=%d) -> %s\n",
|
|
+ i, relaxed_top_n, gsmpl->cur_p.size, (int) gsmpl->cur_p.sorted,
|
|
+ common_token_to_piece(ctx, draft[i]).c_str(), draft[i],
|
|
+ common_token_to_piece(ctx, id).c_str(), id,
|
|
+ accept ? "ACCEPT(relaxed)" : "reject");
|
|
+ }
|
|
+ }
|
|
|
|
- result.push_back(id);
|
|
+ // if the draft token is accepted under the relaxed rule, emit *that* token
|
|
+ // (not `id`) — server-context.cpp fed the whole drafted sequence into the
|
|
+ // target as one batch, so downstream KV state and idxs[i+1]'s logits already
|
|
+ // assume draft[i] was the real token at this position.
|
|
+ const llama_token emitted = accept ? draft[i] : id;
|
|
+
|
|
+ common_sampler_accept(gsmpl, emitted, true);
|
|
+
|
|
+ result.push_back(emitted);
|
|
|
|
- if (draft[i] != id) {
|
|
+ if (!accept) {
|
|
break;
|
|
}
|
|
}
|
|
@@ -705,13 +751,13 @@ std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sample
|
|
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<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const llama_tokens & draft, bool grammar_first, int32_t relaxed_top_n) {
|
|
std::vector<int> idxs(draft.size() + 1);
|
|
for (size_t i = 0; i < idxs.size(); ++i) {
|
|
idxs[i] = i;
|
|
}
|
|
|
|
- return common_sampler_sample_and_accept_n(gsmpl, ctx, idxs, draft, grammar_first);
|
|
+ return common_sampler_sample_and_accept_n(gsmpl, ctx, idxs, draft, grammar_first, relaxed_top_n);
|
|
}
|
|
|
|
uint32_t common_sampler_get_seed(const struct common_sampler * gsmpl) {
|
|
diff --git a/common/sampling.h b/common/sampling.h
|
|
index ced3c83..9627960 100644
|
|
--- a/common/sampling.h
|
|
+++ b/common/sampling.h
|
|
@@ -83,10 +83,13 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co
|
|
//
|
|
// returns at least 1 token, up to idxs.size()
|
|
//
|
|
-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);
|
|
+// relaxed_top_n > 0: accept a draft token whenever it ranks in the target's top-N
|
|
+// candidates by probability, instead of requiring exact equality with the target's
|
|
+// own sampled token. 0 keeps the original strict (distribution-exact) behavior.
|
|
+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, int32_t relaxed_top_n = 0);
|
|
|
|
// 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);
|
|
+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, int32_t relaxed_top_n = 0);
|
|
|
|
uint32_t common_sampler_get_seed(const struct common_sampler * gsmpl);
|
|
|
|
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
|
|
index a9edbd7..74e2f27 100644
|
|
--- a/tools/server/server-context.cpp
|
|
+++ b/tools/server/server-context.cpp
|
|
@@ -3795,7 +3795,8 @@ private:
|
|
common_sampler_ptr smpl_save(common_sampler_clone(slot.smpl.get()));
|
|
|
|
GGML_ASSERT(slot.spec_i_batch.size() == n_draft + 1);
|
|
- auto accepted = common_sampler_sample_and_accept_n(slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft);
|
|
+ auto accepted = common_sampler_sample_and_accept_n(slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft,
|
|
+ /* grammar_first */ false, params_base.speculative.draft.relaxed_top_n);
|
|
slot.spec_i_batch.clear();
|
|
|
|
GGML_ASSERT(accepted.size() >= 1);
|