Rebase and address review comments

This commit is contained in:
Gaurav Garg
2026-08-06 00:27:46 +05:30
parent fea0c3d410
commit 8580cc0b81
2 changed files with 16 additions and 4 deletions
+2 -2
View File
@@ -1274,7 +1274,7 @@ extern "C" {
// [EXPERIMENTAL]
// backend sampling interface:
// return true if the backend supports all ops needed by the sampler and the requested output mode
// return true if the backend supports all ops needed by the sampler and can handle up to n_outputs_per_seq_max outputs per sequence
// note: call once per sampler
bool (*backend_init)(
struct llama_sampler * smpl,
@@ -1298,7 +1298,7 @@ extern "C" {
// called before graph execution to set inputs for the current ubatch
void (*backend_set_input)(struct llama_sampler * smpl);
// called before rebuilding a sampling graph to clear graph-owned tensor references
// called before rebuilding a sampling graph to clear any internal sampler state
void (*backend_reset)(struct llama_sampler * smpl);
};
+14 -2
View File
@@ -2950,9 +2950,15 @@ static void llama_sampler_penalties_free(struct llama_sampler * smpl) {
static bool llama_sampler_penalties_backend_init(
struct llama_sampler * smpl,
ggml_backend_buffer_type_t buft) {
ggml_backend_buffer_type_t buft,
uint32_t n_outputs_per_seq_max) {
auto * sctx = (llama_sampler_penalties *) smpl->ctx;
if (n_outputs_per_seq_max > 1) {
sctx->init(false);
return false;
}
const bool res = llama_sampler_backend_support(smpl, buft);
sctx->init(res);
@@ -3112,6 +3118,12 @@ static void llama_sampler_penalties_backend_set_input(struct llama_sampler * smp
ggml_backend_tensor_set(sctx->inp_counts, sctx->host_counts.data(), 0, sctx->n_max * sizeof(int32_t));
}
static void llama_sampler_penalties_backend_reset(struct llama_sampler * smpl) {
auto * sctx = (llama_sampler_penalties *) smpl->ctx;
sctx->inp_token_ids = nullptr;
sctx->inp_counts = nullptr;
}
static struct llama_sampler_i llama_sampler_penalties_i = {
/* .name = */ llama_sampler_penalties_name,
/* .accept = */ llama_sampler_penalties_accept,
@@ -3123,7 +3135,7 @@ static struct llama_sampler_i llama_sampler_penalties_i = {
/* .backend_accept = */ nullptr,
/* .backend_apply = */ llama_sampler_penalties_backend_apply,
/* .backend_set_input = */ llama_sampler_penalties_backend_set_input,
/* .backend_reset = */ nullptr,
/* .backend_reset = */ llama_sampler_penalties_backend_reset,
};
struct llama_sampler * llama_sampler_init_penalties(