Files
bhethermanandClaude Sonnet 5 627904e130 Add relaxed-MTP llama-server build and wire it into the stack
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>
2026-09-04 11:33:31 -04:00

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);