sampling : check backend support during init

This commit is contained in:
Georgi Gerganov
2025-12-04 17:29:08 +02:00
parent 1bde70785d
commit 6958d41366
8 changed files with 369 additions and 178 deletions

View File

@@ -106,7 +106,6 @@ struct common_sampler {
struct llama_sampler * grmr;
struct llama_sampler * chain;
struct llama_sampler * chain_backend;
ring_buffer<llama_token> prev;
@@ -119,7 +118,6 @@ struct common_sampler {
llama_sampler_reset(grmr);
llama_sampler_reset(chain);
llama_sampler_reset(chain_backend);
}
void set_logits(struct llama_context * ctx, int idx) {
@@ -247,13 +245,12 @@ struct common_sampler * common_sampler_init(const struct llama_model * model, co
}
auto * result = new common_sampler {
/* .params = */ params,
/* .grmr = */ grmr,
/* .chain = */ llama_sampler_chain_init(lparams),
/* .chain_backend = */ llama_sampler_chain_init(lparams),
/* .prev = */ ring_buffer<llama_token>(std::max(32, params.n_prev)),
/* .cur = */ {},
/* .cur_p = */ {},
/* .params = */ params,
/* .grmr = */ grmr,
/* .chain = */ llama_sampler_chain_init(lparams),
/* .prev = */ ring_buffer<llama_token>(std::max(32, params.n_prev)),
/* .cur = */ {},
/* .cur_p = */ {},
};
std::vector<llama_sampler *> samplers;
@@ -318,15 +315,8 @@ struct common_sampler * common_sampler_init(const struct llama_model * model, co
GGML_ASSERT(false && "unknown mirostat version");
}
bool is_backend = params.backend_sampling;
// split in two chains: backend -> CPU
for (auto * smpl : samplers) {
if (!smpl->iface->backend_apply) {
is_backend = false;
}
llama_sampler_chain_add(is_backend ? result->chain_backend : result->chain, smpl);
llama_sampler_chain_add(result->chain, smpl);
}
return result;
@@ -336,7 +326,6 @@ void common_sampler_free(struct common_sampler * gsmpl) {
if (gsmpl) {
llama_sampler_free(gsmpl->grmr);
llama_sampler_free(gsmpl->chain);
llama_sampler_free(gsmpl->chain_backend);
delete gsmpl;
}
@@ -360,13 +349,12 @@ void common_sampler_reset(struct common_sampler * gsmpl) {
struct common_sampler * common_sampler_clone(common_sampler * gsmpl) {
return new common_sampler {
/* .params = */ gsmpl->params,
/* .grmr = */ llama_sampler_clone(gsmpl->grmr),
/* .chain = */ llama_sampler_clone(gsmpl->chain),
/* .chain_backend = */ llama_sampler_clone(gsmpl->chain_backend),
/* .prev = */ gsmpl->prev,
/* .cur = */ gsmpl->cur,
/* .cur_p = */ gsmpl->cur_p,
/* .params = */ gsmpl->params,
/* .grmr = */ llama_sampler_clone(gsmpl->grmr),
/* .chain = */ llama_sampler_clone(gsmpl->chain),
/* .prev = */ gsmpl->prev,
/* .cur = */ gsmpl->cur,
/* .cur_p = */ gsmpl->cur_p,
};
}
@@ -415,8 +403,8 @@ void common_perf_print(const struct llama_context * ctx, const struct common_sam
}
}
struct llama_sampler * common_sampler_chain_backend(const struct common_sampler * gsmpl) {
return gsmpl->chain_backend;
struct llama_sampler * common_sampler_get(const struct common_sampler * gsmpl) {
return gsmpl->chain;
}
llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_context * ctx, int idx, bool grammar_first) {
@@ -424,11 +412,13 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co
// return that token id directly.
{
const llama_token id = llama_get_sampled_token_ith(ctx, idx);
if (id != LLAMA_TOKEN_NULL) {
LOG_DBG("%s: Backend sampler selected token: '%d'. Will not run any CPU samplers\n", __func__, id);
return id;
}
}
llama_synchronize(ctx);
// start measuring sampling time after the llama_context synchronization in order to not measure any ongoing async operations
@@ -556,16 +546,12 @@ llama_token common_sampler_last(const struct common_sampler * gsmpl) {
}
std::string common_sampler_print(const struct common_sampler * gsmpl) {
std::string result = llama_sampler_chain_n(gsmpl->chain_backend) > 0 ? "*logits " : "logits ";
for (int i = 0; i < llama_sampler_chain_n(gsmpl->chain_backend); i++) {
const auto * smpl = llama_sampler_chain_get(gsmpl->chain_backend, i);
result += std::string("-> *") + llama_sampler_name(smpl) + " ";
}
std::string result = "logits ";
for (int i = 0; i < llama_sampler_chain_n(gsmpl->chain); i++) {
const auto * smpl = llama_sampler_chain_get(gsmpl->chain, i);
result += std::string("-> ") + llama_sampler_name(smpl) + " ";
result += std::string("-> ");
result += std::string(llama_sampler_name(smpl)) + " ";
}
return result;