Codeprocastinator/optimized-tinyllama-covalent
0119
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 