Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
sampling.cpp527 linesDownload Raw Back to common
1#include "sampling.h"2 3#include "common.h"4 5#include <cmath>6#include <unordered_map>7 8// the ring buffer works similarly to std::deque, but with a fixed capacity9// TODO: deduplicate with llama-impl.h10template<typename T>11struct ring_buffer {12    ring_buffer(size_t cap) : capacity(cap), data(cap) {}13 14    T & front() {15        if (sz == 0) {16            throw std::runtime_error("ring buffer is empty");17        }18        return data[first];19    }20 21    const T & front() const {22        if (sz == 0) {23            throw std::runtime_error("ring buffer is empty");24        }25        return data[first];26    }27 28    T & back() {29        if (sz == 0) {30            throw std::runtime_error("ring buffer is empty");31        }32        return data[pos];33    }34 35    const T & back() const {36        if (sz == 0) {37            throw std::runtime_error("ring buffer is empty");38        }39        return data[pos];40    }41 42    void push_back(const T & value) {43        if (sz == capacity) {44            // advance the start when buffer is full45            first = (first + 1) % capacity;46        } else {47            sz++;48        }49        data[pos] = value;50        pos = (pos + 1) % capacity;51    }52 53    T pop_front() {54        if (sz == 0) {55            throw std::runtime_error("ring buffer is empty");56        }57        T value = data[first];58        first = (first + 1) % capacity;59        sz--;60        return value;61    }62 63    const T & rat(size_t i) const {64        if (i >= sz) {65            throw std::runtime_error("ring buffer: index out of bounds");66        }67        return data[(first + sz - i - 1) % capacity];68    }69 70    std::vector<T> to_vector() const {71        std::vector<T> result;72        result.reserve(sz);73        for (size_t i = 0; i < sz; i++) {74            result.push_back(data[(first + i) % capacity]);75        }76        return result;77    }78 79    void clear() {80        // here only reset the status of the buffer81        sz = 0;82        first = 0;83        pos = 0;84    }85 86    bool empty() const {87        return sz == 0;88    }89 90    size_t size() const {91        return sz;92    }93 94    size_t capacity = 0;95    size_t sz = 0;96    size_t first = 0;97    size_t pos = 0;98    std::vector<T> data;99};100 101struct common_sampler {102    common_params_sampling params;103 104    struct llama_sampler * grmr;105    struct llama_sampler * chain;106 107    ring_buffer<llama_token> prev;108 109    std::vector<llama_token_data> cur;110 111    llama_token_data_array cur_p;112 113    void set_logits(struct llama_context * ctx, int idx) {114        const auto * logits = llama_get_logits_ith(ctx, idx);115 116        const llama_model * model = llama_get_model(ctx);117        const llama_vocab * vocab = llama_model_get_vocab(model);118 119        const int n_vocab = llama_vocab_n_tokens(vocab);120 121        cur.resize(n_vocab);122 123        for (llama_token token_id = 0; token_id < n_vocab; token_id++) {124            cur[token_id] = llama_token_data{token_id, logits[token_id], 0.0f};125        }126 127        cur_p = { cur.data(), cur.size(), -1, false };128    }129};130 131std::string common_params_sampling::print() const {132    char result[1024];133 134    snprintf(result, sizeof(result),135            "\trepeat_last_n = %d, repeat_penalty = %.3f, frequency_penalty = %.3f, presence_penalty = %.3f\n"136            "\tdry_multiplier = %.3f, dry_base = %.3f, dry_allowed_length = %d, dry_penalty_last_n = %d\n"137            "\ttop_k = %d, top_p = %.3f, min_p = %.3f, xtc_probability = %.3f, xtc_threshold = %.3f, typical_p = %.3f, temp = %.3f\n"138            "\tmirostat = %d, mirostat_lr = %.3f, mirostat_ent = %.3f",139            penalty_last_n, penalty_repeat, penalty_freq, penalty_present,140            dry_multiplier, dry_base, dry_allowed_length, dry_penalty_last_n,141            top_k, top_p, min_p, xtc_probability, xtc_threshold, typ_p, temp,142            mirostat, mirostat_eta, mirostat_tau);143 144    return std::string(result);145}146 147struct common_sampler * common_sampler_init(const struct llama_model * model, const struct common_params_sampling & params) {148    const llama_vocab * vocab = llama_model_get_vocab(model);149 150    llama_sampler_chain_params lparams = llama_sampler_chain_default_params();151 152    lparams.no_perf = params.no_perf;153 154    std::vector<const char *> trigger_words;155    trigger_words.reserve(params.grammar_trigger_words.size());156    for (const auto & str : params.grammar_trigger_words) {157        trigger_words.push_back(str.word.c_str());158    }159 160    struct llama_sampler * grmr;161    if (params.grammar.compare(0, 11, "%llguidance") == 0) {162#ifdef LLAMA_USE_LLGUIDANCE163        grmr = llama_sampler_init_llg(vocab, "lark", params.grammar.c_str());164#else165        GGML_ABORT("llguidance (cmake -DLLAMA_LLGUIDANCE=ON) is not enabled");166#endif // LLAMA_USE_LLGUIDANCE167    } else {168        grmr = params.grammar_lazy169             ? llama_sampler_init_grammar_lazy(vocab, params.grammar.c_str(), "root",170                                               trigger_words.data(), trigger_words.size(),171                                               params.grammar_trigger_tokens.data(), params.grammar_trigger_tokens.size())172             :      llama_sampler_init_grammar(vocab, params.grammar.c_str(), "root");173    }174 175    auto * result = new common_sampler {176        /* .params = */ params,177        /* .grmr   = */ grmr,178        /* .chain  = */ llama_sampler_chain_init(lparams),179        /* .prev   = */ ring_buffer<llama_token>(std::max(32, params.n_prev)),180        /* .cur    = */ {},181        /* .cur_p  = */ {},182    };183 184    llama_sampler_chain_add(result->chain,185            llama_sampler_init_logit_bias(186                llama_vocab_n_tokens(vocab),187                params.logit_bias.size(),188                params.logit_bias.data()));189 190    if (params.mirostat == 0) {191        for (const auto & cnstr : params.samplers) {192            switch (cnstr) {193                case COMMON_SAMPLER_TYPE_DRY:194                    {195                        std::vector<const char *> c_breakers;196                        c_breakers.reserve(params.dry_sequence_breakers.size());197                        for (const auto & str : params.dry_sequence_breakers) {198                            c_breakers.push_back(str.c_str());199                        }200 201                        llama_sampler_chain_add(result->chain, llama_sampler_init_dry      (vocab, llama_model_n_ctx_train(model), params.dry_multiplier, params.dry_base, params.dry_allowed_length, params.dry_penalty_last_n, c_breakers.data(), c_breakers.size()));202                    }203                    break;204                case COMMON_SAMPLER_TYPE_TOP_K:205                    llama_sampler_chain_add(result->chain, llama_sampler_init_top_k    (params.top_k));206                    break;207                case COMMON_SAMPLER_TYPE_TOP_P:208                    llama_sampler_chain_add(result->chain, llama_sampler_init_top_p    (params.top_p, params.min_keep));209                    break;210                case COMMON_SAMPLER_TYPE_MIN_P:211                    llama_sampler_chain_add(result->chain, llama_sampler_init_min_p    (params.min_p, params.min_keep));212                    break;213                case COMMON_SAMPLER_TYPE_XTC:214                    llama_sampler_chain_add(result->chain, llama_sampler_init_xtc      (params.xtc_probability, params.xtc_threshold, params.min_keep, params.seed));215                    break;216                case COMMON_SAMPLER_TYPE_TYPICAL_P:217                    llama_sampler_chain_add(result->chain, llama_sampler_init_typical  (params.typ_p, params.min_keep));218                    break;219                case COMMON_SAMPLER_TYPE_TEMPERATURE:220                    llama_sampler_chain_add(result->chain, llama_sampler_init_temp_ext (params.temp, params.dynatemp_range, params.dynatemp_exponent));221                    break;222                case COMMON_SAMPLER_TYPE_INFILL:223                    llama_sampler_chain_add(result->chain, llama_sampler_init_infill   (vocab));224                    break;225                case COMMON_SAMPLER_TYPE_PENALTIES:226                    llama_sampler_chain_add(result->chain, llama_sampler_init_penalties(params.penalty_last_n, params.penalty_repeat, params.penalty_freq, params.penalty_present));227                    break;228                default:229                    GGML_ASSERT(false && "unknown sampler type");230            }231        }232        llama_sampler_chain_add(result->chain, llama_sampler_init_dist(params.seed));233    } else if (params.mirostat == 1) {234        llama_sampler_chain_add(result->chain, llama_sampler_init_temp(params.temp));235        llama_sampler_chain_add(result->chain, llama_sampler_init_mirostat(llama_vocab_n_tokens(vocab), params.seed, params.mirostat_tau, params.mirostat_eta, 100));236    } else if (params.mirostat == 2) {237        llama_sampler_chain_add(result->chain, llama_sampler_init_temp(params.temp));238        llama_sampler_chain_add(result->chain, llama_sampler_init_mirostat_v2(params.seed, params.mirostat_tau, params.mirostat_eta));239    } else {240        GGML_ASSERT(false && "unknown mirostat version");241    }242 243    return result;244}245 246void common_sampler_free(struct common_sampler * gsmpl) {247    if (gsmpl) {248        llama_sampler_free(gsmpl->grmr);249 250        llama_sampler_free(gsmpl->chain);251 252        delete gsmpl;253    }254}255 256void common_sampler_accept(struct common_sampler * gsmpl, llama_token token, bool accept_grammar) {257    if (accept_grammar) {258        llama_sampler_accept(gsmpl->grmr, token);259    }260 261    llama_sampler_accept(gsmpl->chain, token);262 263    gsmpl->prev.push_back(token);264}265 266void common_sampler_reset(struct common_sampler * gsmpl) {267    llama_sampler_reset(gsmpl->grmr);268 269    llama_sampler_reset(gsmpl->chain);270}271 272struct common_sampler * common_sampler_clone(common_sampler * gsmpl) {273    return new common_sampler {274        /* .params = */ gsmpl->params,275        /* .grmr   = */ llama_sampler_clone(gsmpl->grmr),276        /* .chain  = */ llama_sampler_clone(gsmpl->chain),277        /* .prev   = */ gsmpl->prev,278        /* .cur    = */ gsmpl->cur,279        /* .cur_p  = */ gsmpl->cur_p,280    };281}282 283void common_perf_print(const struct llama_context * ctx, const struct common_sampler * gsmpl) {284    // TODO: measure grammar performance285 286    if (gsmpl) {287        llama_perf_sampler_print(gsmpl->chain);288    }289    if (ctx) {290        llama_perf_context_print(ctx);291    }292}293 294llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_context * ctx, int idx, bool grammar_first) {295    gsmpl->set_logits(ctx, idx);296 297    auto & grmr  = gsmpl->grmr;298    auto & chain = gsmpl->chain;299    auto & cur_p = gsmpl->cur_p; // initialized by set_logits300 301    if (grammar_first) {302        llama_sampler_apply(grmr, &cur_p);303    }304 305    llama_sampler_apply(chain, &cur_p);306 307    GGML_ASSERT(cur_p.selected != -1 && "no selected token during sampling - check your sampling configuration");308 309    const llama_token id = cur_p.data[cur_p.selected].id;310 311    if (grammar_first) {312        return id;313    }314 315    // check if it the sampled token fits the grammar316    {317        llama_token_data       single_token_data       = { id, 1.0f, 0.0f };318        llama_token_data_array single_token_data_array = { &single_token_data, 1, -1, false };319 320        llama_sampler_apply(grmr, &single_token_data_array);321 322        const bool is_valid = single_token_data_array.data[0].logit != -INFINITY;323        if (is_valid) {324            return id;325        }326    }327 328    // resampling:329    // if the token is not valid, sample again, but first apply the grammar sampler and then the sampling chain330    gsmpl->set_logits(ctx, idx);331 332    llama_sampler_apply(grmr,  &cur_p);333    llama_sampler_apply(chain, &cur_p);334 335    GGML_ASSERT(cur_p.selected != -1 && "no selected token during re-sampling - check your sampling configuration");336 337    return cur_p.data[cur_p.selected].id;338}339 340std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const std::vector<int> & idxs, const llama_tokens & draft, bool grammar_first) {341    GGML_ASSERT(idxs.size() == draft.size() + 1 && "idxs.size() must be draft.size() + 1");342 343    std::vector<llama_token> result;344    result.reserve(idxs.size());345 346    size_t i = 0;347    for (; i < draft.size(); i++) {348        const llama_token id = common_sampler_sample(gsmpl, ctx, idxs[i], grammar_first);349 350        common_sampler_accept(gsmpl, id, true);351 352        result.push_back(id);353 354        if (draft[i] != id) {355            break;356        }357    }358 359    if (i == draft.size()) {360        const llama_token id = common_sampler_sample(gsmpl, ctx, idxs[i], grammar_first);361 362        common_sampler_accept(gsmpl, id, true);363 364        result.push_back(id);365    }366 367    return result;368}369 370std::vector<llama_token> common_sampler_sample_and_accept_n(struct common_sampler * gsmpl, struct llama_context * ctx, const llama_tokens & draft, bool grammar_first) {371    std::vector<int> idxs(draft.size() + 1);372    for (size_t i = 0; i < idxs.size(); ++i) {373        idxs[i] = i;374    }375 376    return common_sampler_sample_and_accept_n(gsmpl, ctx, idxs, draft, grammar_first);377}378 379uint32_t common_sampler_get_seed(const struct common_sampler * gsmpl) {380    return llama_sampler_get_seed(gsmpl->chain);381}382 383// helpers384 385llama_token_data_array * common_sampler_get_candidates(struct common_sampler * gsmpl) {386    return &gsmpl->cur_p;387}388 389llama_token common_sampler_last(const struct common_sampler * gsmpl) {390    return gsmpl->prev.rat(0);391}392 393std::string common_sampler_print(const struct common_sampler * gsmpl) {394    std::string result = "logits ";395 396    for (int i = 0; i < llama_sampler_chain_n(gsmpl->chain); i++) {397        const auto * smpl = llama_sampler_chain_get(gsmpl->chain, i);398        result += std::string("-> ") + llama_sampler_name(smpl) + " ";399    }400 401    return result;402}403 404std::string common_sampler_prev_str(common_sampler * gsmpl, llama_context * ctx_main, int n) {405    n = std::min(n, (int) gsmpl->prev.size());406 407    if (n <= 0) {408        return "";409    }410 411    std::string result;412    result.reserve(8*n); // 8 is the average length of a token [citation needed], TODO: compute this from the vocab413 414    for (int i = n - 1; i >= 0; i--) {415        const llama_token id = gsmpl->prev.rat(i);416 417        GGML_ASSERT(id != LLAMA_TOKEN_NULL && "null token in the sampling history - should not happen");418 419        result += common_token_to_piece(ctx_main, id);420    }421 422    return result;423}424 425char common_sampler_type_to_chr(enum common_sampler_type cnstr) {426    switch (cnstr) {427        case COMMON_SAMPLER_TYPE_DRY:         return 'd';428        case COMMON_SAMPLER_TYPE_TOP_K:       return 'k';429        case COMMON_SAMPLER_TYPE_TYPICAL_P:   return 'y';430        case COMMON_SAMPLER_TYPE_TOP_P:       return 'p';431        case COMMON_SAMPLER_TYPE_MIN_P:       return 'm';432        case COMMON_SAMPLER_TYPE_TEMPERATURE: return 't';433        case COMMON_SAMPLER_TYPE_XTC:         return 'x';434        case COMMON_SAMPLER_TYPE_INFILL:      return 'i';435        case COMMON_SAMPLER_TYPE_PENALTIES:   return 'e';436        default : return '?';437    }438}439 440std::string common_sampler_type_to_str(enum common_sampler_type cnstr) {441    switch (cnstr) {442        case COMMON_SAMPLER_TYPE_DRY:         return "dry";443        case COMMON_SAMPLER_TYPE_TOP_K:       return "top_k";444        case COMMON_SAMPLER_TYPE_TYPICAL_P:   return "typ_p";445        case COMMON_SAMPLER_TYPE_TOP_P:       return "top_p";446        case COMMON_SAMPLER_TYPE_MIN_P:       return "min_p";447        case COMMON_SAMPLER_TYPE_TEMPERATURE: return "temperature";448        case COMMON_SAMPLER_TYPE_XTC:         return "xtc";449        case COMMON_SAMPLER_TYPE_INFILL:      return "infill";450        case COMMON_SAMPLER_TYPE_PENALTIES:   return "penalties";451        default : return "";452    }453}454 455std::vector<common_sampler_type> common_sampler_types_from_names(const std::vector<std::string> & names, bool allow_alt_names) {456    std::unordered_map<std::string, common_sampler_type> sampler_canonical_name_map {457        { "dry",         COMMON_SAMPLER_TYPE_DRY },458        { "top_k",       COMMON_SAMPLER_TYPE_TOP_K },459        { "top_p",       COMMON_SAMPLER_TYPE_TOP_P },460        { "typ_p",       COMMON_SAMPLER_TYPE_TYPICAL_P },461        { "min_p",       COMMON_SAMPLER_TYPE_MIN_P },462        { "temperature", COMMON_SAMPLER_TYPE_TEMPERATURE },463        { "xtc",         COMMON_SAMPLER_TYPE_XTC },464        { "infill",      COMMON_SAMPLER_TYPE_INFILL },465        { "penalties",   COMMON_SAMPLER_TYPE_PENALTIES },466    };467 468    // since samplers names are written multiple ways469    // make it ready for both system names and input names470    std::unordered_map<std::string, common_sampler_type> sampler_alt_name_map {471        { "top-k",       COMMON_SAMPLER_TYPE_TOP_K },472        { "top-p",       COMMON_SAMPLER_TYPE_TOP_P },473        { "nucleus",     COMMON_SAMPLER_TYPE_TOP_P },474        { "typical-p",   COMMON_SAMPLER_TYPE_TYPICAL_P },475        { "typical",     COMMON_SAMPLER_TYPE_TYPICAL_P },476        { "typ-p",       COMMON_SAMPLER_TYPE_TYPICAL_P },477        { "typ",         COMMON_SAMPLER_TYPE_TYPICAL_P },478        { "min-p",       COMMON_SAMPLER_TYPE_MIN_P },479        { "temp",        COMMON_SAMPLER_TYPE_TEMPERATURE },480    };481 482    std::vector<common_sampler_type> samplers;483    samplers.reserve(names.size());484 485    for (const auto & name : names) {486        auto sampler = sampler_canonical_name_map.find(name);487        if (sampler != sampler_canonical_name_map.end()) {488            samplers.push_back(sampler->second);489        } else {490            if (allow_alt_names) {491                sampler = sampler_alt_name_map.find(name);492                if (sampler != sampler_alt_name_map.end()) {493                    samplers.push_back(sampler->second);494                }495            }496        }497    }498 499    return samplers;500}501 502std::vector<common_sampler_type> common_sampler_types_from_chars(const std::string & chars) {503    std::unordered_map<char, common_sampler_type> sampler_name_map = {504        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_DRY),         COMMON_SAMPLER_TYPE_DRY },505        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_TOP_K),       COMMON_SAMPLER_TYPE_TOP_K },506        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_TYPICAL_P),   COMMON_SAMPLER_TYPE_TYPICAL_P },507        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_TOP_P),       COMMON_SAMPLER_TYPE_TOP_P },508        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_MIN_P),       COMMON_SAMPLER_TYPE_MIN_P },509        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_TEMPERATURE), COMMON_SAMPLER_TYPE_TEMPERATURE },510        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_XTC),         COMMON_SAMPLER_TYPE_XTC },511        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_INFILL),      COMMON_SAMPLER_TYPE_INFILL },512        { common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_PENALTIES),   COMMON_SAMPLER_TYPE_PENALTIES },513    };514 515    std::vector<common_sampler_type> samplers;516    samplers.reserve(chars.size());517 518    for (const auto & c : chars) {519        const auto sampler = sampler_name_map.find(c);520        if (sampler != sampler_name_map.end()) {521            samplers.push_back(sampler->second);522        }523    }524 525    return samplers;526}527