Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes479downloads
test-reasoning-budget.cpp238 linesDownload Raw Back to tests
1#include "reasoning-budget.h"2#include "unicode.h"3 4#include "llama.h"5#include "ggml.h"6 7#ifdef NDEBUG8#undef NDEBUG9#endif10 11#include <cmath>12#include <cstddef>13#include <cstdio>14#include <string>15#include <vector>16 17// Reasoning budget sampler test helper18// These tests use nullptr vocab which safely falls back to treating all tokens as complete19// (The UTF-8 boundary detection logic is tested separately in test_utf8_boundary_detection)20static void test_reasoning_budget(21    const char * test_name,22    const std::vector<llama_token> & sequence,23    const std::vector<llama_token> & start_tokens,24    const std::vector<llama_token> & end_tokens,25    const std::vector<llama_token> & forced_tokens,26    int32_t budget,27    common_reasoning_budget_state initial_state,28    size_t expected_force_start,   // token index where forcing should start (SIZE_MAX = never)29    size_t expected_force_end      // token index where forcing should end (after this, no more forcing)30) {31    // Find the maximum token ID to ensure our vocab covers all tokens32    llama_token max_token = 0;33    for (auto t : sequence) max_token = std::max(max_token, t);34    for (auto t : start_tokens) max_token = std::max(max_token, t);35    for (auto t : end_tokens) max_token = std::max(max_token, t);36    for (auto t : forced_tokens) max_token = std::max(max_token, t);37 38    // Create a minimal sampler with mock vocabulary39    // For this test, we use nullptr as vocab since we're testing state transitions40    // The UTF-8 boundary check will treat all tokens as complete (safe fallback)41    auto * sampler = common_reasoning_budget_init(42        nullptr,  // vocab - not used for basic state machine tests43        start_tokens,44        end_tokens,45        forced_tokens,46        budget,47        initial_state48    );49 50    // Create a test token data array for checking forcing behavior51    // Vocab size must be large enough to include all tokens (start, end, forced, sequence)52    std::vector<llama_token_data> cur;53    const size_t n_vocab = (size_t)max_token + 1;54    for (size_t i = 0; i < n_vocab; i++) {55        cur.emplace_back(llama_token_data{(llama_token)i, logf((float)(i+1)), 0.0f});56    }57    llama_token_data_array cur_p = { cur.data(), cur.size(), -1, false };58 59    size_t actual_force_start = SIZE_MAX;60    size_t actual_force_end = SIZE_MAX;61 62    // Feed the sequence and track when forcing occurs63    for (size_t i = 0; i < sequence.size(); i++) {64        // Check if we're in forcing state by applying and seeing if logits are modified65        cur_p.selected = -1;66        for (size_t j = 0; j < cur.size(); j++) {67            cur[j].logit = logf((float)(j+1));  // reset logits68        }69 70        llama_sampler_apply(sampler, &cur_p);71 72        // Check if forcing is active (all logits except one should be -INFINITY)73        size_t finite_count = 0;74        llama_token finite_token = -1;75        for (size_t j = 0; j < cur.size(); j++) {76            if (std::isfinite(cur[j].logit)) {77                finite_count++;78                finite_token = cur[j].id;79            }80        }81 82        llama_sampler_accept(sampler, sequence[i]);83 84        fprintf(stderr, "    i=%zu: token=%d, finite_count=%zu, finite_token=%d\n", i, (int)sequence[i], finite_count, (int)finite_token);85 86        if (finite_count == 1) {87            if (actual_force_start == SIZE_MAX) {88                actual_force_start = i;89            }90            actual_force_end = i;91        } else if (actual_force_start != SIZE_MAX && actual_force_end != SIZE_MAX) {92            // Forcing stopped93            break;94        }95    }96 97    llama_sampler_free(sampler);98 99    // Verify forcing occurred at expected positions100    if (expected_force_start == SIZE_MAX) {101        if (actual_force_start != SIZE_MAX) {102            fprintf(stderr, "Test '%s' FAILED: Expected no forcing, but forcing occurred at %zu\n", test_name, actual_force_start);103            GGML_ASSERT(false && "Expected no forcing, but forcing occurred");104        }105    } else {106        if (actual_force_start == SIZE_MAX) {107            fprintf(stderr, "Test '%s' FAILED: Expected forcing but none occurred\n", test_name);108            GGML_ASSERT(false && "Expected forcing but none occurred");109        }110        if (actual_force_start != expected_force_start) {111            fprintf(stderr, "Test '%s' FAILED: Forcing started at %zu, expected %zu\n", test_name, actual_force_start, expected_force_start);112            GGML_ASSERT(false && "Forcing started at wrong position");113        }114    }115 116    if (expected_force_end != SIZE_MAX) {117        if (actual_force_end < expected_force_end) {118            fprintf(stderr, "Test '%s' FAILED: Forcing ended at %zu, expected >= %zu\n", test_name, actual_force_end, expected_force_end);119            GGML_ASSERT(false && "Forcing ended too early");120        }121    }122 123    fprintf(stderr, "  Test '%s' passed (force_start=%zu, force_end=%zu)\n", test_name, actual_force_start, actual_force_end);124    (void)sequence;125}126 127// UTF-8 boundary detection unit test128// Tests common_utf8_is_complete() from reasoning-budget.h129static void test_utf8_boundary_detection() {130    // Complete sequences131    GGML_ASSERT(common_utf8_is_complete("hello"));132    GGML_ASSERT(common_utf8_is_complete(""));133    GGML_ASSERT(common_utf8_is_complete("\xC2\xA0"));            // complete 2-byte UTF-8 (U+00A0)134    GGML_ASSERT(common_utf8_is_complete("\xE2\x80\x9C"));        // complete 3-byte UTF-8 (left double quote)135    GGML_ASSERT(common_utf8_is_complete("\xF0\x9F\x98\x80"));    // complete 4-byte UTF-8 (emoji)136    GGML_ASSERT(common_utf8_is_complete("abc\xC3\xA9"));         // ASCII + complete 2-byte137 138    // Incomplete sequences139    GGML_ASSERT(!common_utf8_is_complete(std::string("\xC2", 1)));            // 2-byte start, missing continuation140    GGML_ASSERT(!common_utf8_is_complete(std::string("\xE2\x80", 2)));        // 3-byte start + 1 cont, missing 1141    GGML_ASSERT(!common_utf8_is_complete(std::string("\xE2", 1)));            // 3-byte start, missing 2142    GGML_ASSERT(!common_utf8_is_complete(std::string("\xF0\x9F\x98", 3)));    // 4-byte start + 2 cont, missing 1143    GGML_ASSERT(!common_utf8_is_complete(std::string("\xF0\x9F", 2)));        // 4-byte start + 1 cont, missing 2144    GGML_ASSERT(!common_utf8_is_complete(std::string("\xF0", 1)));            // 4-byte start, missing 3145    GGML_ASSERT(!common_utf8_is_complete(std::string("\x80", 1)));            // orphan continuation byte146 147    // Mixed: ASCII followed by start of multi-byte148    GGML_ASSERT(!common_utf8_is_complete(std::string("hello\xC3", 6)));       // ASCII + incomplete 2-byte149    GGML_ASSERT(common_utf8_is_complete(std::string("hello\xC3\xA9", 7)));    // ASCII + complete 2-byte150}151 152int main(void) {153    // Reasoning budget sampler tests154    printf("Testing reasoning budget sampler... ");155 156    // Test 1: Basic budget with start/end tokens - no forcing (natural end before budget exhausted)157    {158        const std::vector<llama_token> start = {100};  // start token159        const std::vector<llama_token> end = {101};    // end token160        const std::vector<llama_token> forced = {102}; // forced token (not used in this test)161        const std::vector<llama_token> sequence = {100, 50, 51, 101, 52}; // start, two tokens, end, one more162 163        test_reasoning_budget("natural end before budget exhausted", sequence, start, end, forced,164            5,      // budget of 5 tokens165            REASONING_BUDGET_IDLE,166            SIZE_MAX, SIZE_MAX); // no forcing expected (natural end)167    }168 169    // Test 2: Budget exhausted, forcing should occur170    // Flow: i=0 apply()->passthrough, accept(100)->COUNTING; i=1 accept(50)->remaining=1171    // i=2 accept(51)->remaining=0->FORCING; i=3 apply() forces token[0]; i=4 apply() forces token[1]172    // At i=4, accept() advances force_pos to 2 which equals forced_tokens.size(), so state becomes DONE173    {174        const std::vector<llama_token> start = {100};175        const std::vector<llama_token> end = {101};176        const std::vector<llama_token> forced = {102, 101}; // forced message + end177        const std::vector<llama_token> sequence = {100, 50, 51, 52, 53}; // start + 4 tokens (budget=2)178 179        test_reasoning_budget("budget exhausted forcing", sequence, start, end, forced,180            2,      // budget of 2 tokens181            REASONING_BUDGET_IDLE,182            3,      // forcing starts at i=3 (accept at i=2 depletes budget, apply at i=3 forces)183            4);     // forcing continues through i=4 (accept at i=4 transitions to DONE)184    }185 186    // Test 3: Activate immediately with budget=0, forcing should start right away187    // Flow: init promotes COUNTING+budget=0 to FORCING, so apply() sees FORCING at i=0188    {189        const std::vector<llama_token> start = {100};190        const std::vector<llama_token> end = {101};191        const std::vector<llama_token> forced = {102, 101};192        const std::vector<llama_token> sequence = {100, 50, 51, 52}; // start token first, then 3 tokens193 194        test_reasoning_budget("activate immediately budget=0", sequence, start, end, forced,195            0,      // budget of 0 tokens196            REASONING_BUDGET_COUNTING, // starts counting, promoted to FORCING since budget=0197            0,      // forcing starts at i=0 (initialized in FORCING, apply forces immediately)198            1);     // forcing continues through i=1 (accept at i=1 transitions to DONE)199    }200 201    // Test 4: No start/end tokens configured - passthrough (no forcing)202    {203        const std::vector<llama_token> start = {};204        const std::vector<llama_token> end = {};205        const std::vector<llama_token> forced = {102};206        const std::vector<llama_token> sequence = {50, 51, 52, 53};207 208        test_reasoning_budget("no start/end configured", sequence, start, end, forced,209            2,      // budget210            REASONING_BUDGET_IDLE,211            SIZE_MAX, SIZE_MAX); // no forcing (no start/end configured)212    }213 214    // Test 5: Activate immediately with budget > 0, count down then force215    // Flow: i=0 accept(50)->remaining=1, i=1 accept(51)->remaining=0->FORCING216    // Forcing starts at i=2 (apply sees FORCING after accept at i=1 transitioned)217    {218        const std::vector<llama_token> start = {100};219        const std::vector<llama_token> end = {101};220        const std::vector<llama_token> forced = {102, 101};221        const std::vector<llama_token> sequence = {50, 51, 52, 53};222 223        test_reasoning_budget("activate immediately with budget", sequence, start, end, forced,224            2,      // budget of 2 tokens225            REASONING_BUDGET_COUNTING,226            2,      // forcing starts at i=2 (after 2 accepts deplete budget, apply at i=2 forces)227            3);     // forcing continues through i=3228    }229 230    printf("OK (5 tests passed)\n");231 232    printf("Testing UTF-8 boundary detection... ");233    test_utf8_boundary_detection();234    printf("OK\n");235 236    return 0;237}238