common : simplify sampler chain initialization

This commit is contained in:
Georgi Gerganov
2025-12-01 17:11:11 +02:00
parent 217469f07f
commit 4032ce2378
7 changed files with 171 additions and 178 deletions
+131 -12
View File
@@ -349,7 +349,7 @@ static uint32_t get_rng_seed(uint32_t seed) {
// llama_sampler API
struct llama_sampler * llama_sampler_init(
const struct llama_sampler_i * iface,
struct llama_sampler_i * iface,
llama_sampler_context_t ctx) {
return new llama_sampler {
/* .iface = */ iface,
@@ -468,6 +468,42 @@ llama_token llama_sampler_sample(struct llama_sampler * smpl, struct llama_conte
return token;
}
// backend sampling (empty iface)
static void llama_sampler_empty_backend_init(
struct llama_sampler * smpl,
ggml_backend_buffer_type_t buft) {
GGML_UNUSED(smpl);
GGML_UNUSED(buft);
}
static void llama_sampler_empty_backend_accept(
struct llama_sampler * smpl,
ggml_context * ctx,
ggml_cgraph * gf,
struct ggml_tensor * selected_token) {
GGML_UNUSED(smpl);
GGML_UNUSED(ctx);
GGML_UNUSED(gf);
GGML_UNUSED(selected_token);
}
static void llama_sampler_empty_backend_apply(
struct llama_sampler * smpl,
struct ggml_context * ctx,
struct ggml_cgraph * gf,
struct llama_sampler_data * data) {
GGML_UNUSED(smpl);
GGML_UNUSED(ctx);
GGML_UNUSED(gf);
GGML_UNUSED(data);
}
static void llama_sampler_empty_backend_set_input(struct llama_sampler * smpl) {
GGML_UNUSED(smpl);
}
// sampler chain
static const char * llama_sampler_chain_name(const struct llama_sampler * /*smpl*/) {
@@ -1171,7 +1207,7 @@ static void llama_sampler_top_p_backend_apply(
ggml_set_output(data->candidates);
ggml_build_forward_expand(gf, data->candidates);
ggml_set_output(data->logits);
ggml_build_forward_expand(gf, data->logits);
}
@@ -1446,13 +1482,24 @@ static struct llama_sampler_i llama_sampler_typical_i = {
};
struct llama_sampler * llama_sampler_init_typical(float p, size_t min_keep) {
return llama_sampler_init(
auto * res = llama_sampler_init(
/* .iface = */ &llama_sampler_typical_i,
/* .ctx = */ new llama_sampler_typical {
/* .p = */ p,
/* .min_keep = */ min_keep,
}
);
const bool is_empty = (p >= 1.0f);
if (is_empty) {
res->iface->backend_init = llama_sampler_empty_backend_init;
res->iface->backend_accept = llama_sampler_empty_backend_accept;
res->iface->backend_apply = llama_sampler_empty_backend_apply;
res->iface->backend_set_input = llama_sampler_empty_backend_set_input;
}
return res;
}
// temp
@@ -1615,6 +1662,27 @@ static void llama_sampler_temp_ext_free(struct llama_sampler * smpl) {
delete (llama_sampler_temp_ext *) smpl->ctx;
}
static void llama_sampler_temp_ext_backend_apply(
struct llama_sampler * smpl,
struct ggml_context * ctx,
struct ggml_cgraph * gf,
struct llama_sampler_data * data) {
auto * ctx_data = (llama_sampler_temp *) smpl->ctx;
if (ctx_data->temp <= 0.0f) {
return;
}
struct ggml_tensor * scaled = ggml_scale(ctx, data->logits, 1.0f / ctx_data->temp);
ggml_set_name(scaled, "temp_scaled");
// Make sure the scaled tensor is contiguous for subsequent operations
data->logits = ggml_cont(ctx, scaled);
ggml_set_name(data->logits, "temp_scaled_logits");
ggml_build_forward_expand(gf, data->logits);
}
static struct llama_sampler_i llama_sampler_temp_ext_i = {
/* .name = */ llama_sampler_temp_ext_name,
/* .accept = */ nullptr,
@@ -1629,7 +1697,7 @@ static struct llama_sampler_i llama_sampler_temp_ext_i = {
};
struct llama_sampler * llama_sampler_init_temp_ext(float temp, float delta, float exponent) {
return llama_sampler_init(
auto * res = llama_sampler_init(
/* .iface = */ &llama_sampler_temp_ext_i,
/* .ctx = */ new llama_sampler_temp_ext {
/* .temp = */ temp,
@@ -1637,6 +1705,14 @@ struct llama_sampler * llama_sampler_init_temp_ext(float temp, float delta, floa
/* .exponent = */ exponent,
}
);
const bool is_backend = delta <= 0.0f;
if (is_backend) {
res->iface->backend_apply = llama_sampler_temp_ext_backend_apply;
}
return res;
}
// xtc
@@ -1727,8 +1803,9 @@ static struct llama_sampler_i llama_sampler_xtc_i = {
};
struct llama_sampler * llama_sampler_init_xtc(float p, float t, size_t min_keep, uint32_t seed) {
auto seed_cur = get_rng_seed(seed);
return llama_sampler_init(
const auto seed_cur = get_rng_seed(seed);
auto * res = llama_sampler_init(
/* .iface = */ &llama_sampler_xtc_i,
/* .ctx = */ new llama_sampler_xtc {
/* .probability = */ p,
@@ -1739,6 +1816,17 @@ struct llama_sampler * llama_sampler_init_xtc(float p, float t, size_t min_keep,
/* .rng = */ std::mt19937(seed_cur),
}
);
const bool is_empty = (p <= 0.0f || t > 0.5f);
if (is_empty) {
res->iface->backend_init = llama_sampler_empty_backend_init;
res->iface->backend_accept = llama_sampler_empty_backend_accept;
res->iface->backend_apply = llama_sampler_empty_backend_apply;
res->iface->backend_set_input = llama_sampler_empty_backend_set_input;
}
return res;
}
// mirostat
@@ -2280,7 +2368,7 @@ struct llama_sampler * llama_sampler_init_penalties(
float penalty_present) {
penalty_last_n = std::max(penalty_last_n, 0);
return llama_sampler_init(
auto * res = llama_sampler_init(
/* .iface = */ &llama_sampler_penalties_i,
/* .ctx = */ new llama_sampler_penalties {
/* .penalty_last_n = */ penalty_last_n,
@@ -2291,6 +2379,17 @@ struct llama_sampler * llama_sampler_init_penalties(
/* .token_count = */ {},
}
);
const bool is_empty = (penalty_last_n == 0 || (penalty_repeat == 1.0f && penalty_freq == 0.0f && penalty_present == 0.0f));
if (is_empty) {
res->iface->backend_init = llama_sampler_empty_backend_init;
res->iface->backend_accept = llama_sampler_empty_backend_accept;
res->iface->backend_apply = llama_sampler_empty_backend_apply;
res->iface->backend_set_input = llama_sampler_empty_backend_set_input;
}
return res;
}
// top-n-sigma
@@ -2317,9 +2416,7 @@ static void llama_sampler_top_n_sigma_apply(struct llama_sampler * smpl, llama_t
for (size_t i = 0; i < cur_p->size; ++i) {
// Only count non-negative infinity values
if (cur_p->data[i].logit != -INFINITY) {
if (cur_p->data[i].logit > max) {
max = cur_p->data[i].logit;
}
max = std::max(max, cur_p->data[i].logit);
logits_sum += cur_p->data[i].logit;
valid_count++;
}
@@ -2369,12 +2466,23 @@ static struct llama_sampler_i llama_sampler_top_n_sigma_i = {
};
struct llama_sampler * llama_sampler_init_top_n_sigma(float n) {
return llama_sampler_init(
auto * res = llama_sampler_init(
/* .iface = */ &llama_sampler_top_n_sigma_i,
/* .ctx = */ new llama_sampler_top_n_sigma {
/* .n = */ n,
}
);
const bool is_empty = (n <= 0.0f);
if (is_empty) {
res->iface->backend_init = llama_sampler_empty_backend_init;
res->iface->backend_accept = llama_sampler_empty_backend_accept;
res->iface->backend_apply = llama_sampler_empty_backend_apply;
res->iface->backend_set_input = llama_sampler_empty_backend_set_input;
}
return res;
}
// DRY
@@ -2733,7 +2841,7 @@ struct llama_sampler * llama_sampler_init_dry(const struct llama_vocab * vocab,
}
}
return llama_sampler_init(
auto * res = llama_sampler_init(
/* .iface = */ &llama_sampler_dry_i,
/* .ctx = */ new llama_sampler_dry {
/* .total_context_size = */ n_ctx_train,
@@ -2747,6 +2855,15 @@ struct llama_sampler * llama_sampler_init_dry(const struct llama_vocab * vocab,
/* .last_tokens = */ dry_enabled ? ring_buffer<llama_token>(effective_dry_penalty_last_n) : ring_buffer<llama_token>(0),
}
);
if (!dry_enabled) {
res->iface->backend_init = llama_sampler_empty_backend_init;
res->iface->backend_accept = llama_sampler_empty_backend_accept;
res->iface->backend_apply = llama_sampler_empty_backend_apply;
res->iface->backend_set_input = llama_sampler_empty_backend_set_input;
}
return res;
}
// wrapper for test-sampling.cpp
@@ -2854,6 +2971,8 @@ static void llama_sampler_logit_bias_backend_apply(
// Add the sparse logit logit_bias to the logits
struct ggml_tensor * logit_biased = ggml_add_inplace(ctx, data->logits, sctx->inp_logit_bias);
data->logits = logit_biased;
ggml_build_forward_expand(gf, logit_biased);
}