echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0479
1#include "reasoning-budget.h"2#include "common.h"3#include "unicode.h"4 5#include "log.h"6 7#include <cmath>8#include <cstdint>9#include <string>10#include <vector>11 12struct token_matcher {13 std::vector<llama_token> tokens;14 size_t pos = 0;15 16 bool advance(llama_token token) {17 if (tokens.empty()) {18 return false;19 }20 21 if (token == tokens[pos]) {22 pos++;23 if (pos >= tokens.size()) {24 pos = 0;25 return true;26 }27 } else {28 pos = 0;29 if (token == tokens[0]) {30 pos = 1;31 }32 }33 return false;34 }35 36 void reset() { pos = 0; }37};38 39struct common_reasoning_budget_ctx {40 const llama_vocab * vocab;41 42 token_matcher start_matcher;43 token_matcher end_matcher;44 std::vector<llama_token> forced_tokens;45 46 int32_t budget; // maximum tokens in reasoning block47 int32_t remaining; // tokens remaining in budget48 49 common_reasoning_budget_state state;50 51 // for forcing52 size_t force_pos; // next position in forced_tokens to force53};54 55static const char * common_reasoning_budget_name(const struct llama_sampler * /*smpl*/) {56 return "reasoning-budget";57}58 59static void common_reasoning_budget_accept(struct llama_sampler * smpl, llama_token token) {60 auto * ctx = (common_reasoning_budget_ctx *) smpl->ctx;61 62 switch (ctx->state) {63 case REASONING_BUDGET_IDLE:64 {65 if (ctx->start_matcher.advance(token)) {66 ctx->state = REASONING_BUDGET_COUNTING;67 ctx->remaining = ctx->budget;68 LOG_INF("reasoning-budget: activated, budget=%d tokens\n", ctx->budget);69 70 if (ctx->remaining <= 0) {71 ctx->state = REASONING_BUDGET_FORCING;72 ctx->force_pos = 0;73 LOG_INF("reasoning-budget: budget=0, forcing immediately\n");74 }75 }76 break;77 }78 case REASONING_BUDGET_COUNTING:79 case REASONING_BUDGET_WAITING_UTF8:80 {81 if (ctx->end_matcher.advance(token)) {82 ctx->state = REASONING_BUDGET_DONE;83 LOG_INF("reasoning-budget: deactivated (natural end)\n");84 break;85 }86 87 bool utf8_complete = true;88 if (ctx->vocab != nullptr) {89 const std::string piece = common_token_to_piece(ctx->vocab, token, false);90 utf8_complete = common_utf8_is_complete(piece);91 }92 93 if (ctx->state == REASONING_BUDGET_WAITING_UTF8) {94 if (utf8_complete) {95 ctx->state = REASONING_BUDGET_FORCING;96 ctx->force_pos = 0;97 ctx->end_matcher.reset();98 LOG_INF("reasoning-budget: UTF-8 complete, now forcing end sequence\n");99 }100 } else if (ctx->state == REASONING_BUDGET_COUNTING) {101 ctx->remaining--;102 if (ctx->remaining <= 0) {103 if (utf8_complete) {104 ctx->state = REASONING_BUDGET_FORCING;105 ctx->force_pos = 0;106 ctx->end_matcher.reset();107 LOG_INF("reasoning-budget: budget exhausted, forcing end sequence\n");108 } else {109 ctx->state = REASONING_BUDGET_WAITING_UTF8;110 ctx->end_matcher.reset();111 LOG_INF("reasoning-budget: budget exhausted, waiting for UTF-8 completion\n");112 }113 }114 }115 break;116 }117 case REASONING_BUDGET_FORCING:118 ctx->force_pos++;119 if (ctx->force_pos >= ctx->forced_tokens.size()) {120 ctx->state = REASONING_BUDGET_DONE;121 LOG_INF("reasoning-budget: forced sequence complete, done\n");122 }123 break;124 case REASONING_BUDGET_DONE:125 break;126 }127}128 129static void common_reasoning_budget_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {130 auto * ctx = (common_reasoning_budget_ctx *) smpl->ctx;131 132 if (ctx->state != REASONING_BUDGET_FORCING) {133 // passthrough — don't modify logits134 return;135 }136 137 if (ctx->force_pos >= ctx->forced_tokens.size()) {138 return;139 }140 141 const llama_token forced = ctx->forced_tokens[ctx->force_pos];142 143 // set all logits to -inf except the forced token144 for (size_t i = 0; i < cur_p->size; i++) {145 if (cur_p->data[i].id != forced) {146 cur_p->data[i].logit = -INFINITY;147 }148 }149}150 151static void common_reasoning_budget_reset(struct llama_sampler * smpl) {152 auto * ctx = (common_reasoning_budget_ctx *) smpl->ctx;153 ctx->state = REASONING_BUDGET_IDLE;154 ctx->remaining = ctx->budget;155 ctx->start_matcher.reset();156 ctx->end_matcher.reset();157 ctx->force_pos = 0;158}159 160// forward declaration for use in clone161static struct llama_sampler * common_reasoning_budget_init_state(162 const struct llama_vocab * vocab, const std::vector<llama_token> & start_tokens,163 const std::vector<llama_token> & end_tokens, const std::vector<llama_token> & forced_tokens,164 int32_t budget, common_reasoning_budget_state initial_state);165 166static struct llama_sampler * common_reasoning_budget_clone(const struct llama_sampler * smpl) {167 const auto * ctx = (const common_reasoning_budget_ctx *) smpl->ctx;168 return common_reasoning_budget_init_state(169 ctx->vocab,170 ctx->start_matcher.tokens,171 ctx->end_matcher.tokens,172 ctx->forced_tokens,173 ctx->budget,174 ctx->state);175}176 177static void common_reasoning_budget_free(struct llama_sampler * smpl) {178 delete (common_reasoning_budget_ctx *) smpl->ctx;179}180 181static struct llama_sampler_i common_reasoning_budget_i = {182 /* .name = */ common_reasoning_budget_name,183 /* .accept = */ common_reasoning_budget_accept,184 /* .apply = */ common_reasoning_budget_apply,185 /* .reset = */ common_reasoning_budget_reset,186 /* .clone = */ common_reasoning_budget_clone,187 /* .free = */ common_reasoning_budget_free,188 /* .backend_init = */ nullptr,189 /* .backend_accept = */ nullptr,190 /* .backend_apply = */ nullptr,191 /* .backend_set_input = */ nullptr,192};193 194static struct llama_sampler * common_reasoning_budget_init_state(195 const struct llama_vocab * vocab,196 const std::vector<llama_token> & start_tokens,197 const std::vector<llama_token> & end_tokens,198 const std::vector<llama_token> & forced_tokens,199 int32_t budget,200 common_reasoning_budget_state initial_state) {201 // promote COUNTING with budget <= 0 to FORCING202 if (initial_state == REASONING_BUDGET_COUNTING && budget <= 0) {203 initial_state = REASONING_BUDGET_FORCING;204 }205 206 return llama_sampler_init(207 /* .iface = */ &common_reasoning_budget_i,208 /* .ctx = */ new common_reasoning_budget_ctx {209 /* .vocab = */ vocab,210 /* .start_matcher = */ { start_tokens, 0 },211 /* .end_matcher = */ { end_tokens, 0 },212 /* .forced_tokens = */ forced_tokens,213 /* .budget = */ budget,214 /* .remaining = */ budget,215 /* .state = */ initial_state,216 /* .force_pos = */ 0,217 }218 );219}220 221struct llama_sampler * common_reasoning_budget_init(222 const struct llama_vocab * vocab,223 const std::vector<llama_token> & start_tokens,224 const std::vector<llama_token> & end_tokens,225 const std::vector<llama_token> & forced_tokens,226 int32_t budget,227 const std::vector<llama_token> & prefill_tokens) {228 // Determine initial state from prefill: COUNTING if the prefill begins with229 // the start sequence but does not also contain the end sequence after it.230 common_reasoning_budget_state initial_state = REASONING_BUDGET_IDLE;231 if (!prefill_tokens.empty() && !start_tokens.empty() &&232 prefill_tokens.size() >= start_tokens.size() &&233 std::equal(start_tokens.begin(), start_tokens.end(), prefill_tokens.begin())) {234 initial_state = REASONING_BUDGET_COUNTING;235 // If the end sequence also follows the start in the prefill, reasoning236 // was opened and immediately closed — stay IDLE.237 if (!end_tokens.empty() &&238 prefill_tokens.size() >= start_tokens.size() + end_tokens.size()) {239 auto end_start = prefill_tokens.end() - (ptrdiff_t) end_tokens.size();240 if (end_start >= prefill_tokens.begin() + (ptrdiff_t) start_tokens.size() &&241 std::equal(end_tokens.begin(), end_tokens.end(), end_start)) {242 initial_state = REASONING_BUDGET_IDLE;243 }244 }245 }246 return common_reasoning_budget_init_state(vocab, start_tokens, end_tokens, forced_tokens, budget, initial_state);247}248 249struct llama_sampler * common_reasoning_budget_init(250 const struct llama_vocab * vocab,251 const std::vector<llama_token> & start_tokens,252 const std::vector<llama_token> & end_tokens,253 const std::vector<llama_token> & forced_tokens,254 int32_t budget,255 common_reasoning_budget_state initial_state) {256 return common_reasoning_budget_init_state(vocab, start_tokens, end_tokens, forced_tokens, budget, initial_state);257}258 259common_reasoning_budget_state common_reasoning_budget_get_state(const struct llama_sampler * smpl) {260 if (!smpl) {261 return REASONING_BUDGET_IDLE;262 }263 return ((const common_reasoning_budget_ctx *)smpl->ctx)->state;264}265 