Commit d0b490f25 for llama.cpp
commit d0b490f25edec147b4dcd69cce31c7d1099dff1d
Author: Hrishith Thadicherla <99313418+hthadicherla@users.noreply.github.com>
Date: Wed Oct 7 03:34:18 2026 -0700
sampling : use greedy selection for eligible temperature-zero chains (#29797)
* sampling : use greedy selection for eligible temperature-zero chains
Assisted-by: OpenAI Codex
* Apply suggestion from @ggerganov
Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
* sampling: simplify zero-temperature greedy eligibility
Allow the same greedy selection on CPU and grammar/reasoning-budget paths.
Keep distribution sampling for dynamic temperature and requested probabilities.
Cover the common sampler selection and probability behavior in the existing sampler tests.
Assisted-by: OpenAI Codex
* sampling : use greedy selection after final top-k with k=1
Assisted-by: OpenAI Codex
* cont : clean-up
---------
Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
diff --git a/common/sampling.cpp b/common/sampling.cpp
index d9c508049..6cc8872f2 100644
--- a/common/sampling.cpp
+++ b/common/sampling.cpp
@@ -399,8 +399,11 @@ struct common_sampler * common_sampler_init(
// only if user explicitly included adaptive-p sampler
samplers.push_back(llama_sampler_init_adaptive_p(params.adaptive_target, params.adaptive_decay, params.seed));
} else {
- // default: sample from distribution
- samplers.push_back(llama_sampler_init_dist(params.seed));
+ // Keep distribution sampling when callers request probabilities.
+ const bool greedy = params.n_probs == 0 && !params.samplers.empty() &&
+ ((params.samplers.back() == COMMON_SAMPLER_TYPE_TEMPERATURE && params.temp == 0.0f && params.dynatemp_range == 0.0f) ||
+ (params.samplers.back() == COMMON_SAMPLER_TYPE_TOP_K && params.top_k == 1));
+ samplers.push_back(greedy ? llama_sampler_init_greedy() : llama_sampler_init_dist(params.seed));
}
} else if (params.mirostat == 1) {
samplers.push_back(llama_sampler_init_temp(params.temp));
diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp
index 6c958a287..521c5718e 100644
--- a/src/llama-sampler.cpp
+++ b/src/llama-sampler.cpp
@@ -1086,6 +1086,12 @@ static void llama_sampler_greedy_backend_apply(
struct ggml_tensor * curl = ggml_argmax(ctx, logits);
ggml_set_name(curl, "greedy_argmax");
+ if (data->candidates != nullptr) {
+ struct ggml_tensor * candidates = ggml_reshape_2d(ctx, data->candidates, 1, ggml_nelements(data->candidates));
+ curl = ggml_get_rows(ctx, candidates, curl);
+ ggml_set_name(curl, "greedy_sampled_token");
+ }
+
data->sampled = curl;
}
diff --git a/tests/test-backend-sampler.cpp b/tests/test-backend-sampler.cpp
index 56736ac46..62c5ce40b 100644
--- a/tests/test-backend-sampler.cpp
+++ b/tests/test-backend-sampler.cpp
@@ -2,6 +2,7 @@
#include "llama.h"
#include "llama-cpp.h"
#include "common.h"
+#include "sampling.h"
#ifdef NDEBUG
#undef NDEBUG
@@ -316,7 +317,7 @@ static llama_sampler * test_single_output_backend_sampler_init(
return llama_sampler_init(&test_single_output_backend_sampler_i, ctx);
}
-static void test_backend_greedy_sampling(const test_params & params) {
+static void test_greedy(const test_params & params) {
const int seq_id = 0;
struct llama_sampler_chain_params backend_sampler_params = llama_sampler_chain_default_params();
@@ -351,7 +352,95 @@ static void test_backend_greedy_sampling(const test_params & params) {
}
}
-static void test_backend_top_k_sampling(const test_params & params) {
+
+static void test_greedy_filtered_common(const test_params & params) {
+ auto check = [&](common_params_sampling sp, bool greedy) {
+ common_sampler_ptr sampler(common_sampler_init(params.model.get(), sp));
+ auto * chain = common_sampler_get(sampler.get());
+ auto * last = llama_sampler_chain_get(chain, llama_sampler_chain_n(chain) - 1);
+ GGML_ASSERT(strcmp(llama_sampler_name(last), greedy ? "greedy" : "dist") == 0);
+
+ llama_token_data tokens[] = {{0, -2.0f, 0.0f}, {1, -1.0f, 0.0f}, {2, 0.0f, 0.0f}, {3, 1.0f, 0.0f}};
+ llama_token_data_array candidates = {tokens, 4, -1, false};
+ llama_sampler_apply(chain, &candidates);
+ GGML_ASSERT(candidates.selected >= 0);
+ if (greedy || sp.n_probs > 0) {
+ GGML_ASSERT(candidates.data[candidates.selected].id == 3);
+ }
+ if (sp.n_probs > 0) {
+ GGML_ASSERT(candidates.data[candidates.selected].p == 1.0f);
+ }
+ if (sp.dynatemp_range > 0.0f) {
+ int positive = 0;
+ for (size_t i = 0; i < candidates.size; ++i) {
+ positive += candidates.data[i].p > 0.0f;
+ }
+ GGML_ASSERT(positive > 1);
+ }
+ };
+
+ common_params_sampling sp;
+ sp.temp = 0.0f;
+ sp.samplers = {COMMON_SAMPLER_TYPE_TOP_K, COMMON_SAMPLER_TYPE_TEMPERATURE};
+ for (bool backend : {false, true}) {
+ sp.backend_sampling = backend;
+ check(sp, true);
+ auto probabilities = sp;
+ probabilities.n_probs = 4;
+ check(probabilities, false);
+ auto dynamic = sp;
+ dynamic.dynatemp_range = 1.0f;
+ check(dynamic, false);
+ }
+ auto grammar = sp;
+ grammar.grammar = {COMMON_GRAMMAR_TYPE_USER, "root ::= [a-z]+"};
+ check(grammar, true);
+ auto budget = sp;
+ budget.reasoning_budget_start = {1};
+ budget.reasoning_budget_end = {{2}};
+ budget.reasoning_budget_tokens = 1;
+ check(budget, true);
+ sp.samplers = {COMMON_SAMPLER_TYPE_TEMPERATURE, COMMON_SAMPLER_TYPE_TOP_K};
+ check(sp, false);
+ sp.samplers.clear();
+ check(sp, false);
+
+ sp.temp = 0.8f;
+ sp.samplers = {COMMON_SAMPLER_TYPE_TEMPERATURE, COMMON_SAMPLER_TYPE_TOP_K};
+ for (bool backend : {false, true}) {
+ sp.backend_sampling = backend;
+ sp.top_k = 1;
+ check(sp, true);
+ auto probabilities = sp;
+ probabilities.n_probs = 4;
+ check(probabilities, false);
+ sp.top_k = 8;
+ check(sp, false);
+ }
+}
+
+static void test_greedy_filtered(const test_params & params) {
+ for (int k : {1, 8}) {
+ llama_sampler_ptr chain(llama_sampler_chain_init(llama_sampler_chain_default_params()));
+ llama_sampler_chain_add(chain.get(), llama_sampler_init_top_k(k));
+ llama_sampler_chain_add(chain.get(), llama_sampler_init_greedy());
+ std::vector<llama_sampler_seq_config> configs = {{0, chain.get()}};
+ test_context test_ctx(params, configs);
+ GGML_ASSERT(test_ctx.decode({{0, "Write a Python function"}}));
+ for (int step = 0; step < 4; ++step) {
+ const int idx = test_ctx.idx_for_seq(0);
+ const auto * logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), idx);
+ const auto * ids = llama_get_sampled_candidates_ith(test_ctx.ctx.get(), idx);
+ GGML_ASSERT(llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), idx) == (uint32_t) k);
+ GGML_ASSERT(llama_get_sampled_candidates_count_ith(test_ctx.ctx.get(), idx) == (uint32_t) k);
+ const auto expected = ids[std::max_element(logits, logits + k) - logits];
+ GGML_ASSERT(llama_get_sampled_token_ith(test_ctx.ctx.get(), idx) == expected);
+ GGML_ASSERT(test_ctx.decode_token(expected));
+ }
+ }
+}
+
+static void test_top_k(const test_params & params) {
const int seq_id = 0;
const int32_t k = 8;
struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();
@@ -393,7 +482,7 @@ static void test_backend_top_k_sampling(const test_params & params) {
printf("backend top-k hybrid sampling test PASSED\n");
}
-static void test_backend_temp_sampling(const test_params & params) {
+static void test_temp(const test_params & params) {
{
const float temp_0 = 0.8f;
struct llama_sampler_chain_params backend_chain_params_0 = llama_sampler_chain_default_params();
@@ -481,7 +570,7 @@ static void test_backend_temp_sampling(const test_params & params) {
printf("backend temp sampling test PASSED\n");
}
-static void test_backend_temp_ext_sampling(const test_params & params) {
+static void test_temp_ext(const test_params & params) {
{
int seq_id = 0;
const float temp = 0.8f;
@@ -546,7 +635,7 @@ static void test_backend_temp_ext_sampling(const test_params & params) {
printf("backend temp_ext sampling test PASSED\n");
}
-static void test_backend_min_p_sampling(const test_params & params) {
+static void test_min_p(const test_params & params) {
const int seq_id = 0;
const float p = 0.1;
struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();
@@ -598,7 +687,7 @@ static void test_backend_min_p_sampling(const test_params & params) {
printf("min-p sampling test PASSED\n");
}
-static void test_backend_top_p_sampling(const test_params & params) {
+static void test_top_p(const test_params & params) {
const int seq_id = 0;
const float p = 0.9;
struct llama_sampler_chain_params backend_chain_params = llama_sampler_chain_default_params();
@@ -648,7 +737,7 @@ static void test_backend_top_p_sampling(const test_params & params) {
printf("top-p sampling test PASSED\n");
}
-static void test_backend_multi_sequence_sampling(const test_params & params) {
+static void test_multi_sequence(const test_params & params) {
struct llama_sampler_chain_params chain_params_0 = llama_sampler_chain_default_params();
llama_sampler_ptr sampler_chain_0(llama_sampler_chain_init(chain_params_0));
llama_sampler_chain_add(sampler_chain_0.get(), llama_sampler_init_greedy());
@@ -714,7 +803,7 @@ static void test_backend_multi_sequence_sampling(const test_params & params) {
printf("backend multi-sequence sampling test PASSED\n");
}
-static void test_backend_dist_sampling(const test_params & params) {
+static void test_dist(const test_params & params) {
const int seq_id = 0;
const int32_t seed = 88;
@@ -742,7 +831,7 @@ static void test_backend_dist_sampling(const test_params & params) {
printf("backend dist sampling test PASSED\n");
}
-static void test_backend_dist_sampling_and_cpu(const test_params & params) {
+static void test_dist_and_cpu(const test_params & params) {
const int seq_id = 0;
const int32_t seed = 88;
@@ -772,7 +861,7 @@ static void test_backend_dist_sampling_and_cpu(const test_params & params) {
printf("backend dist & cpu sampling test PASSED\n");
}
-static void test_backend_logit_bias_sampling(const test_params & params) {
+static void test_logit_bias(const test_params & params) {
const auto * model = params.model.get();
const auto * vocab = llama_model_get_vocab(model);
@@ -1093,7 +1182,7 @@ static void compare_penalties_logits(
GGML_ASSERT(stats.n_mismatch == 0);
}
-static void test_penalty_parameter_values(const test_params & params) {
+static void check_penalty_parameter_values(const test_params & params) {
struct penalty_test_case {
const char * name;
float repeat;
@@ -1292,7 +1381,7 @@ static void compare_masking_penalties_logits(
GGML_ASSERT(stats.n_mismatch == 0);
}
-static void test_backend_penalties_sampling(const test_params & params) {
+static void test_penalties(const test_params & params) {
printf("Testing backend penalties (repeat + freq + presence)\n");
compare_penalties_logits(params, 64, 1.1f, 0.5f, 0.25f, "Hello Hello world");
@@ -1371,14 +1460,14 @@ static void test_backend_penalties_sampling(const test_params & params) {
}, 64, 1.0f, 0.0f, 0.25f, "Hello", penalties_position::after_filter, true);
printf("Testing backend penalty parameter values\n");
- test_penalty_parameter_values(params);
+ check_penalty_parameter_values(params);
printf("backend penalties sampling test PASSED\n");
}
// This test verifies that it is possible to have two different backend samplers,
// one that uses the backend dist sampler, and another that uses CPU dist sampler.
-static void test_backend_mixed_sampling(const test_params & params) {
+static void test_mixed(const test_params & params) {
struct llama_sampler_chain_params chain_params_0 = llama_sampler_chain_default_params();
llama_sampler_ptr sampler_chain_0(llama_sampler_chain_init(chain_params_0));
llama_sampler_chain_add(sampler_chain_0.get(), llama_sampler_init_dist(88));
@@ -1428,7 +1517,7 @@ static void test_backend_mixed_sampling(const test_params & params) {
printf("backend mixed sampling test PASSED\n");
}
-static void test_backend_set_sampler(const test_params & params) {
+static void test_set_sampler(const test_params & params) {
const int seq_id = 0;
const int32_t seed = 88;
@@ -1493,7 +1582,7 @@ static void test_backend_set_sampler(const test_params & params) {
printf("backend set sampler test PASSED\n");
}
-static void test_backend_cpu_mixed_batch(const test_params & params) {
+static void test_cpu_mixed(const test_params & params) {
// Sequence 0 uses backend sampling
struct llama_sampler_chain_params chain_params_0 = llama_sampler_chain_default_params();
llama_sampler_ptr sampler_chain_0(llama_sampler_chain_init(chain_params_0));
@@ -1581,7 +1670,7 @@ static void test_backend_cpu_mixed_batch(const test_params & params) {
printf("backend-cpu mixed batch test PASSED\n");
}
-static void test_backend_multi_output_limit(const test_params & params) {
+static void test_multi_output_limit(const test_params & params) {
const llama_seq_id seq_id = 0;
llama_sampler_ptr chain(llama_sampler_chain_init(llama_sampler_chain_default_params()));
@@ -1594,15 +1683,15 @@ static void test_backend_multi_output_limit(const test_params & params) {
batch.add(llama_vocab_bos(test_ctx.vocab), i, seq_id, true);
}
- printf(">>> test_backend_multi_output_limit expected error start:\n");
+ printf(">>> test_multi_output_limit expected error start:\n");
const int ret = llama_process(test_ctx.ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get());
GGML_ASSERT(ret != 0 && "llama_decode should reject outputs above the per-sequence limit");
- printf("<<< test_backend_multi_output_limit expected error end.\n");
+ printf("<<< test_multi_output_limit expected error end.\n");
printf("backend multi-output limit test PASSED\n");
}
-static void test_backend_multi_sequence_multi_output_dist(const test_params & params) {
+static void test_multi_output_multi_sequence_dist(const test_params & params) {
const llama_vocab * vocab = llama_model_get_vocab(params.model.get());
const int32_t n_vocab = llama_vocab_n_tokens(vocab);
const uint32_t seeds[] = { 88, 1337 };
@@ -1697,7 +1786,7 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
printf("backend multi-sequence multi-output dist test PASSED\n");
}
-static void test_backend_multi_output_dist_transaction(const test_params & params) {
+static void test_multi_output_dist_transaction(const test_params & params) {
const llama_seq_id seq_id = 0;
const uint32_t seed = 95;
const llama_vocab * vocab = llama_model_get_vocab(params.model.get());
@@ -1762,7 +1851,7 @@ static void test_backend_multi_output_dist_transaction(const test_params & param
printf("backend multi-output dist transaction test PASSED\n");
}
-static void test_backend_multi_output_sampling_chain(const test_params & params) {
+static void test_multi_output_sampling_chain(const test_params & params) {
const llama_seq_id seq_id = 0;
const uint32_t seed = 88;
const float p = 0.9f;
@@ -1913,7 +2002,7 @@ static void test_backend_multi_output_sampling_chain(const test_params & params)
printf("backend multi-output sampling chain test PASSED\n");
}
-static void test_backend_multi_output_cpu_suffix(const test_params & params) {
+static void test_multi_output_cpu_suffix(const test_params & params) {
const llama_seq_id seq_id = 0;
const int32_t k = 8;
const llama_vocab * vocab = llama_model_get_vocab(params.model.get());
@@ -1977,28 +2066,35 @@ struct backend_test_case {
bool enabled_by_default;
};
+// note: test names are "test_<suffix>" and match the function implementing them
static const backend_test_case BACKEND_TESTS[] = {
- { "greedy", test_backend_greedy_sampling, true },
- { "logit_bias", test_backend_logit_bias_sampling, true },
- { "penalties", test_backend_penalties_sampling, true },
- { "temp", test_backend_temp_sampling, true },
- { "temp_ext", test_backend_temp_ext_sampling, true },
- { "top_k", test_backend_top_k_sampling, true },
- { "multi_sequence", test_backend_multi_sequence_sampling, true },
- { "dist", test_backend_dist_sampling, true },
- { "dist_and_cpu", test_backend_dist_sampling_and_cpu, true },
- { "set_sampler", test_backend_set_sampler, true },
- { "multi_output_limit", test_backend_multi_output_limit, true },
- { "multi_sequence_multi_output_dist", test_backend_multi_sequence_multi_output_dist, true },
- { "multi_output_dist_transaction", test_backend_multi_output_dist_transaction, true },
- { "multi_output_sampling_chain", test_backend_multi_output_sampling_chain, true },
- { "multi_output_cpu", test_backend_multi_output_cpu_suffix, true },
- { "mixed", test_backend_mixed_sampling, true },
- { "min_p", test_backend_min_p_sampling, true },
- { "cpu_mixed", test_backend_cpu_mixed_batch, true },
- { "top_p", test_backend_top_p_sampling, true },
+ // single sampler
+ { "test_greedy", test_greedy, true },
+ { "test_greedy_filtered", test_greedy_filtered, true },
+ { "test_greedy_filtered_common", test_greedy_filtered_common, true },
+ { "test_temp", test_temp, true },
+ { "test_temp_ext", test_temp_ext, true },
+ { "test_top_k", test_top_k, true },
+ { "test_top_p", test_top_p, true },
+ { "test_min_p", test_min_p, true },
+ { "test_logit_bias", test_logit_bias, true },
+ { "test_penalties", test_penalties, true },
+ // multiple sequences
+ { "test_multi_sequence", test_multi_sequence, true },
+ { "test_dist", test_dist, true },
+ { "test_dist_and_cpu", test_dist_and_cpu, true },
+ { "test_mixed", test_mixed, true },
+ { "test_cpu_mixed", test_cpu_mixed, true },
+ { "test_set_sampler", test_set_sampler, true },
+ // multiple outputs per sequence
+ { "test_multi_output_limit", test_multi_output_limit, true },
+ { "test_multi_output_multi_sequence_dist", test_multi_output_multi_sequence_dist, true },
+ { "test_multi_output_dist_transaction", test_multi_output_dist_transaction, true },
+ { "test_multi_output_sampling_chain", test_multi_output_sampling_chain, true },
+ { "test_multi_output_cpu_suffix", test_multi_output_cpu_suffix, true },
};
+// TODO: add usage, examples
static test_args parse_cli(int argc, char ** argv) {
test_args out;
@@ -2082,10 +2178,12 @@ static std::vector<const backend_test_case *> collect_tests_to_run(const std::st
}
#ifdef GGML_USE_HIP
// TODO: remove this when https://github.com/ggml-org/llama.cpp/pull/26592 is merged
- if (test.name == "penalties" || test.name == "set_sampler" ||
- test.name == "mixed" || test.name == "top_p" ||
- test.name == "multi_output_sampling_chain" ||
- test.name == "multi_output_cpu") {
+ if (test.name == "test_penalties" ||
+ test.name == "test_set_sampler" ||
+ test.name == "test_mixed" ||
+ test.name == "test_top_p" ||
+ test.name == "test_multi_output_sampling_chain" ||
+ test.name == "test_multi_output_cpu_suffix") {
fprintf(stderr, "Skipping test '%s' on HIP backend (no backend TOP_K support)\n", test.name.c_str());
continue;
}