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:
@@ -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));
|
||||
|
||||
@@ -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 },
|
||||
|
||||
Reference in New Issue
Block a user