CoolFace
Modelpublic

Codeprocastinator/optimized-tinyllama-covalent

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