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
+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 },