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