Match dist between CPU and GPU

This commit is contained in:
Gaurav Garg
2026-07-07 20:54:05 +05:30
parent 813c39b225
commit a24097dc57
7 changed files with 170 additions and 74 deletions
+2
View File
@@ -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);
+5
View File
@@ -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
View File
@@ -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 {
+1
View File
@@ -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,
+85 -48
View File
@@ -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 },
+4 -16
View File
@@ -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;
}
+2 -2
View File
@@ -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,