sampling : introduce sampling_info struct

This commit introduces a sampling_info struct to encapsulate all
backend sampling related data within the llama_context class.

It also updates to use more descriptive names for sampled tokens and
candidates in the backend sampler ggml data structure.
This commit is contained in:
Daniel Bevenius
2025-11-20 14:45:56 +01:00
parent ed4345bdd9
commit 0d28b16bdc
6 changed files with 91 additions and 85 deletions
+13 -13
View File
@@ -504,10 +504,10 @@ void llm_graph_result::reset() {
t_logits = nullptr;
t_embd = nullptr;
t_embd_pooled = nullptr;
t_sampled_tokens.clear();
t_sampled.clear();
t_sampled_probs.clear();
t_sampled_logits.clear();
t_sampled_token_ids.clear();
t_candidates.clear();
params = {};
@@ -2098,17 +2098,17 @@ void llm_graph_context::build_sampling() const {
ggml_format_name(logits_seq, "logits_seq_%d", seq_id);
struct llama_sampler_ggml_data ggml_data = {
/*.logits =*/ logits_seq,
/*.probs =*/ nullptr,
/*.sampled_token =*/ nullptr,
/*.filtered_ids =*/ nullptr,
/*.logits =*/ logits_seq,
/*.probs =*/ nullptr,
/*.sampled =*/ nullptr,
/*.candidates =*/ nullptr,
};
llama_sampler_apply_ggml(sampler, ctx0, gf, &ggml_data);
if (ggml_data.sampled_token != nullptr) {
res->t_sampled_tokens[seq_id] = ggml_data.sampled_token;
ggml_build_forward_expand(gf, ggml_data.sampled_token);
if (ggml_data.sampled != nullptr) {
res->t_sampled[seq_id] = ggml_data.sampled;
ggml_build_forward_expand(gf, ggml_data.sampled);
}
if (ggml_data.probs != nullptr) {
@@ -2121,16 +2121,16 @@ void llm_graph_context::build_sampling() const {
ggml_build_forward_expand(gf, res->t_sampled_logits[seq_id]);
}
if (ggml_data.filtered_ids != nullptr) {
res->t_sampled_token_ids[seq_id] = ggml_data.filtered_ids;
ggml_build_forward_expand(gf, ggml_data.filtered_ids);
if (ggml_data.candidates != nullptr) {
res->t_candidates[seq_id] = ggml_data.candidates;
ggml_build_forward_expand(gf, ggml_data.candidates);
}
}
// TODO: Call llama_sampler_accept_ggml after all samplers have been applied.
/*
for (const auto & [seq_id, sampler] : samplers) {
if (auto it = res->t_sampled_tokens.find(seq_id); it != res->t_sampled_tokens.end()) {
if (auto it = res->t_sampled.find(seq_id); it != res->t_sampled.end()) {
ggml_tensor * selected_token = it->second;
if (selected_token != nullptr) {
llama_sampler_accept_ggml(sampler, ctx0, gf, selected_token);