llama : move n_vocab from llama_sampler_data to penalty_sampler (#26520)

This matches how it is done for logit_bias and mirostat samplers, see
https://github.com/ggml-org/llama.cpp/pull/25262#discussion_r3703951151
This commit is contained in:
Oliver Simons
2026-08-04 09:02:49 +03:00
committed by GitHub
parent 22dc605c4e
commit 935cad6497
6 changed files with 16 additions and 13 deletions
+6 -5
View File
@@ -823,6 +823,7 @@ enum class penalties_position {
static void add_filter_and_penalties(
llama_sampler * chain,
const sampler_init_fn & init_filter,
int32_t n_vocab,
int32_t penalty_last_n,
float penalty_repeat,
float penalty_freq,
@@ -830,7 +831,7 @@ static void add_filter_and_penalties(
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));
n_vocab, penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
};
if (position == penalties_position::before_filter) {
@@ -1006,7 +1007,7 @@ static sampler_comparison_output run_penalties_comparison(
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));
llama_vocab_n_tokens(vocab), penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
};
const auto accept_history = [&](llama_sampler * chain) {
accept_prompt(chain, vocab, prompt);
@@ -1105,7 +1106,7 @@ static void compare_top_k_penalties_logits(
GGML_ASSERT(excluded_history_token != LLAMA_TOKEN_NULL);
const auto add_samplers = [&](llama_sampler * chain) {
add_filter_and_penalties(chain, init_top_k,
add_filter_and_penalties(chain, init_top_k, n_vocab,
penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position);
};
@@ -1190,7 +1191,7 @@ static void compare_masking_penalties_logits(
GGML_ASSERT(masked_token != LLAMA_TOKEN_NULL);
const auto add_samplers = [&](llama_sampler * chain) {
add_filter_and_penalties(chain, init_filter,
add_filter_and_penalties(chain, init_filter, n_vocab,
penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position);
};
auto accept_history = [&](llama_sampler * smpl) {
@@ -1218,7 +1219,7 @@ static void compare_masking_penalties_logits(
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));
n_vocab, 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()));
+1 -1
View File
@@ -144,7 +144,7 @@ static void test_penalties(
sampler_tester tester(probs, probs_expected);
auto * sampler = llama_sampler_init_penalties(last_tokens.size(), repeat_penalty, alpha_frequency, alpha_presence);
auto * sampler = llama_sampler_init_penalties((int32_t) probs.size(), (int32_t) last_tokens.size(), repeat_penalty, alpha_frequency, alpha_presence);
for (size_t i = 0; i < last_tokens.size(); i++) {
llama_sampler_accept(sampler, last_tokens[i]);