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
reasoning-budget.cpp265 linesDownload Raw Back to common
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