KBaba7/llama.cpp
0
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 