diff --git a/include/llama.h b/include/llama.h index 5e7880dd8d..f33e9ee8c1 100644 --- a/include/llama.h +++ b/include/llama.h @@ -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); }; diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp index cc9630c832..37c6eb6d50 100644 --- a/src/llama-sampler.cpp +++ b/src/llama-sampler.cpp @@ -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(