]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
CUDA: Add backend sampler for penalties sampler (#25262)
authorKonrad Moren <redacted>
Mon, 3 Aug 2026 12:26:09 +0000 (14:26 +0200)
committerGitHub <redacted>
Mon, 3 Aug 2026 12:26:09 +0000 (14:26 +0200)
* sampling: enhance penalty handling in common_sampler_init

- Set default value for penalty_last_n based on model context if not specified.
- Ensure penalty_last_n and n_prev are non-negative.
- Update llama_sampler_penalties structure to inherit from llama_sampler_backend and add backend input handling for penalties.
- Implement backend initialization and application logic for penalties, including frequency and presence adjustments.

* tests: add backend penalties sampling tests and utility functions

- Introduced `accept_prompt` and `unique_prompt_tokens` functions to handle prompt acceptance and token uniqueness.
- Implemented `compare_penalties_logits` to compare logits from backend and CPU samplers with penalties.
- Added `test_backend_penalties_sampling` to validate backend penalties with various configurations.
- Enhanced the test suite for better coverage of penalty handling in sampling.

* sampling: add support for top-k penalties in backend sampling

* sampling: add fix to ensure  stable numerical results. Preserve masked logits as -Inf and no longer generate NaN.

* sampling: enhance penalty comparison tests with masking penalties logic

* add comments on padding

* sampling: add comments on modifications

* add the unit test to cover masked-out token as -INF

* validate repeat penalty to ensure it is finite and greater than 0; add tests for invalid values

* refactor: test functions to share logic and be less verbose

* add test to cover case where previously penalized token is not part of candidates

* remove comments

* remove redundant penalty_last_n initialization and validation in common_sampler_init

* add support for penalties in sampler chain with configurable positions

* add validation for penalty parameters and enhance tests for non-finite values

* add context parameter to common_sampler_init and set default for penalty_last_n

* add llama_n_ctx parameter to common_sampler_init for improved sampler initialization

* replace penalty_last_n x n_candidates comparison matrix with a vocabulary-sized count tensor

* add tests for backend penalties sampling without filler entries , token_count.size() == n_active == n_max == 64

* add test for backend penalties sampling  after top-p with large history window

* remove as unused

* add is_disabled method, tensor logits reshape, add rest review suggestions

* clarify comment

common/arg.cpp
common/common.cpp
common/sampling.cpp
common/sampling.h
include/llama.h
src/llama-graph.cpp
src/llama-sampler.cpp
tests/test-arg-parser.cpp
tests/test-backend-sampler.cpp
tools/server/server-context.cpp

index 772422f6820a2c897027a359ce234be34e6465f1..305938fcb2b4efc3a7d0c7719d1953e8b32192fb 100644 (file)
@@ -27,6 +27,7 @@
 #include <algorithm>
 #include <cinttypes>
 #include <climits>
+#include <cmath>
 #include <cstdarg>
 #include <filesystem>
 #include <fstream>
@@ -2036,7 +2037,13 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
         {"--repeat-penalty"}, "N",
         string_format("penalize repeat sequence of tokens (default: %.2f, 1.0 = disabled)", (double)params.sampling.penalty_repeat),
         [](common_params & params, const std::string & value) {
-            params.sampling.penalty_repeat = std::stof(value);
+            const float penalty_repeat = std::stof(value);
+            if (!std::isfinite(penalty_repeat) ||
+                penalty_repeat <= 0.0f ||
+                !std::isfinite(1.0f/penalty_repeat)) {
+                throw std::runtime_error("error: repeat-penalty must be finite and greater than 0\n");
+            }
+            params.sampling.penalty_repeat = penalty_repeat;
             params.sampling.user_sampling_config |= common_params_sampling_config::COMMON_PARAMS_SAMPLING_CONFIG_PENALTY_REPEAT;
         }
     ).set_sampling());
