CUDA: Add backend sampler for penalties sampler (#25262)

* 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
This commit is contained in:
Konrad Moren
2026-08-03 14:26:09 +02:00
committed by GitHub
parent 9bd4c09ea5
commit 96278e39fc
10 changed files with 860 additions and 31 deletions
+18 -3
View 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(
+2 -1
View 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()) };
}
+19 -2
View 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;
+4 -1
View 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);
+4 -3
View 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(
+1
View 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);
+221 -20
View 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)
);
}
+28
View 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));
+561
View 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 },
+2 -1
View 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);