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 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) { +// 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 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, int32_t relaxed_top_n) { GGML_ASSERT(idxs.size() == draft.size() + 1 && "idxs.size() must be draft.size() + 1"); std::vector result; @@ -685,11 +712,30 @@ std::vector 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 common_sampler_sample_and_accept_n(struct common_sample 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 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 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 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); +// 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 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, int32_t relaxed_top_n = 0); // 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); +std::vector 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);