@@ -2044,14 +2051,22 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
         {"--presence-penalty"}, "N",
         string_format("repeat alpha presence penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_present),
         [](common_params & params, const std::string & value) {
-            params.sampling.penalty_present = std::stof(value);
+            const float penalty_present = std::stof(value);
+            if (!std::isfinite(penalty_present)) {
+                throw std::runtime_error("error: presence-penalty must be finite\n");
+            }
+            params.sampling.penalty_present = penalty_present;
         }
     ).set_sampling());
     add_opt(common_arg(
         {"--frequency-penalty"}, "N",
         string_format("repeat alpha frequency penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_freq),
         [](common_params & params, const std::string & value) {
-            params.sampling.penalty_freq = std::stof(value);
+            const float penalty_freq = std::stof(value);
+            if (!std::isfinite(penalty_freq)) {
+                throw std::runtime_error("error: frequency-penalty must be finite\n");
+            }
+            params.sampling.penalty_freq = penalty_freq;
         }
     ).set_sampling());
     add_opt(common_arg(
index ff27d392fb2e03f9491e5ce57cb1a93eb7d6be3b..c941fd505ab3eb21b1847ab34b12b8f81fad47e8 100644 (file)
@@ -1299,8 +1299,9 @@ common_init_result::common_init_result(common_params & params, bool model_only)
     pimpl->samplers.resize(cparams.n_seq_max);
     pimpl->samplers_seq_config.resize(cparams.n_seq_max);
 
+    const int32_t n_ctx = cparams.n_ctx > 0 ? (int32_t) cparams.n_ctx : llama_model_n_ctx_train(model);
     for (int i = 0; i < (int) cparams.n_seq_max; ++i) {
-        pimpl->samplers[i].reset(common_sampler_init(model, params.sampling));
+        pimpl->samplers[i].reset(common_sampler_init(model, params.sampling, n_ctx));
         pimpl->samplers_seq_config[i] = { i, common_sampler_get(pimpl->samplers[i].get()) };
     }
 
index 256ac161e20f14f28ccde26967380eba28438c86..5698c0263b24d46d6eb9b16380f37e4e3288bd69 100644 (file)
@@ -184,9 +184,26 @@ std::string common_params_sampling::print() const {
     return std::string(result);
 }
 
-struct common_sampler * common_sampler_init(const struct llama_model * model, struct common_params_sampling & params) {
-    const llama_vocab * vocab = llama_model_get_vocab(model);
+struct common_sampler * common_sampler_init(
+        const struct llama_model * model,
+        struct common_params_sampling & params,
+        int32_t n_ctx) {
+    if (!std::isfinite(params.penalty_repeat) ||
+        params.penalty_repeat <= 0.0f ||
+        !std::isfinite(1.0f/params.penalty_repeat)) {
+        throw std::invalid_argument("penalty_repeat must be finite and greater than 0");
+    }
+    if (!std::isfinite(params.penalty_freq)) {
+        throw std::invalid_argument("penalty_freq must be finite");
+    }
+    if (!std::isfinite(params.penalty_present)) {
+        throw std::invalid_argument("penalty_present must be finite");
+    }
+    if (params.penalty_last_n == -1) {
+        params.penalty_last_n = n_ctx > 0 ? n_ctx : llama_model_n_ctx_train(model);
+    }
 
+    const llama_vocab * vocab = llama_model_get_vocab(model);
     llama_sampler_chain_params lparams = llama_sampler_chain_default_params();
 
     lparams.no_perf = params.no_perf;
index 4191988bb87796d93c8e68470d042d30a968ae9e..91e2cea787fb169a1fbf501613919f5e8fbe60e2 100644 (file)
@@ -37,7 +37,10 @@ struct common_sampler;
 // llama_sampler API overloads
 
 // note: can mutate params in some cases
-struct common_sampler * common_sampler_init(const struct llama_model * model, struct common_params_sampling & params);
+struct common_sampler * common_sampler_init(
+        const struct llama_model * model,
+        struct common_params_sampling & params,
+        int32_t n_ctx = 0);
 
 void common_sampler_free(struct common_sampler * gsmpl);
 
index 6e53e2297235ee0e4c4e3589cb85e90727752e8b..f2d7e38858378b2843833d21d31bae310c0adddf 100644 (file)
@@ -1256,6 +1256,7 @@ extern "C" {
         struct ggml_tensor * probs;
         struct ggml_tensor * sampled;
         struct ggml_tensor * candidates;
+        int64_t              n_vocab;
     };
 
     // user code can implement the interface below in order to create custom llama_sampler
@@ -1425,9 +1426,9 @@ extern "C" {
     /// NOTE: Avoid using on the full vocabulary as searching for repeated tokens can become slow. For example, apply top-k or top-p sampling first.
     LLAMA_API struct llama_sampler * llama_sampler_init_penalties(
                              int32_t   penalty_last_n,   // last n tokens to penalize (0 = disable penalty, -1 = context size)
-                               float   penalty_repeat,   // 1.0 = disabled
-                               float   penalty_freq,     // 0.0 = disabled
-                               float   penalty_present); // 0.0 = disabled
+                               float   penalty_repeat,   // must be > 0.0, 1.0 = disabled
+                               float   penalty_freq,     // must be finite, 0.0 = disabled
+                               float   penalty_present); // must be finite, 0.0 = disabled
 
     ///  @details DRY sampler, designed by p-e-w, as described in: https://github.com/oobabooga/text-generation-webui/pull/5677, porting Koboldcpp implementation authored by pi6am: https://github.com/LostRuins/koboldcpp/pull/982
     LLAMA_API struct llama_sampler * llama_sampler_init_dry(
index e12a8cdc2aacf4d1b744f57571d7799bd460fd15..1a35692300c1d28f64b488d892252311fd7707fb 100644 (file)
@@ -3620,6 +3620,7 @@ void llm_graph_context::build_sampling() const {
             /*.probs       =*/ nullptr,
             /*.sampled     =*/ nullptr,
             /*.candidates  =*/ nullptr,
+            /*.n_vocab     =*/ logits_seq->ne[0],
         };
 
         assert(sampler->iface->backend_apply);
index a9cb6bee5fd78728e5c94d5d1d008c3022abf330..b2f1abe7378512faaeb47a26718573ba03bd96a0 100644 (file)
@@ -589,6 +589,7 @@ static bool llama_sampler_backend_support(
         /*.probs      = */ nullptr,
         /*.sampled    = */ nullptr,
         /*.candidates = */ ggml_new_tensor_1d(ctx, GGML_TYPE_I32, n),
+        /*.n_vocab    = */ n,
     };
 
     ggml_cgraph * gf = ggml_new_graph(ctx);
@@ -2638,7 +2639,7 @@ struct llama_sampler * llama_sampler_init_grammar_lazy_patterns(
 
 // penalties
 
-struct llama_sampler_penalties {
+struct llama_sampler_penalties : public llama_sampler_backend {
     const int32_t penalty_last_n;
     const float   penalty_repeat;
     const float   penalty_freq;
@@ -2648,10 +2649,49 @@ struct llama_sampler_penalties {
 
     // a frequency map to count token occurrences
     std::unordered_map<llama_token, int> token_count;
+
+    // backend graph inputs
+    ggml_tensor * inp_token_ids = nullptr;
+    ggml_tensor * inp_counts    = nullptr;
+
+    // backend helpers
+    int32_t n_vocab = 0;
+    int32_t n_max   = 0;
+    bool has_candidates = false;
+
+    std::vector<int32_t> host_token_ids;
+    std::vector<int32_t> host_counts;
+
+    static bool is_disabled(
+            int32_t penalty_last_n,
+            float   penalty_repeat,
+            float   penalty_freq,
+            float   penalty_present) {
+        return penalty_last_n == 0 ||
+            (penalty_repeat == 1.0f && penalty_freq == 0.0f && penalty_present == 0.0f);
+    }
+
+    bool is_disabled() const {
+        return is_disabled(penalty_last_n, penalty_repeat, penalty_freq, penalty_present);
+    }
+
+    llama_sampler_penalties(
+            int32_t penalty_last_n,
+            float   penalty_repeat,
+            float   penalty_freq,
+            float   penalty_present)
+        : llama_sampler_backend("penalties")
+        , penalty_last_n  (penalty_last_n)
+        , penalty_repeat  (penalty_repeat)
+        , penalty_freq    (penalty_freq)
+        , penalty_present (penalty_present)
+        , prev            (penalty_last_n) {
+    }
 };
 
-static const char * llama_sampler_penalties_name(const struct llama_sampler * /*smpl*/) {
-    return "penalties";
+static const char * llama_sampler_penalties_name(const struct llama_sampler * smpl) {
+    auto * ctx = (llama_sampler_penalties *) smpl->ctx;
+    return ctx->get_name();
 }
 
 static void llama_sampler_penalties_accept(struct llama_sampler * smpl, llama_token token) {
@@ -2688,8 +2728,7 @@ static void llama_sampler_penalties_accept(struct llama_sampler * smpl, llama_to
 static void llama_sampler_penalties_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {
     auto * ctx = (llama_sampler_penalties *) smpl->ctx;
 
-    if ((ctx->penalty_last_n == 0) ||
-        (ctx->penalty_repeat == 1.0f && ctx->penalty_freq == 0.0f && ctx->penalty_present == 0.0f)) {
+    if (ctx->is_disabled()) {
         return;
     }
 
@@ -2736,7 +2775,8 @@ static struct llama_sampler * llama_sampler_penalties_clone(const struct llama_s
     {
         auto * result_ctx = (llama_sampler_penalties *) result->ctx;
 
-        result_ctx->prev = ctx->prev;
+        result_ctx->prev        = ctx->prev;
+        result_ctx->token_count = ctx->token_count;
     }
 
     return result;
@@ -2746,6 +2786,171 @@ static void llama_sampler_penalties_free(struct llama_sampler * smpl) {
     delete (llama_sampler_penalties *) smpl->ctx;
 }
 
+static bool llama_sampler_penalties_backend_init(
+        struct llama_sampler       * smpl,
+        ggml_backend_buffer_type_t   buft) {
+    auto * sctx = (llama_sampler_penalties *) smpl->ctx;
+
+    const bool res = llama_sampler_backend_support(smpl, buft);
+
+    sctx->init(res);
+
+    return res;
+}
+
+static void llama_sampler_penalties_backend_apply(
+        struct llama_sampler      * smpl,
+        struct ggml_context       * ctx,
+        struct ggml_cgraph        * gf,
+        struct llama_sampler_data * data) {
+    GGML_UNUSED(gf);
+
+    auto * sctx = (llama_sampler_penalties *) smpl->ctx;
+
+    if (sctx->is_disabled()) {
+        return;
+    }
+
+    GGML_ASSERT(data->n_vocab > 0 && data->n_vocab <= INT32_MAX);
+
+    sctx->has_candidates = data->candidates != nullptr;
+    sctx->n_vocab = (int32_t) data->n_vocab;
+    sctx->n_max   = std::min(sctx->penalty_last_n, sctx->n_vocab);
+
+    sctx->inp_token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, sctx->n_max);
+    ggml_set_name(sctx->inp_token_ids, "penalties_token_ids");
+    ggml_set_input(sctx->inp_token_ids);
+
+    sctx->inp_counts = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, sctx->n_max);
+    ggml_set_name(sctx->inp_counts, "penalties_counts");
+    ggml_set_input(sctx->inp_counts);
+
+    if ((int32_t) sctx->host_token_ids.size() != sctx->n_max) {
+        sctx->host_token_ids.assign(sctx->n_max, 0);
+        sctx->host_counts.assign(sctx->n_max, 0);
+    }
+
+    // flatten
+    ggml_tensor * logits = ggml_reshape_1d(ctx, data->logits, ggml_nelements(data->logits));
+    ggml_tensor * gathered = logits;
+    ggml_tensor * counts_f32 = ggml_cast(ctx, sctx->inp_counts, GGML_TYPE_F32);
+
+    if (sctx->has_candidates) {
+        ggml_tensor * candidates = ggml_reshape_1d(
+                ctx, data->candidates, ggml_nelements(data->candidates));
+        const int64_t n_candidates = candidates->ne[0];
+        GGML_ASSERT(n_candidates == ggml_nelements(logits));
+
+        ggml_tensor * counts_rows = ggml_fill(
+                ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 1, sctx->n_vocab), 0.0f);
+        ggml_tensor * scatter_rows = ggml_reshape_2d(ctx, counts_f32, 1, sctx->n_max);
+        counts_rows = ggml_set_rows(ctx, counts_rows, scatter_rows, sctx->inp_token_ids);
+        counts_f32 = ggml_get_rows(ctx, counts_rows, candidates);
+        counts_f32 = ggml_reshape_1d(ctx, counts_f32, n_candidates);
+    } else {
+        ggml_tensor * logits_rows = ggml_reshape_2d(ctx, logits, 1, ggml_nelements(logits));
+        gathered = ggml_get_rows(ctx, logits_rows, sctx->inp_token_ids);
+        gathered = ggml_reshape_1d(ctx, gathered, sctx->n_max);
+    }
+
+    ggml_tensor * active_mask = ggml_step(ctx, counts_f32);
+    ggml_tensor * inactive_mask = ggml_sub(ctx, ggml_fill(ctx, active_mask, 1.0f), active_mask);
+
+    ggml_tensor * penalized = gathered;
+
+    if (sctx->penalty_repeat != 1.0f) {
+        ggml_tensor * pos_mask = ggml_step(ctx, penalized);
+        ggml_tensor * neg_mask = ggml_sub(ctx, ggml_fill(ctx, pos_mask, 1.0f), pos_mask);
+
+        ggml_tensor * pos_scale = ggml_scale(ctx, pos_mask, 1.0f/sctx->penalty_repeat);
+        ggml_tensor * neg_scale = ggml_scale(ctx, neg_mask, sctx->penalty_repeat);
+        ggml_tensor * repeat_scale = ggml_add(ctx, pos_scale, neg_scale);
+
+        // scale inactive entries with 1 to avoid -INF * 0 = NaN for values masked by top-p
+        repeat_scale = ggml_mul(ctx, repeat_scale, active_mask);
+        repeat_scale = ggml_add(ctx, repeat_scale, inactive_mask);
+        penalized = ggml_mul(ctx, gathered, repeat_scale);
+    }
+
+    if (sctx->penalty_freq != 0.0f) {
+        ggml_tensor * penalty_freq = ggml_scale(ctx, counts_f32, sctx->penalty_freq);
+        penalized = ggml_sub(ctx, penalized, penalty_freq);
+    }
+
+    if (sctx->penalty_present != 0.0f) {
+        ggml_tensor * penalty_present = ggml_scale(ctx, active_mask, sctx->penalty_present);
+        penalized = ggml_sub(ctx, penalized, penalty_present);
+    }
+
+    if (sctx->has_candidates) {
+        data->logits = penalized;
+    } else {
+        ggml_tensor * logits_rows = ggml_reshape_2d(ctx, logits, 1, ggml_nelements(logits));
+        ggml_tensor * scatter_rows = ggml_reshape_2d(ctx, penalized, 1, sctx->n_max);
+        logits_rows = ggml_set_rows(ctx, logits_rows, scatter_rows, sctx->inp_token_ids);
+        data->logits = ggml_reshape_1d(ctx, logits_rows, ggml_nelements(logits));
+    }
+}
+
+static void llama_sampler_penalties_backend_set_input(struct llama_sampler * smpl) {
+    auto * sctx = (llama_sampler_penalties *) smpl->ctx;
+
+    if (!sctx->inp_token_ids || !sctx->inp_counts || sctx->n_max <= 0 || sctx->n_vocab <= 0) {
+        return;
+    }
+
+    if (sctx->is_disabled()) {
+        return;
+    }
+
+    // fill active entries from the map
+    int32_t n_active = 0;
+
+    for (const auto & it : sctx->token_count) {
+        GGML_ASSERT(n_active < sctx->n_max);
+        sctx->host_token_ids[n_active] = it.first;
+        sctx->host_counts   [n_active] = it.second;
+        ++n_active;
+    }
+
+    // Sorting is required because backend_apply uses ggml_set_rows (a scatter-back operation)
+    std::vector<std::pair<int32_t, int32_t>> entries;
+    entries.reserve(n_active);
+    for (int32_t i = 0; i < n_active; ++i) {
+        entries.emplace_back(sctx->host_token_ids[i], sctx->host_counts[i]);
+    }
+    std::sort(entries.begin(), entries.end(), [](const auto & a, const auto & b) {
+        return a.first < b.first;
+    });
+    for (int32_t i = 0; i < n_active; ++i) {
+        sctx->host_token_ids[i] = entries[i].first;
+        sctx->host_counts   [i] = entries[i].second;
+    }
+
+    // Padding: Finds a filler token id that is not present in token_count.
+    // Use it to do padding for the arrays, it avoids resizing every time.
+    // The arrays must always have exactly n_max entries (the GPU tensor is a fixed size).
+    int32_t filler = 0;
+    if (n_active < sctx->n_max) {
+        while (sctx->token_count.find(filler) != sctx->token_count.end()) {
+            ++filler;
+        }
+        GGML_ASSERT(filler < sctx->n_vocab);
+    }
+
+    // Fill the rest of the arrays with the filler token id and count 0.
+    // Inactive slots are padded with a unique dummy token ID (count = 0).
+    // The uniqueness matters because ggml_set_rows with duplicate indices can produce non-deterministic or incorrect results.
+    // Using a filler token with count 0 that isn't in the active set is safe, because the active_mask step in backend_apply filters them out via ggml_step(counts_f32)
+    for (int32_t i = n_active; i < sctx->n_max; ++i) {
+        sctx->host_token_ids[i] = filler;
+        sctx->host_counts   [i] = 0;
+    }
+
+    ggml_backend_tensor_set(sctx->inp_token_ids, sctx->host_token_ids.data(), 0, sctx->n_max * sizeof(int32_t));
+    ggml_backend_tensor_set(sctx->inp_counts,    sctx->host_counts.data(),    0, sctx->n_max * sizeof(int32_t));
+}
+
 static struct llama_sampler_i llama_sampler_penalties_i = {
     /* .name              = */ llama_sampler_penalties_name,
     /* .accept            = */ llama_sampler_penalties_accept,
@@ -2753,10 +2958,10 @@ static struct llama_sampler_i llama_sampler_penalties_i = {
     /* .reset             = */ llama_sampler_penalties_reset,
     /* .clone             = */ llama_sampler_penalties_clone,
     /* .free              = */ llama_sampler_penalties_free,
-    /* .backend_init      = */ nullptr,
+    /* .backend_init      = */ llama_sampler_penalties_backend_init,
     /* .backend_accept    = */ nullptr,
-    /* .backend_apply     = */ nullptr,
-    /* .backend_set_input = */ nullptr,
+    /* .backend_apply     = */ llama_sampler_penalties_backend_apply,
+    /* .backend_set_input = */ llama_sampler_penalties_backend_set_input,
 };
 
 struct llama_sampler * llama_sampler_init_penalties(
@@ -2766,22 +2971,18 @@ struct llama_sampler * llama_sampler_init_penalties(
         float penalty_present) {
     penalty_last_n = std::max(penalty_last_n, 0);
 
-    const bool is_empty = (penalty_last_n == 0 || (penalty_repeat == 1.0f && penalty_freq == 0.0f && penalty_present == 0.0f));
-
-    if (is_empty) {
+    if (llama_sampler_penalties::is_disabled(
+                penalty_last_n, penalty_repeat, penalty_freq, penalty_present)) {
         return llama_sampler_init_empty("?penalties");
     }
 
     return llama_sampler_init(
         /* .iface = */ &llama_sampler_penalties_i,
-        /* .ctx   = */ new llama_sampler_penalties {
-            /* .penalty_last_n  = */ penalty_last_n,
-            /* .penalty_repeat  = */ penalty_repeat,
-            /* .penalty_freq    = */ penalty_freq,
-            /* .penalty_present = */ penalty_present,
-            /* .prev            = */ ring_buffer<llama_token>(penalty_last_n),
-            /* .token_count     = */ {},
-        }
+        /* .ctx   = */ new llama_sampler_penalties(
+            penalty_last_n,
+            penalty_repeat,
+            penalty_freq,
+            penalty_present)
     );
 }
 
index 1d3584f903c42161a37701a320952c8c1dcb123f..fd5adb740eab632505cd0a4d999fb55a093a5f84 100644 (file)
@@ -99,6 +99,34 @@ static void test(void) {
     argv = {"binary_name", "-sm", "hello"};
     assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));
 
+    {
+        common_params penalty_params;
+
+        argv = {"binary_name", "--repeat-penalty", "0"};
+        assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
+
+        argv = {"binary_name", "--repeat-penalty", "-1"};
+        assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
+
+        argv = {"binary_name", "--repeat-penalty", "nan"};
+        assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
+
+        argv = {"binary_name", "--repeat-penalty", "inf"};
+        assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
+
+        argv = {"binary_name", "--repeat-penalty", "-inf"};
+        assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
+
+        const char * penalty_options[] = {"--frequency-penalty", "--presence-penalty"};
+        const char * nonfinite_values[] = {"nan", "inf", "-inf"};
+        for (const char * option : penalty_options) {
+            for (const char * value : nonfinite_values) {
+                argv = {"binary_name", option, value};
+                assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
+            }
+        }
+    }
+
     // non-existence arg in specific example (--draft cannot be used outside llama-speculative)
     argv = {"binary_name", "--draft", "123"};
     assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_EMBEDDING));
index c24076e313723930a91f33703a36f65b14cfcea0..1a46468ba2641a2cb4673bf5680696ee2c2947cd 100644 (file)
@@ -8,12 +8,15 @@
 #endif
 
 #include <algorithm>
+#include <cmath>
 #include <cstdlib>
 #include <cstring>
 #include <fstream>
+#include <functional>
 #include <map>
 #include <string>
 #include <unordered_map>
+#include <unordered_set>
 #include <vector>
 
 struct test_args {
@@ -761,6 +764,563 @@ static void test_backend_logit_bias_sampling(const test_params & params) {
     printf("backend logit bias sampling test PASSED\n");
 }
 
+static void accept_prompt(llama_sampler * smpl, const llama_vocab * vocab, const std::string & prompt) {
+    const llama_token bos = llama_vocab_bos(vocab);
+    if (bos != LLAMA_TOKEN_NULL) {
+        llama_sampler_accept(smpl, bos);
+    }
+
+    std::vector<llama_token> tokens(64);
+    int32_t n_tokens = llama_tokenize(vocab, prompt.c_str(), (int32_t) prompt.size(),
+            tokens.data(), (int32_t) tokens.size(), false, false);
+    if (n_tokens < 0) {
+        tokens.resize(-n_tokens);
+        n_tokens = llama_tokenize(vocab, prompt.c_str(), (int32_t) prompt.size(),
+                tokens.data(), (int32_t) tokens.size(), false, false);
+    }
+
+    for (int32_t i = 0; i < n_tokens; ++i) {
+        llama_sampler_accept(smpl, tokens[i]);
+    }
+}
+
+static std::vector<float> decode_raw_logits(const test_params & params, const std::string & prompt) {
+    const int seq_id = 0;
+    const int n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(params.model.get()));
+    std::vector<llama_sampler_seq_config> empty_configs;
+    test_context ctx(params, empty_configs);
+
+    GGML_ASSERT(ctx.decode({{ seq_id, prompt }}));
+
+    float * logits = llama_get_logits_ith(ctx.ctx.get(), ctx.idx_for_seq(seq_id));
+    GGML_ASSERT(logits != nullptr);
+    return std::vector<float>(logits, logits + n_vocab);
+}
+
+static std::vector<llama_token_data> apply_cpu_sampler(
+        const std::vector<float> & raw_logits,
+        llama_sampler * sampler) {
+    std::vector<llama_token_data> data;
+    data.reserve(raw_logits.size());
+    for (llama_token token = 0; token < (llama_token) raw_logits.size(); ++token) {
+        data.push_back({ token, raw_logits[token], 0.0f });
+    }
+
+    llama_token_data_array cur_p = { data.data(), data.size(), -1, false };
+    llama_sampler_apply(sampler, &cur_p);
+    data.resize(cur_p.size);
+    return data;
+}
+
+using sampler_setup_fn = std::function<void(llama_sampler *)>;
+using sampler_init_fn = std::function<llama_sampler *()>;
+
+enum class penalties_position {
+    before_filter,
+    after_filter,
+};
+
+static void add_filter_and_penalties(
+        llama_sampler * chain,
+        const sampler_init_fn & init_filter,
+        int32_t penalty_last_n,
+        float penalty_repeat,
+        float penalty_freq,
+        float penalty_present,
+        penalties_position position) {
+    const auto add_penalties = [&]() {
+        llama_sampler_chain_add(chain, llama_sampler_init_penalties(
+                    penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
+    };
+
+    if (position == penalties_position::before_filter) {
+        add_penalties();
+        llama_sampler_chain_add(chain, init_filter());
+    } else {
+        llama_sampler_chain_add(chain, init_filter());
+        add_penalties();
+    }
+}
+
+static llama_sampler_ptr make_sampler_chain(
+        const sampler_setup_fn & add_samplers,
+        const sampler_setup_fn & accept_history) {
+    llama_sampler_ptr chain(llama_sampler_chain_init(llama_sampler_chain_default_params()));
+    add_samplers(chain.get());
+    accept_history(chain.get());
+    return chain;
+}
+
+struct backend_sampler_output {
+    std::vector<float> logits;
+    std::vector<llama_token> candidates;
+};
+
+static backend_sampler_output run_backend_sampler(
+        const test_params & params,
+        const std::string & prompt,
+        llama_sampler * sampler) {
+    const int seq_id = 0;
+    std::vector<llama_sampler_seq_config> configs = {{ seq_id, sampler }};
+    test_context ctx(params, configs);
+
+    GGML_ASSERT(ctx.decode({{ seq_id, prompt }}));
+    llama_synchronize(ctx.ctx.get());
+
+    const int32_t idx = ctx.idx_for_seq(seq_id);
+    const uint32_t n_logits = llama_get_sampled_logits_count_ith(ctx.ctx.get(), idx);
+    const uint32_t n_candidates = llama_get_sampled_candidates_count_ith(ctx.ctx.get(), idx);
+    float * logits = llama_get_sampled_logits_ith(ctx.ctx.get(), idx);
+    llama_token * candidates = llama_get_sampled_candidates_ith(ctx.ctx.get(), idx);
+    GGML_ASSERT(logits != nullptr);
+
+    backend_sampler_output result;
+    result.logits.assign(logits, logits + n_logits);
+    result.candidates.resize(n_logits);
+
+    if (n_candidates == 0) {
+        for (uint32_t i = 0; i < n_logits; ++i) {
+            result.candidates[i] = (llama_token) i;
+        }
+    } else {
+        GGML_ASSERT(candidates != nullptr);
+        GGML_ASSERT(n_candidates == n_logits);
+        std::memcpy(result.candidates.data(), candidates, n_candidates * sizeof(llama_token));
+    }
+
+    return result;
+}
+
+struct sampler_comparison_output {
+    std::vector<llama_token_data> expected;
+    backend_sampler_output actual;
+};
+
+static sampler_comparison_output run_sampler_comparison(
+        const test_params & params,
+        const std::string & prompt,
+        const std::vector<float> & raw_logits,
+        const sampler_setup_fn & add_samplers,
+        const sampler_setup_fn & accept_history) {
+    llama_sampler_ptr cpu_chain = make_sampler_chain(add_samplers, accept_history);
+    llama_sampler_ptr backend_chain = make_sampler_chain(add_samplers, accept_history);
+    return {
+        apply_cpu_sampler(raw_logits, cpu_chain.get()),
+        run_backend_sampler(params, prompt, backend_chain.get()),
+    };
+}
+
+static std::unordered_map<llama_token, float> map_logits(const std::vector<llama_token_data> & data) {
+    std::unordered_map<llama_token, float> result;
+    result.reserve(data.size());
+    for (const auto & item : data) {
+        result[item.id] = item.logit;
+    }
+    return result;
+}
+
+struct sampler_comparison_stats {
+    int n_mismatch = 0;
+    int n_masked = 0;
+    float max_diff = 0.0f;
+};
+
+static sampler_comparison_stats compare_sampler_outputs(
+        const char * name,
+        const std::unordered_map<llama_token, float> & expected,
+        const backend_sampler_output & actual,
+        bool allow_extra_candidates = false) {
+    GGML_ASSERT(actual.logits.size() == actual.candidates.size());
+
+    sampler_comparison_stats result;
+    std::unordered_set<llama_token> seen;
+    seen.reserve(actual.candidates.size());
+
+    for (size_t i = 0; i < actual.logits.size(); ++i) {
+        const llama_token token = actual.candidates[i];
+        const float logit = actual.logits[i];
+        if (!seen.insert(token).second || std::isnan(logit)) {
+            if (result.n_mismatch < 5) {
+                printf("%s token %d has invalid backend output\n", name, token);
+            }
+            ++result.n_mismatch;
+            continue;
+        }
+
+        const auto it = expected.find(token);
+        if (it == expected.end()) {
+            if (std::isinf(logit) && logit < 0.0f) {
+                ++result.n_masked;
+            } else if (!allow_extra_candidates) {
+                if (result.n_mismatch < 5) {
+                    printf("%s token %d was not masked\n", name, token);
+                }
+                ++result.n_mismatch;
+            }
+            continue;
+        }
+
+        const float diff = fabsf(it->second - logit);
+        result.max_diff = std::max(result.max_diff, diff);
+        if (!std::isfinite(logit) || diff > 1e-3f) {
+            if (result.n_mismatch < 5) {
+                printf("%s mismatch token %d: cpu=%.6f backend=%.6f diff=%.6f\n",
+                        name, token, it->second, logit, diff);
+            }
+            ++result.n_mismatch;
+        }
+    }
+
+    for (const auto & item : expected) {
+        if (seen.find(item.first) == seen.end()) {
+            if (result.n_mismatch < 5) {
+                printf("%s missing backend token %d\n", name, item.first);
+            }
+            ++result.n_mismatch;
+        }
+    }
+
+    printf("%s logits: max_diff=%.6f n_masked=%d n_mismatch=%d\n",
+            name, result.max_diff, result.n_masked, result.n_mismatch);
+    return result;
+}
+
+static float find_backend_logit(const backend_sampler_output & output, llama_token token) {
+    for (size_t i = 0; i < output.candidates.size(); ++i) {
+        if (output.candidates[i] == token) {
+            return output.logits[i];
+        }
+    }
+    GGML_ABORT("backend token not found");
+}
+
+static sampler_comparison_output run_penalties_comparison(
+        const test_params & params,
+        int32_t penalty_last_n,
+        float penalty_repeat,
+        float penalty_freq,
+        float penalty_present,
+        const std::string & prompt,
+        const std::function<void(llama_sampler *)> & extra_accept = {}) {
+    const auto * vocab = llama_model_get_vocab(params.model.get());
+    const std::vector<float> raw_logits = decode_raw_logits(params, prompt);
+    const auto add_samplers = [&](llama_sampler * chain) {
+        llama_sampler_chain_add(chain, llama_sampler_init_penalties(
+                    penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
+    };
+    const auto accept_history = [&](llama_sampler * chain) {
+        accept_prompt(chain, vocab, prompt);
+        if (extra_accept) {
+            extra_accept(chain);
+        }
+    };
+
+    return run_sampler_comparison(
+            params, prompt, raw_logits, add_samplers, accept_history);
+}
+
+static void compare_penalties_logits(
+        const test_params & params,
+        int32_t penalty_last_n,
+        float penalty_repeat,
+        float penalty_freq,
+        float penalty_present,
+        const std::string & prompt,
+        const std::function<void(llama_sampler *)> & extra_accept = {}) {
+    const sampler_comparison_output output = run_penalties_comparison(
+            params, penalty_last_n, penalty_repeat, penalty_freq, penalty_present, prompt, extra_accept);
+
+    GGML_ASSERT(output.expected.size() == output.actual.logits.size());
+
+    const sampler_comparison_stats stats = compare_sampler_outputs(
+            "penalties", map_logits(output.expected), output.actual);
+    GGML_ASSERT(stats.n_masked == 0);
+    GGML_ASSERT(stats.n_mismatch == 0);
+}
+
+static void test_penalty_parameter_values(const test_params & params) {
+    struct penalty_test_case {
+        const char * name;
+        float repeat;
+        float frequency;
+        float presence;
+    };
+
+    const penalty_test_case cases[] = {
+        { "frequency -1",   1.0f, -1.0f,     0.0f },
+        { "frequency 0",    1.0f,  0.0f,     0.0f },
+        { "frequency 1",    1.0f,  1.0f,     0.0f },
+        { "presence -1",    1.0f,  0.0f,    -1.0f },
+        { "presence 0",     1.0f,  0.0f,     0.0f },
+        { "presence 1",     1.0f,  0.0f,     1.0f },
+        { "repeat 1",       1.0f,  0.0f,     0.0f },
+    };
+
+    int n_failed = 0;
+    for (const auto & test : cases) {
+        const sampler_comparison_output output = run_penalties_comparison(
+                params, 64, test.repeat, test.frequency, test.presence, "Hello Hello world");
+        GGML_ASSERT(output.expected.size() == output.actual.logits.size());
+        const sampler_comparison_stats stats = compare_sampler_outputs(
+                test.name, map_logits(output.expected), output.actual);
+        n_failed += stats.n_mismatch != 0;
+    }
+
+    GGML_ASSERT(n_failed == 0);
+}
+
+static void compare_top_k_penalties_logits(
+        const test_params & params,
+        int32_t k,
+        int32_t penalty_last_n,
+        float penalty_repeat,
+        float penalty_freq,
+        float penalty_present,
+        const std::string & prompt,
+        penalties_position position) {
+    const auto * vocab = llama_model_get_vocab(params.model.get());
+    const std::vector<float> raw_logits = decode_raw_logits(params, prompt);
+    const int n_vocab = (int) raw_logits.size();
+
+    GGML_ASSERT(n_vocab > k);
+
+    const sampler_init_fn init_top_k = [k]() {
+        return llama_sampler_init_top_k(k);
+    };
+    llama_sampler_ptr top_k(init_top_k());
+    const std::vector<llama_token_data> top_k_data = apply_cpu_sampler(raw_logits, top_k.get());
+    GGML_ASSERT(top_k_data.size() == (size_t) k);
+    const llama_token retained_history_token = top_k_data[0].id;
+
+    llama_token excluded_history_token = LLAMA_TOKEN_NULL;
+    for (llama_token token = 0; token < n_vocab; ++token) {
+        const auto it = std::find_if(top_k_data.begin(), top_k_data.end(), [token](const llama_token_data & data) {
+            return data.id == token;
+        });
+        if (it == top_k_data.end()) {
+            excluded_history_token = token;
+            break;
+        }
+    }
+    GGML_ASSERT(excluded_history_token != LLAMA_TOKEN_NULL);
+
+    const auto add_samplers = [&](llama_sampler * chain) {
+        add_filter_and_penalties(chain, init_top_k,
+                penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position);
+    };
+
+    auto accept_history = [&](llama_sampler * smpl) {
+        accept_prompt(smpl, vocab, prompt);
+        llama_sampler_accept(smpl, excluded_history_token);
+        llama_sampler_accept(smpl, excluded_history_token);
+        llama_sampler_accept(smpl, retained_history_token);
+        llama_sampler_accept(smpl, retained_history_token);
+    };
+
+    const sampler_comparison_output output = run_sampler_comparison(
+            params, prompt, raw_logits, add_samplers, accept_history);
+
+    GGML_ASSERT(output.expected.size() == (size_t) k);
+    GGML_ASSERT(output.actual.logits.size() == (size_t) k);
+
+    const std::unordered_map<llama_token, float> expected_logits = map_logits(output.expected);
+
+    if (position == penalties_position::after_filter) {
+        GGML_ASSERT(expected_logits.find(retained_history_token) != expected_logits.end());
+        GGML_ASSERT(fabsf(expected_logits.at(retained_history_token) - raw_logits[retained_history_token]) > 1e-6f);
+        GGML_ASSERT(expected_logits.find(excluded_history_token) == expected_logits.end());
+        GGML_ASSERT(std::find(output.actual.candidates.begin(), output.actual.candidates.end(),
+                    excluded_history_token) == output.actual.candidates.end());
+    } else {
+        const std::unordered_map<llama_token, float> unpenalized_logits = map_logits(top_k_data);
+        bool changed = false;
+        for (const auto & item : expected_logits) {
+            const auto it = unpenalized_logits.find(item.first);
+            if (it == unpenalized_logits.end() || fabsf(it->second - item.second) > 1e-6f) {
+                changed = true;
+                break;
+            }
+        }
+        GGML_ASSERT(changed);
+    }
+
+    const char * name = position == penalties_position::before_filter
+        ? "penalties top-k"
+        : "top-k penalties";
+    const sampler_comparison_stats stats = compare_sampler_outputs(
+            name, expected_logits, output.actual);
+    GGML_ASSERT(stats.n_masked == 0);
+    GGML_ASSERT(stats.n_mismatch == 0);
+}
+
+static void compare_masking_penalties_logits(
+        const test_params & params,
+        const char * filter_name,
+        const sampler_init_fn & init_filter,
+        int32_t penalty_last_n,
+        float penalty_repeat,
+        float penalty_freq,
+        float penalty_present,
+        const std::string & prompt,
+        penalties_position position,
+        bool allow_extra_candidates,
+        bool add_history = true) {
+    const auto * vocab = llama_model_get_vocab(params.model.get());
+    const std::vector<float> raw_logits = decode_raw_logits(params, prompt);
+    const int n_vocab = (int) raw_logits.size();
+    llama_sampler_ptr filter(init_filter());
+    const std::vector<llama_token_data> filtered_data = apply_cpu_sampler(raw_logits, filter.get());
+    GGML_ASSERT(!filtered_data.empty());
+    GGML_ASSERT(filtered_data.size() < (size_t) n_vocab);
+
+    const llama_token penalized_token = filtered_data[0].id;
+    std::unordered_set<llama_token> retained_tokens;
+    retained_tokens.reserve(filtered_data.size());
+    for (const auto & data : filtered_data) {
+        retained_tokens.insert(data.id);
+    }
+
+    llama_token masked_token = LLAMA_TOKEN_NULL;
+    for (llama_token token = 0; token < n_vocab; ++token) {
+        if (retained_tokens.find(token) == retained_tokens.end()) {
+            masked_token = token;
+            break;
+        }
+    }
+    GGML_ASSERT(masked_token != LLAMA_TOKEN_NULL);
+
+    const auto add_samplers = [&](llama_sampler * chain) {
+        add_filter_and_penalties(chain, init_filter,
+                penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position);
+    };
+    auto accept_history = [&](llama_sampler * smpl) {
+        if (!add_history) {
+            return;
+        }
+        accept_prompt(smpl, vocab, prompt);
+        llama_sampler_accept(smpl, penalized_token);
+        llama_sampler_accept(smpl, penalized_token);
+        llama_sampler_accept(smpl, masked_token);
+        llama_sampler_accept(smpl, masked_token);
+    };
+
+    const sampler_comparison_output output = run_sampler_comparison(
+            params, prompt, raw_logits, add_samplers, accept_history);
+
+    GGML_ASSERT(output.actual.logits.size() == (size_t) n_vocab);
+
+    const std::unordered_map<llama_token, float> expected_logits = map_logits(output.expected);
+
+    GGML_ASSERT(expected_logits.find(masked_token) == expected_logits.end());
+    if (add_history) {
+        if (position == penalties_position::after_filter) {
+            GGML_ASSERT(expected_logits.find(penalized_token) != expected_logits.end());
+            GGML_ASSERT(fabsf(expected_logits.at(penalized_token) - raw_logits[penalized_token]) > 1e-6f);
+        } else {
+            llama_sampler_ptr penalties(llama_sampler_init_penalties(
+                        penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
+            accept_history(penalties.get());
+            const std::unordered_map<llama_token, float> penalized_logits =
+                map_logits(apply_cpu_sampler(raw_logits, penalties.get()));
+            GGML_ASSERT(fabsf(penalized_logits.at(penalized_token) - raw_logits[penalized_token]) > 1e-6f);
+        }
+    }
+
+    const std::string name = position == penalties_position::before_filter
+        ? "penalties " + std::string(filter_name)
+        : std::string(filter_name) + " penalties";
+    const sampler_comparison_stats stats = compare_sampler_outputs(
+            name.c_str(), expected_logits, output.actual, allow_extra_candidates);
+    const float masked_logit = find_backend_logit(output.actual, masked_token);
+    GGML_ASSERT(stats.n_masked > 0);
+    GGML_ASSERT(std::isinf(masked_logit) && masked_logit < 0.0f);
+    GGML_ASSERT(stats.n_mismatch == 0);
+}
+
+static void test_backend_penalties_sampling(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");
+
+    printf("Testing backend penalties with penalty_last_n > 64\n");
+    const auto * vocab = llama_model_get_vocab(params.model.get());
+    std::vector<llama_token> tokens(8);
+    int32_t n_tok = llama_tokenize(vocab, "a", 1, tokens.data(), (int32_t) tokens.size(), false, false);
+    if (n_tok < 0) {
+        tokens.resize(-n_tok);
+        n_tok = llama_tokenize(vocab, "a", 1, tokens.data(), (int32_t) tokens.size(), false, false);
+    }
+    GGML_ASSERT(n_tok > 0);
+    const llama_token tok = tokens[0];
+
+    compare_penalties_logits(params, 80, 1.15f, 0.1f, 0.05f, "a", [tok](llama_sampler * smpl) {
+        // accept_prompt already accepted BOS + one 'a'; fill the ring to n=80
+        for (int i = 0; i < 78; ++i) {
+            llama_sampler_accept(smpl, tok);
+        }
+    });
+
+    printf("Testing backend penalties without filler entries\n");
+    compare_penalties_logits(params, 64, 1.1f, 0.5f, 0.25f, "Hello", [](llama_sampler * smpl) {
+        for (llama_token token = 0; token < 64; ++token) {
+            llama_sampler_accept(smpl, token);
+        }
+    });
+
+    printf("Testing backend top-k followed by penalties\n");
+    compare_top_k_penalties_logits(params, 8, 64, 1.1f, 0.5f, 0.25f, "Hello",
+            penalties_position::after_filter);
+
+    printf("Testing backend penalties followed by top-k\n");
+    compare_top_k_penalties_logits(params, 8, 64, 1.1f, 0.5f, 0.25f, "Hello",
+            penalties_position::before_filter);
+
+    printf("Testing backend top-p followed by penalties\n");
+    compare_masking_penalties_logits(params, "top-p", []() {
+        return llama_sampler_init_top_p(0.9f, 0);
+    }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true);
+
+    printf("Testing backend top-p followed by penalties with a large history window\n");
+    compare_masking_penalties_logits(params, "top-p large-window", []() {
+        return llama_sampler_init_top_p(0.9f, 0);
+    }, 4096, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true);
+
+    printf("Testing backend penalties followed by top-p\n");
+    compare_masking_penalties_logits(params, "top-p", []() {
+        return llama_sampler_init_top_p(0.9f, 0);
+    }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::before_filter, true);
+
+    printf("Testing backend min-p followed by penalties\n");
+    compare_masking_penalties_logits(params, "min-p", []() {
+        return llama_sampler_init_min_p(0.1f, 0);
+    }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, false);
+
+    printf("Testing backend penalties followed by min-p\n");
+    compare_masking_penalties_logits(params, "min-p", []() {
+        return llama_sampler_init_min_p(0.1f, 0);
+    }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::before_filter, false);
+
+    printf("Testing backend top-p followed by penalties with empty history\n");
+    compare_masking_penalties_logits(params, "top-p empty", []() {
+        return llama_sampler_init_top_p(0.9f, 0);
+    }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true, false);
+
+    printf("Testing backend top-p followed by individual penalties\n");
+    compare_masking_penalties_logits(params, "top-p repeat", []() {
+        return llama_sampler_init_top_p(0.9f, 0);
+    }, 64, 1.1f, 0.0f, 0.0f, "Hello", penalties_position::after_filter, true);
+    compare_masking_penalties_logits(params, "top-p frequency", []() {
+        return llama_sampler_init_top_p(0.9f, 0);
+    }, 64, 1.0f, 0.5f, 0.0f, "Hello", penalties_position::after_filter, true);
+    compare_masking_penalties_logits(params, "top-p presence", []() {
+        return llama_sampler_init_top_p(0.9f, 0);
+    }, 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);
+
+    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) {
@@ -1014,6 +1574,7 @@ struct backend_test_case {
 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  },
index 4655b518e21f5ccafbd32274e3c022b5b300715f..5d2798cc14e9295646bb8e570bbec166c9ecc72c 100644 (file)
@@ -1807,7 +1807,8 @@ private:
         // initialize samplers
         if (task.need_sampling()) {
             try {
-                slot.smpl.reset(common_sampler_init(model_tgt, task.params.sampling));
+                slot.smpl.reset(common_sampler_init(
+                        model_tgt, task.params.sampling, (int32_t) llama_n_ctx(ctx_tgt)));
             } catch (std::exception & e) {
                 std::string err_msg = std::string("Failed to initialize samplers: ") + e.what();
                 send_error(task, err_msg, ERROR_TYPE_INVALID_REQUEST);