From 6a4b205620d8de1c0270e32137cbac69588ff6e6 Mon Sep 17 00:00:00 2001 From: sb32445 Date: Sat, 3 Oct 2026 08:58:26 +0200 Subject: [PATCH 1/2] common : take the top-k straight from the logits when top-k starts the chain common_sampler_sample() built a token array for the whole vocabulary (set_logits) and then let the top-k sampler partial-sort it. When every sampler before top-k does nothing with the current settings (penalties, dry and top-n-sigma off, no logit bias, no mirostat, no backend sampling), no grammar applied first and the reasoning budget not forcing, select the k largest logits directly. The selection repeats the heap steps of libstdc++ std::partial_sort, so entries and the order of equal logits are the same as before. Checked against std::partial_sort on 120000 random arrays (ties, NaN, inf) and on server outputs (sampling, tool calls, penalties, top_k 1/200, logprobs): identical. LLAMA_SAMPLER_FAST_TOPK=0 turns it off. Decode with MTP n-max 2: +1.9 % (greedy benchmark), +1.85 % with the Hermes sampling settings (thinking, budget 16384, top_k 20). Co-Authored-By: Claude Sonnet 5.5 --- common/sampling.cpp | 146 +++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 145 insertions(+), 1 deletion(-) diff --git a/common/sampling.cpp b/common/sampling.cpp index dd5b14823c76..0d5bd281c93d 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -109,6 +109,76 @@ struct ring_buffer { std::vector data; }; +// Top k of the logits without building the full token array. It repeats the steps of libstdc++ std::partial_sort +// (make_heap on the first k, replace the heap top by every larger logit in index order, sort_heap) so that tokens with +// equal logits end up in the same order as with llama_sampler_top_k. +static void topk_from_logits(const float * logits, int32_t n_vocab, int32_t k, std::vector & out) { + static const auto comp = [](const llama_token_data & a, const llama_token_data & b) { + return a.logit > b.logit; + }; + + out.resize(k); + for (int32_t i = 0; i < k; i++) { + out[i] = llama_token_data{ i, logits[i], 0.0f }; + } + std::make_heap(out.begin(), out.end(), comp); + + llama_token_data * heap = out.data(); + + // std::__pop_heap(first, middle, result) with value = the new token: sift down to a leaf, then sift up + const auto replace_top = [&](const llama_token_data & value) { + ptrdiff_t hole = 0; + ptrdiff_t child = 0; + const ptrdiff_t len = k; + while (child < (len - 1) / 2) { + child = 2 * (child + 1); + if (comp(heap[child], heap[child - 1])) { + child--; + } + heap[hole] = heap[child]; + hole = child; + } + if ((len & 1) == 0 && child == (len - 2) / 2) { + child = 2 * (child + 1); + heap[hole] = heap[child - 1]; + hole = child - 1; + } + ptrdiff_t parent = (hole - 1) / 2; + while (hole > 0 && comp(heap[parent], value)) { + heap[hole] = heap[parent]; + hole = parent; + parent = (hole - 1) / 2; + } + heap[hole] = value; + }; + + constexpr int32_t BLOCK = 16; + float thr = heap[0].logit; + int32_t i = k; + for (; i + BLOCK <= n_vocab; i += BLOCK) { + bool any = false; + for (int32_t j = 0; j < BLOCK; j++) { + any |= logits[i + j] > thr; + } + if (!any) { + continue; + } + for (int32_t j = 0; j < BLOCK; j++) { + if (logits[i + j] > heap[0].logit) { + replace_top(llama_token_data{ i + j, logits[i + j], 0.0f }); + } + } + thr = heap[0].logit; + } + for (; i < n_vocab; i++) { + if (logits[i] > heap[0].logit) { + replace_top(llama_token_data{ i, logits[i], 0.0f }); + } + } + + std::sort_heap(out.begin(), out.end(), comp); +} + struct common_sampler { common_params_sampling params; @@ -122,6 +192,9 @@ struct common_sampler { llama_token_data_array cur_p; + // > 0: the chain starts with top-k (all earlier samplers are no-ops), so cur_p can be built from the top k logits only + int32_t fast_topk = 0; + void reset() { prev.clear(); @@ -162,6 +235,27 @@ struct common_sampler { cur_p = { cur.data(), cur.size(), -1, false }; } + // cur_p = the k largest logits in descending order, same entries and order as top-k on the full array (set_logits + std::partial_sort) + // returns false if the logits are not plain CPU logits (backend sampling) + bool set_logits_topk(struct llama_context * ctx, int idx, int32_t k) { + if (llama_get_sampled_probs_ith(ctx, idx) || llama_get_sampled_logits_ith(ctx, idx)) { + return false; + } + + const float * logits = llama_get_logits_ith(ctx, idx); + GGML_ASSERT(logits != nullptr); + + const int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx))); + if (n_vocab <= k) { + return false; + } + + topk_from_logits(logits, n_vocab, k, cur); + + cur_p = { cur.data(), cur.size(), -1, true }; + return true; + } + common_time_meas tm() { return common_time_meas(t_total_us, params.no_perf); } @@ -185,6 +279,41 @@ std::string common_params_sampling::print() const { return std::string(result); } +// k > 0 if the sampler chain starts with top-k (every sampler before it does nothing with these settings); LLAMA_SAMPLER_FAST_TOPK=0 turns it off +static int32_t common_sampler_fast_topk(const common_params_sampling & params, bool has_logit_bias) { + static const bool enabled = [] { const char * e = getenv("LLAMA_SAMPLER_FAST_TOPK"); return !e || atoi(e) != 0; }(); + + if (!enabled || has_logit_bias || params.mirostat != 0 || params.backend_sampling || params.top_k <= 0 || params.top_k > 128) { + return 0; + } + + for (const auto cnstr : params.samplers) { + switch (cnstr) { + case COMMON_SAMPLER_TYPE_TOP_K: + return params.top_k; + case COMMON_SAMPLER_TYPE_PENALTIES: + if (params.penalty_last_n != 0 && (params.penalty_repeat != 1.0f || params.penalty_freq != 0.0f || params.penalty_present != 0.0f)) { + return 0; + } + break; + case COMMON_SAMPLER_TYPE_DRY: + if (params.dry_multiplier != 0.0f && params.dry_base >= 1.0f && params.dry_penalty_last_n != 0) { + return 0; + } + break; + case COMMON_SAMPLER_TYPE_TOP_N_SIGMA: + if (params.top_n_sigma > 0.0f) { + return 0; + } + break; + default: + return 0; + } + } + + return 0; +} + struct common_sampler * common_sampler_init( const struct llama_model * model, struct common_params_sampling & params) { @@ -323,6 +452,8 @@ struct common_sampler * common_sampler_init( } } + bool has_logit_bias = false; + // logit bias: user biases + model suppress tokens (-INFINITY) { std::vector merged = params.logit_bias; @@ -334,6 +465,7 @@ struct common_sampler * common_sampler_init( } if (!merged.empty()) { + has_logit_bias = true; samplers.push_back(llama_sampler_init_logit_bias(llama_vocab_n_tokens(vocab), merged.size(), merged.data())); } } @@ -435,6 +567,8 @@ struct common_sampler * common_sampler_init( /* .cur_p = */ {}, }; + result->fast_topk = common_sampler_fast_topk(result->params, has_logit_bias); + return result; } @@ -516,6 +650,7 @@ struct common_sampler * common_sampler_clone(common_sampler * gsmpl) { /* .prev = */ gsmpl->prev, /* .cur = */ gsmpl->cur, /* .cur_p = */ gsmpl->cur_p, + /* .fast_topk = */ gsmpl->fast_topk, }; } @@ -536,6 +671,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->fast_topk = src->fast_topk; dst->t_total_us = src->t_total_us; } @@ -605,7 +741,15 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co auto & chain = gsmpl->chain; auto & cur_p = gsmpl->cur_p; // initialized by set_logits - gsmpl->set_logits(ctx, idx); + // top-k directly from the logits when nothing before the top-k sampler reads the full array + const bool fast = gsmpl->fast_topk > 0 && + !(grammar_first && grammar_should_apply(gsmpl)) && + !(rbudget && common_reasoning_budget_get_state(rbudget) == REASONING_BUDGET_FORCING) && + gsmpl->set_logits_topk(ctx, idx, gsmpl->fast_topk); + + if (!fast) { + gsmpl->set_logits(ctx, idx); + } // Check if a backend sampler has already sampled a token in which case we // return that token id directly. From a51ea62e79cae2e4b0c5a95278187d834c7b2f6e Mon Sep 17 00:00:00 2001 From: sb32445 Date: Sat, 3 Oct 2026 14:58:06 +0200 Subject: [PATCH 2/2] common : remove the LLAMA_SAMPLER_FAST_TOPK switch The environment switch of the previous commit was only there to measure the change; the direct top-k selection is now used whenever its conditions hold. Co-Authored-By: Claude Sonnet 5.5 --- common/sampling.cpp | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/common/sampling.cpp b/common/sampling.cpp index 0d5bd281c93d..853e9d068554 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -279,11 +279,9 @@ std::string common_params_sampling::print() const { return std::string(result); } -// k > 0 if the sampler chain starts with top-k (every sampler before it does nothing with these settings); LLAMA_SAMPLER_FAST_TOPK=0 turns it off +// k > 0 if the sampler chain starts with top-k (every sampler before it does nothing with these settings) static int32_t common_sampler_fast_topk(const common_params_sampling & params, bool has_logit_bias) { - static const bool enabled = [] { const char * e = getenv("LLAMA_SAMPLER_FAST_TOPK"); return !e || atoi(e) != 0; }(); - - if (!enabled || has_logit_bias || params.mirostat != 0 || params.backend_sampling || params.top_k <= 0 || params.top_k > 128) { + if (has_logit_bias || params.mirostat != 0 || params.backend_sampling || params.top_k <= 0 || params.top_k > 128) { return 0; }