diff --git a/include/llama.h b/include/llama.h index a14498925f..2fe09bc615 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1310,6 +1310,8 @@ extern "C" { LLAMA_API void llama_sampler_apply ( struct llama_sampler * smpl, llama_token_data_array * cur_p); LLAMA_API void llama_sampler_reset ( struct llama_sampler * smpl); LLAMA_API struct llama_sampler * llama_sampler_clone (const struct llama_sampler * smpl); + // copy the state of src into dst; both must be samplers of the same type + LLAMA_API void llama_sampler_copy ( struct llama_sampler * dst, const struct llama_sampler * src); // important: do not free if the sampler has been added to a llama_sampler_chain (via llama_sampler_chain_add) LLAMA_API void llama_sampler_free ( struct llama_sampler * smpl); diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp index e550fbe4ae..7cad9837c6 100644 --- a/src/llama-sampler.cpp +++ b/src/llama-sampler.cpp @@ -417,6 +417,27 @@ struct llama_sampler * llama_sampler_clone(const struct llama_sampler * smpl) { GGML_ABORT("the sampler does not support cloning"); } +void llama_sampler_copy(struct llama_sampler * dst, const struct llama_sampler * src) { + if (!dst || !src) { + return; + } + + GGML_ASSERT(dst->iface == src->iface && "llama_sampler_copy: cannot copy between different sampler types"); + + // build a temporary sampler carrying src's current state + llama_sampler * tmp = llama_sampler_clone(src); + + // free dst's old state (frees dst->ctx, including children for a chain) + if (dst->iface->free) { + dst->iface->free(dst); + } + + // transplant tmp's state into dst, then destroy the (now empty) temp shell + dst->ctx = tmp->ctx; + tmp->ctx = nullptr; + delete tmp; +} + void llama_sampler_free(struct llama_sampler * smpl) { if (smpl == nullptr) { return;