mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-18 16:55:05 +02:00
Match dist between CPU and GPU
This commit is contained in:
@@ -1055,6 +1055,8 @@ extern "C" {
|
||||
//
|
||||
|
||||
// Get the backend sampled token for the ith token.
|
||||
// With multiple outputs, sampler state advances when the token is accepted,
|
||||
// not when it is read through this function.
|
||||
// Returns LLAMA_TOKEN_NULL if no token was sampled.
|
||||
LLAMA_API llama_token llama_get_sampled_token_ith(struct llama_context * ctx, int32_t i);
|
||||
|
||||
|
||||
@@ -1784,6 +1784,11 @@ int llama_context::decode(const llama_batch & batch_inp) {
|
||||
return -2;
|
||||
};
|
||||
|
||||
// start a new sampling transaction for this logical batch
|
||||
for (const auto & entry : sampling.samplers) {
|
||||
llama_sampler_backend_begin(entry.second);
|
||||
}
|
||||
|
||||
int64_t n_outputs_prev = 0;
|
||||
int64_t n_tokens_prev = 0;
|
||||
|
||||
|
||||
+71
-8
@@ -869,6 +869,7 @@ llama_token llama_sampler_sample(struct llama_sampler * smpl, struct llama_conte
|
||||
// If a backend sampler has already sampled a token, return it.
|
||||
if (sampled_token != LLAMA_TOKEN_NULL) {
|
||||
LLAMA_LOG_DEBUG("%s: Backend sampler selected token for idx %d. Skipping CPU samplers\n", __func__, idx);
|
||||
llama_sampler_accept(smpl, sampled_token);
|
||||
return sampled_token;
|
||||
}
|
||||
|
||||
@@ -1087,6 +1088,12 @@ struct llama_sampler_dist : public llama_sampler_backend {
|
||||
|
||||
std::mt19937 rng;
|
||||
|
||||
// multi-output backend draws are committed when their tokens are accepted
|
||||
bool backend_transactional;
|
||||
std::mt19937 rng_backend;
|
||||
size_t n_backend_generated;
|
||||
size_t n_backend_accepted;
|
||||
|
||||
// inputs for the current sampling graph
|
||||
std::vector<ggml_tensor *> inp_uniforms;
|
||||
};
|
||||
@@ -1172,6 +1179,9 @@ static void llama_sampler_dist_reset(struct llama_sampler * smpl) {
|
||||
auto * ctx = (llama_sampler_dist *) smpl->ctx;
|
||||
ctx->seed_cur = get_rng_seed(ctx->seed);
|
||||
ctx->rng.seed(ctx->seed_cur);
|
||||
ctx->rng_backend = ctx->rng;
|
||||
ctx->n_backend_generated = 0;
|
||||
ctx->n_backend_accepted = 0;
|
||||
}
|
||||
|
||||
static struct llama_sampler * llama_sampler_dist_clone(const struct llama_sampler * smpl) {
|
||||
@@ -1182,7 +1192,11 @@ static struct llama_sampler * llama_sampler_dist_clone(const struct llama_sample
|
||||
{
|
||||
auto * result_ctx = (llama_sampler_dist *) result->ctx;
|
||||
|
||||
result_ctx->rng = ctx->rng;
|
||||
result_ctx->rng = ctx->rng;
|
||||
result_ctx->backend_transactional = ctx->backend_transactional;
|
||||
result_ctx->rng_backend = ctx->rng_backend;
|
||||
result_ctx->n_backend_generated = ctx->n_backend_generated;
|
||||
result_ctx->n_backend_accepted = ctx->n_backend_accepted;
|
||||
}
|
||||
|
||||
return result;
|
||||
@@ -1197,11 +1211,14 @@ static bool llama_sampler_dist_backend_init(
|
||||
ggml_backend_buffer_type_t buft,
|
||||
uint32_t n_outputs_per_seq_max) {
|
||||
auto * sctx = (llama_sampler_dist *) smpl->ctx;
|
||||
GGML_UNUSED(n_outputs_per_seq_max);
|
||||
|
||||
const bool res = llama_sampler_backend_support(smpl, buft);
|
||||
|
||||
sctx->init(res);
|
||||
sctx->backend_transactional = n_outputs_per_seq_max > 1;
|
||||
sctx->rng_backend = sctx->rng;
|
||||
sctx->n_backend_generated = 0;
|
||||
sctx->n_backend_accepted = 0;
|
||||
|
||||
return res;
|
||||
}
|
||||
@@ -1282,10 +1299,17 @@ static void llama_sampler_dist_backend_set_input(struct llama_sampler * smpl) {
|
||||
// different sequences).
|
||||
std::uniform_real_distribution<double> dist(0.0f, 1.0f);
|
||||
|
||||
auto & rng = sctx->backend_transactional ? sctx->rng_backend : sctx->rng;
|
||||
|
||||
for (auto * inp_uniform : sctx->inp_uniforms) {
|
||||
GGML_ASSERT(inp_uniform != nullptr);
|
||||
const float rnd = dist(sctx->rng);
|
||||
|
||||
const float rnd = dist(rng);
|
||||
ggml_backend_tensor_set(inp_uniform, &rnd, 0, sizeof(float));
|
||||
|
||||
if (sctx->backend_transactional) {
|
||||
++sctx->n_backend_generated;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1294,9 +1318,23 @@ static void llama_sampler_dist_backend_reset(struct llama_sampler * smpl) {
|
||||
sctx->inp_uniforms.clear();
|
||||
}
|
||||
|
||||
static void llama_sampler_dist_accept(struct llama_sampler * smpl, llama_token token) {
|
||||
GGML_UNUSED(token);
|
||||
|
||||
auto * sctx = (llama_sampler_dist *) smpl->ctx;
|
||||
|
||||
if (!sctx->backend_transactional || sctx->n_backend_accepted >= sctx->n_backend_generated) {
|
||||
return;
|
||||
}
|
||||
|
||||
std::uniform_real_distribution<double> dist(0.0f, 1.0f);
|
||||
dist(sctx->rng);
|
||||
++sctx->n_backend_accepted;
|
||||
}
|
||||
|
||||
static struct llama_sampler_i llama_sampler_dist_i = {
|
||||
/* .name = */ llama_sampler_dist_name,
|
||||
/* .accept = */ nullptr,
|
||||
/* .accept = */ llama_sampler_dist_accept,
|
||||
/* .apply = */ llama_sampler_dist_apply,
|
||||
/* .reset = */ llama_sampler_dist_reset,
|
||||
/* .clone = */ llama_sampler_dist_clone,
|
||||
@@ -1314,14 +1352,39 @@ struct llama_sampler * llama_sampler_init_dist(uint32_t seed) {
|
||||
/* .iface = */ &llama_sampler_dist_i,
|
||||
/* .ctx = */ new llama_sampler_dist {
|
||||
("dist"),
|
||||
/* .seed = */ seed,
|
||||
/* .seed_cur = */ seed_cur,
|
||||
/* .rng = */ std::mt19937(seed_cur),
|
||||
/* .inp_uniforms = */ {},
|
||||
/* .seed = */ seed,
|
||||
/* .seed_cur = */ seed_cur,
|
||||
/* .rng = */ std::mt19937(seed_cur),
|
||||
/* .backend_transactional = */ false,
|
||||
/* .rng_backend = */ std::mt19937(seed_cur),
|
||||
/* .n_backend_generated = */ 0,
|
||||
/* .n_backend_accepted = */ 0,
|
||||
/* .inp_uniforms = */ {},
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
void llama_sampler_backend_begin(llama_sampler * sampler) {
|
||||
GGML_ASSERT(sampler != nullptr);
|
||||
|
||||
if (sampler->iface == &llama_sampler_chain_i) {
|
||||
auto * chain = (llama_sampler_chain *) sampler->ctx;
|
||||
for (auto & entry : chain->samplers) {
|
||||
if (!entry.is_backend) {
|
||||
break;
|
||||
}
|
||||
llama_sampler_backend_begin(entry.ptr);
|
||||
}
|
||||
} else if (sampler->iface == &llama_sampler_dist_i) {
|
||||
auto * ctx = (llama_sampler_dist *) sampler->ctx;
|
||||
if (ctx->backend_transactional) {
|
||||
ctx->rng_backend = ctx->rng;
|
||||
ctx->n_backend_generated = 0;
|
||||
ctx->n_backend_accepted = 0;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// top-k
|
||||
|
||||
struct llama_sampler_top_k : public llama_sampler_backend {
|
||||
|
||||
@@ -36,6 +36,7 @@ struct llama_sampler_chain {
|
||||
};
|
||||
|
||||
uint32_t llama_sampler_backend_n_nodes(const llama_sampler * sampler);
|
||||
void llama_sampler_backend_begin(llama_sampler * sampler);
|
||||
|
||||
struct llama_sampler * llama_sampler_init_dry_testing(
|
||||
float dry_multiplier,
|
||||
|
||||
@@ -1621,51 +1621,6 @@ static void test_backend_multi_output_limit(const test_params & params) {
|
||||
printf("backend multi-output limit test PASSED\n");
|
||||
}
|
||||
|
||||
// greedy is a stateless terminal selector; verify multi-output backend argmax
|
||||
// matches the per-row argmax of the reference logits.
|
||||
static void test_backend_multi_output_greedy(const test_params & params) {
|
||||
const llama_seq_id seq_id = 0;
|
||||
const llama_vocab * vocab = llama_model_get_vocab(params.model.get());
|
||||
const int32_t n_vocab = llama_vocab_n_tokens(vocab);
|
||||
|
||||
llama_sampler_ptr chain(llama_sampler_chain_init(llama_sampler_chain_default_params()));
|
||||
llama_sampler_chain_add(chain.get(), llama_sampler_init_greedy());
|
||||
std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
|
||||
test_context test_ctx(params, configs, 4, 4, 0, 4);
|
||||
|
||||
std::vector<llama_sampler_seq_config> reference_configs;
|
||||
test_context reference_ctx(params, reference_configs, 1, 4);
|
||||
|
||||
llama_batch batch = llama_batch_init(4, 0, 1);
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
common_batch_add(batch, llama_vocab_bos(vocab), i, { seq_id }, true);
|
||||
}
|
||||
|
||||
GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0 &&
|
||||
"multi-output backend greedy sampling should succeed");
|
||||
GGML_ASSERT(llama_decode(reference_ctx.ctx.get(), batch) == 0);
|
||||
|
||||
for (int i = 0; i < batch.n_tokens; ++i) {
|
||||
const llama_token backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), i);
|
||||
GGML_ASSERT(backend_token >= 0 && backend_token < n_vocab);
|
||||
|
||||
const float * logits = llama_get_logits_ith(reference_ctx.ctx.get(), i);
|
||||
GGML_ASSERT(logits != nullptr);
|
||||
llama_token argmax = 0;
|
||||
for (llama_token t = 1; t < n_vocab; ++t) {
|
||||
if (logits[t] > logits[argmax]) {
|
||||
argmax = t;
|
||||
}
|
||||
}
|
||||
printf("row %d: backend greedy=%d argmax=%d\n", i, backend_token, argmax);
|
||||
GGML_ASSERT(backend_token == argmax);
|
||||
}
|
||||
|
||||
llama_batch_free(batch);
|
||||
|
||||
printf("backend multi-output greedy test PASSED\n");
|
||||
}
|
||||
|
||||
static void test_backend_multi_sequence_multi_output_dist(const test_params & params) {
|
||||
const llama_vocab * vocab = llama_model_get_vocab(params.model.get());
|
||||
const int32_t n_vocab = llama_vocab_n_tokens(vocab);
|
||||
@@ -1714,7 +1669,8 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
|
||||
GGML_ASSERT(seq_id == 0 || seq_id == 1);
|
||||
outputs_per_seq[seq_id]++;
|
||||
|
||||
const llama_token backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), i);
|
||||
llama_sampler * chain = seq_id == 0 ? chain_0.get() : chain_1.get();
|
||||
const llama_token backend_token = llama_sampler_sample(chain, test_ctx.ctx.get(), i);
|
||||
const float * sampled_logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), i);
|
||||
const float * sampled_probs = llama_get_sampled_probs_ith(test_ctx.ctx.get(), i);
|
||||
const uint32_t n_logits = llama_get_sampled_logits_count_ith(test_ctx.ctx.get(), i);
|
||||
@@ -1759,6 +1715,87 @@ static void test_backend_multi_sequence_multi_output_dist(const test_params & pa
|
||||
printf("backend multi-sequence multi-output dist test PASSED\n");
|
||||
}
|
||||
|
||||
static void test_backend_multi_output_dist_transaction(const test_params & params) {
|
||||
const llama_seq_id seq_id = 0;
|
||||
const uint32_t seed = 95;
|
||||
const llama_vocab * vocab = llama_model_get_vocab(params.model.get());
|
||||
|
||||
llama_sampler_ptr chain(llama_sampler_chain_init(llama_sampler_chain_default_params()));
|
||||
llama_sampler_chain_add(chain.get(), llama_sampler_init_temp(10.0f));
|
||||
llama_sampler_chain_add(chain.get(), llama_sampler_init_dist(seed));
|
||||
std::vector<llama_sampler_seq_config> configs = {{ seq_id, chain.get() }};
|
||||
test_context test_ctx(params, configs, 1, 3, 2, 3);
|
||||
|
||||
auto verify_random = [&](int32_t row, float rnd, bool accept = true) {
|
||||
const llama_token token = accept ?
|
||||
llama_sampler_sample(chain.get(), test_ctx.ctx.get(), row) :
|
||||
llama_get_sampled_token_ith(test_ctx.ctx.get(), row);
|
||||
const float * probs = llama_get_sampled_probs_ith(test_ctx.ctx.get(), row);
|
||||
|
||||
GGML_ASSERT(token >= 0 && token < llama_vocab_n_tokens(vocab));
|
||||
GGML_ASSERT(probs != nullptr);
|
||||
|
||||
float cumsum_before = 0.0f;
|
||||
for (llama_token i = 0; i < token; ++i) {
|
||||
cumsum_before += probs[i];
|
||||
}
|
||||
|
||||
const float cumsum_sampled = cumsum_before + probs[token];
|
||||
GGML_ASSERT(rnd >= cumsum_before - 1e-4f);
|
||||
GGML_ASSERT(rnd <= cumsum_sampled + 1e-4f);
|
||||
};
|
||||
|
||||
std::mt19937 rng(seed);
|
||||
std::uniform_real_distribution<double> dist(0.0, 1.0);
|
||||
float randoms[7];
|
||||
for (float & rnd : randoms) {
|
||||
rnd = dist(rng);
|
||||
}
|
||||
|
||||
int32_t pos = 0;
|
||||
auto decode = [&]() {
|
||||
llama_batch batch = llama_batch_init(3, 0, 1);
|
||||
for (int32_t i = 0; i < 3; ++i) {
|
||||
common_batch_add(batch, llama_vocab_bos(vocab), pos++, { seq_id }, true);
|
||||
}
|
||||
GGML_ASSERT(llama_decode(test_ctx.ctx.get(), batch) == 0);
|
||||
return batch;
|
||||
};
|
||||
|
||||
llama_batch batch = decode();
|
||||
verify_random(0, randoms[0], false);
|
||||
llama_batch_free(batch);
|
||||
|
||||
batch = decode();
|
||||
verify_random(0, randoms[0]);
|
||||
verify_random(1, randoms[1]);
|
||||
llama_batch_free(batch);
|
||||
|
||||
batch = decode();
|
||||
verify_random(0, randoms[2]);
|
||||
verify_random(1, randoms[3]);
|
||||
verify_random(2, randoms[4]);
|
||||
llama_batch_free(batch);
|
||||
|
||||
batch = decode();
|
||||
verify_random(0, randoms[5]);
|
||||
llama_batch_free(batch);
|
||||
|
||||
batch = decode();
|
||||
llama_sampler_ptr saved(llama_sampler_clone(chain.get()));
|
||||
verify_random(0, randoms[6]);
|
||||
llama_batch_free(batch);
|
||||
|
||||
GGML_ASSERT(llama_set_sampler(test_ctx.ctx.get(), seq_id, saved.get()));
|
||||
chain = std::move(saved);
|
||||
|
||||
batch = decode();
|
||||
verify_random(0, randoms[6]);
|
||||
llama_batch_free(batch);
|
||||
|
||||
printf("backend multi-output dist transaction test PASSED\n");
|
||||
}
|
||||
|
||||
static void test_backend_multi_output_sampling_chain(const test_params & params) {
|
||||
const llama_seq_id seq_id = 0;
|
||||
const int32_t seed = 88;
|
||||
@@ -1806,7 +1843,7 @@ static void test_backend_multi_output_sampling_chain(const test_params & params)
|
||||
GGML_ASSERT(llama_decode(reference_ctx.ctx.get(), batch) == 0);
|
||||
|
||||
for (int i = 0; i < batch.n_tokens; ++i) {
|
||||
const llama_token backend_token = llama_get_sampled_token_ith(test_ctx.ctx.get(), i);
|
||||
const llama_token backend_token = llama_sampler_sample(chain.get(), test_ctx.ctx.get(), i);
|
||||
const float * sampled_logits = llama_get_sampled_logits_ith(test_ctx.ctx.get(), i);
|
||||
const float * sampled_probs = llama_get_sampled_probs_ith(test_ctx.ctx.get(), i);
|
||||
const llama_token * sampled_candidates = llama_get_sampled_candidates_ith(test_ctx.ctx.get(), i);
|
||||
@@ -1979,8 +2016,8 @@ static const backend_test_case BACKEND_TESTS[] = {
|
||||
{ "dist_and_cpu", test_backend_dist_sampling_and_cpu, true },
|
||||
{ "set_sampler", test_backend_set_sampler, true },
|
||||
{ "multi_output_limit", test_backend_multi_output_limit, true },
|
||||
{ "multi_output_greedy", test_backend_multi_output_greedy, true },
|
||||
{ "multi_sequence_multi_output_dist", test_backend_multi_sequence_multi_output_dist, true },
|
||||
{ "multi_output_dist_transaction", test_backend_multi_output_dist_transaction, true },
|
||||
{ "multi_output_sampling_chain", test_backend_multi_output_sampling_chain, true },
|
||||
{ "multi_output_cpu", test_backend_multi_output_cpu_suffix, true },
|
||||
{ "mixed", test_backend_mixed_sampling, true },
|
||||
|
||||
@@ -222,7 +222,6 @@ struct server_slot {
|
||||
std::vector<int32_t> spec_i_batch;
|
||||
common_prompt_checkpoint spec_ckpt;
|
||||
bool spec_is_replay = false;
|
||||
common_sampler_ptr spec_smpl_save;
|
||||
|
||||
// TODO: move members that belong to the task (such as `generated_text`, `has_new_line`) to task_results_state
|
||||
// see https://github.com/ggml-org/llama.cpp/pull/18283#issuecomment-3710175837
|
||||
@@ -360,7 +359,6 @@ struct server_slot {
|
||||
spec_draft.clear();
|
||||
spec_i_batch.clear();
|
||||
spec_ckpt.clear();
|
||||
spec_smpl_save.reset();
|
||||
}
|
||||
generated_tokens.clear();
|
||||
generated_token_probs.clear();
|
||||
@@ -3110,11 +3108,6 @@ private:
|
||||
|
||||
// update the batch with the sampled/drafted tokens
|
||||
iterate(generating, [&](server_slot & slot) {
|
||||
GGML_ASSERT(!slot.spec_smpl_save);
|
||||
if (!slot.spec_draft.empty()) {
|
||||
// backend sampling advances the sampler during llama_decode()
|
||||
slot.spec_smpl_save.reset(common_sampler_clone(slot.smpl.get()));
|
||||
}
|
||||
slot.handle_last_sampled_token(batch);
|
||||
});
|
||||
|
||||
@@ -3894,8 +3887,7 @@ private:
|
||||
|
||||
// verify and try to accept the draft
|
||||
{
|
||||
GGML_ASSERT(slot.spec_smpl_save);
|
||||
common_sampler_ptr smpl_save = std::move(slot.spec_smpl_save);
|
||||
common_sampler_ptr smpl_save(common_sampler_clone(slot.smpl.get()));
|
||||
|
||||
GGML_ASSERT(slot.spec_i_batch.size() == n_draft + 1);
|
||||
auto accepted = common_sampler_sample_and_accept_n(slot.smpl.get(), slot.ctx_tgt, slot.spec_i_batch, slot.spec_draft);
|
||||
@@ -3933,16 +3925,12 @@ private:
|
||||
slot.mem.seq_rm(slot.id, ckpt.pos_max + 1, -1);
|
||||
|
||||
slot.prompt.tokens.keep_first(ckpt.n_tokens);
|
||||
const bool restore_backend_sampler = slot.backend_sampling;
|
||||
if (restore_backend_sampler) {
|
||||
llama_set_sampler(slot.ctx_tgt, slot.id, nullptr);
|
||||
if (slot.backend_sampling) {
|
||||
slot.backend_sampling = llama_set_sampler(
|
||||
slot.ctx_tgt, slot.id, common_sampler_get(smpl_save.get()));
|
||||
}
|
||||
|
||||
slot.smpl = std::move(smpl_save);
|
||||
if (restore_backend_sampler) {
|
||||
slot.backend_sampling = llama_set_sampler(
|
||||
slot.ctx_tgt, slot.id, common_sampler_get(slot.smpl.get()));
|
||||
}
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -27,8 +27,8 @@ def test_with_and_without_draft():
|
||||
global server
|
||||
request = {
|
||||
"prompt": "I believe the meaning of life is",
|
||||
"temperature": 0.0,
|
||||
"top_k": 1,
|
||||
"temperature": 0.8,
|
||||
"top_k": 40,
|
||||
"seed": 4242,
|
||||
"n_predict": 16,
|
||||
"return_tokens": True,
|
||||
|
||||
Reference in New Issue
Block a user