diff --git a/common/arg.cpp b/common/arg.cpp index 772422f68..305938fcb 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -27,6 +27,7 @@ #include #include #include +#include #include #include #include @@ -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( diff --git a/common/common.cpp b/common/common.cpp index ff27d392f..c941fd505 100644 --- a/common/common.cpp +++ b/common/common.cpp @@ -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()) }; } diff --git a/common/sampling.cpp b/common/sampling.cpp index 256ac161e..5698c0263 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -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; diff --git a/common/sampling.h b/common/sampling.h index 4191988bb..91e2cea78 100644 --- a/common/sampling.h +++ b/common/sampling.h @@ -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); diff --git a/include/llama.h b/include/llama.h index 6e53e2297..f2d7e3885 100644 --- a/include/llama.h +++ b/include/llama.h @@ -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( diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index e12a8cdc2..1a3569230 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -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); diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp index a9cb6bee5..b2f1abe73 100644 --- a/src/llama-sampler.cpp +++ b/src/llama-sampler.cpp @@ -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 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 host_token_ids; + std::vector 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> 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(penalty_last_n), - /* .token_count = */ {}, - } + /* .ctx = */ new llama_sampler_penalties( + penalty_last_n, + penalty_repeat, + penalty_freq, + penalty_present) ); } diff --git a/tests/test-arg-parser.cpp b/tests/test-arg-parser.cpp index 1d3584f90..fd5adb740 100644 --- a/tests/test-arg-parser.cpp +++ b/tests/test-arg-parser.cpp @@ -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)); diff --git a/tests/test-backend-sampler.cpp b/tests/test-backend-sampler.cpp index c24076e31..1a46468ba 100644 --- a/tests/test-backend-sampler.cpp +++ b/tests/test-backend-sampler.cpp @@ -8,12 +8,15 @@ #endif #include +#include #include #include #include +#include #include #include #include +#include #include 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 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 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 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(logits, logits + n_vocab); +} + +static std::vector apply_cpu_sampler( + const std::vector & raw_logits, + llama_sampler * sampler) { + std::vector 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; +using sampler_init_fn = std::function; + +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 logits; + std::vector 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 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 expected; + backend_sampler_output actual; +}; + +static sampler_comparison_output run_sampler_comparison( + const test_params & params, + const std::string & prompt, + const std::vector & 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 map_logits(const std::vector & data) { + std::unordered_map 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 & 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 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 & extra_accept = {}) { + const auto * vocab = llama_model_get_vocab(params.model.get()); + const std::vector 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 & 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 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 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 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 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 raw_logits = decode_raw_logits(params, prompt); + const int n_vocab = (int) raw_logits.size(); + llama_sampler_ptr filter(init_filter()); + const std::vector 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 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 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 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 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 }, diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 4655b518e..5d2798cc1 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -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);