echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0479
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 