mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 10:57:33 -05:00
Rebase and address review comments
This commit is contained in:
+2
-2
@@ -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
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user