diff --git a/include/llama.h b/include/llama.h index 75dc50652e..35cab4d39d 100644 --- a/include/llama.h +++ b/include/llama.h @@ -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); diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 170f1a5e58..bda3554d65 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -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; diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp index e54cdcff07..690f768ca4 100644 --- a/src/llama-sampler.cpp +++ b/src/llama-sampler.cpp @@ -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 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 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 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 { diff --git a/src/llama-sampler.h b/src/llama-sampler.h index 9cd84dccb4..e5db2982bd 100644 --- a/src/llama-sampler.h +++ b/src/llama-sampler.h @@ -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, diff --git a/tests/test-backend-sampler.cpp b/tests/test-backend-sampler.cpp index 5a10156499..da3100f263 100644 --- a/tests/test-backend-sampler.cpp +++ b/tests/test-backend-sampler.cpp @@ -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 configs = {{ seq_id, chain.get() }}; - test_context test_ctx(params, configs, 4, 4, 0, 4); - - std::vector 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 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 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 }, diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index aa7f4037a9..e22696ea78 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -222,7 +222,6 @@ struct server_slot { std::vector 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; } diff --git a/tools/server/tests/unit/test_speculative.py b/tools/server/tests/unit/test_speculative.py index 12e391eedc..0a5b217951 100644 --- a/tools/server/tests/unit/test_speculative.py +++ b/tools/server/tests/unit/test_speculative.py @@ -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,