mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-08-30 17:10:50 +02:00
spec: Add benchmark-only synthetic speculative acceptance options (#27711)
* Add benchmark-only synthetic speculative acceptance to llama-server and llama-cli * Address review comments * Address review comments * Add some comments in the code
This commit is contained in:
@@ -4,6 +4,7 @@
|
||||
#include "llama.h"
|
||||
#include "speculative.h"
|
||||
|
||||
#include <cmath>
|
||||
#include <limits>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
@@ -34,6 +35,62 @@ static void test(void) {
|
||||
std::numeric_limits<int32_t>::max(),
|
||||
std::numeric_limits<int32_t>::max());
|
||||
|
||||
{
|
||||
common_params_speculative spec;
|
||||
spec.synth_len = 3.4;
|
||||
|
||||
auto assert_invalid = [](const common_params_speculative & value, int32_t n_max) {
|
||||
try {
|
||||
common_speculative_synth_rates_resolve(&value, n_max);
|
||||
assert(false);
|
||||
} catch (const std::invalid_argument &) {
|
||||
}
|
||||
};
|
||||
|
||||
const auto rates = common_speculative_synth_rates_resolve(&spec, 4);
|
||||
assert(rates.size() == 4);
|
||||
assert(std::abs(rates[0] - 0.80581) < 1e-5);
|
||||
assert(std::abs(rates[1] - 0.64933) < 1e-5);
|
||||
assert(std::abs(rates[2] - 0.52323) < 1e-5);
|
||||
assert(std::abs(rates[3] - 0.42163) < 1e-5);
|
||||
assert(std::abs(1.0 + rates[0] + rates[1] + rates[2] + rates[3] - 3.4) < 1e-8);
|
||||
|
||||
spec.synth_len = 1.0;
|
||||
assert(common_speculative_synth_rates_resolve(&spec, 4) == std::vector<double>({0.0, 0.0, 0.0, 0.0}));
|
||||
|
||||
spec.synth_len = 5.0;
|
||||
assert(common_speculative_synth_rates_resolve(&spec, 4) == std::vector<double>({1.0, 1.0, 1.0, 1.0}));
|
||||
|
||||
spec.synth_len = 5.1;
|
||||
assert_invalid(spec, 4);
|
||||
|
||||
spec.synth_len = std::numeric_limits<double>::quiet_NaN();
|
||||
assert_invalid(spec, 4);
|
||||
|
||||
spec.synth_len = 0.0;
|
||||
assert_invalid(spec, 4);
|
||||
|
||||
spec.synth_len = -1.0;
|
||||
spec.synth_rates = {0.8, 0.6, 0.4};
|
||||
assert_invalid(spec, 4);
|
||||
|
||||
spec.synth_rates = {0.8, 0.6, 0.4, 0.2};
|
||||
assert(common_speculative_synth_rates_resolve(&spec, 4) == spec.synth_rates);
|
||||
|
||||
spec.synth_rates = {0.8, 0.9, 0.4, 0.2};
|
||||
assert_invalid(spec, 4);
|
||||
|
||||
spec.synth_rates = {0.8, std::numeric_limits<double>::quiet_NaN(), 0.4, 0.2};
|
||||
assert_invalid(spec, 4);
|
||||
|
||||
spec.synth_rates = {0.8, 0.6, 0.4, -0.2};
|
||||
assert_invalid(spec, 4);
|
||||
|
||||
spec.synth_rates = {0.8, 0.6, 0.4, 0.2};
|
||||
spec.synth_len = 3.0;
|
||||
assert_invalid(spec, 4);
|
||||
}
|
||||
|
||||
{
|
||||
common_params base;
|
||||
base.n_parallel = 4;
|
||||
@@ -197,6 +254,26 @@ static void test(void) {
|
||||
assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_SPECULATIVE));
|
||||
assert(params.speculative.draft.n_max == 123);
|
||||
|
||||
{
|
||||
common_params synth_params;
|
||||
argv = {"binary_name", "--spec-synth-len", "3.4"};
|
||||
assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), synth_params, LLAMA_EXAMPLE_SERVER));
|
||||
assert(synth_params.speculative.synth_len == 3.4);
|
||||
}
|
||||
|
||||
{
|
||||
common_params synth_params;
|
||||
argv = {"binary_name", "--spec-synth-rates", "0.8,0.6,0.2"};
|
||||
assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), synth_params, LLAMA_EXAMPLE_SERVER));
|
||||
assert(synth_params.speculative.synth_rates == std::vector<double>({0.8, 0.6, 0.2}));
|
||||
}
|
||||
|
||||
{
|
||||
common_params synth_params;
|
||||
argv = {"binary_name", "--spec-synth-len", "3.4x"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), synth_params, LLAMA_EXAMPLE_SERVER));
|
||||
}
|
||||
|
||||
argv = {"binary_name", "-lm", "none"};
|
||||
assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));
|
||||
assert(params.load_mode == LLAMA_LOAD_MODE_NONE);
|
||||
|
||||
Reference in New Issue
Block a user