KBaba7/llama.cpp
0
1#include "llama-sampling.h"2 3#include "llama-impl.h"4#include "llama-vocab.h"5#include "llama-grammar.h"6 7#include <algorithm>8#include <cassert>9#include <cfloat>10#include <chrono>11#include <cmath>12#include <cstdlib>13#include <cstring>14#include <ctime>15#include <numeric>16#include <random>17#include <unordered_map>18#include <stdexcept>19 20// the ring buffer works similarly to std::deque, but with a fixed capacity21template<typename T>22struct ring_buffer {23 ring_buffer(size_t cap) : capacity(cap), data(cap) {}24 25 T & front() {26 if (sz == 0) {27 throw std::runtime_error("ring buffer is empty");28 }29 return data[first];30 }31 32 const T & front() const {33 if (sz == 0) {34 throw std::runtime_error("ring buffer is empty");35 }36 return data[first];37 }38 39 T & back() {40 if (sz == 0) {41 throw std::runtime_error("ring buffer is empty");42 }43 return data[pos];44 }45 46 const T & back() const {47 if (sz == 0) {48 throw std::runtime_error("ring buffer is empty");49 }50 return data[pos];51 }52 53 void push_back(const T & value) {54 if (capacity == 0) {55 throw std::runtime_error("ring buffer: capacity is zero");56 }57 58 if (sz == capacity) {59 // advance the start when buffer is full60 first = (first + 1) % capacity;61 } else {62 sz++;63 }64 data[pos] = value;65 pos = (pos + 1) % capacity;66 }67 68 T pop_front() {69 if (sz == 0) {70 throw std::runtime_error("ring buffer is empty");71 }72 T value = data[first];73 first = (first + 1) % capacity;74 sz--;75 return value;76 }77 78 //T & operator[](size_t i) {79 // if (i >= sz) {80 // throw std::runtime_error("ring buffer: index out of bounds");81 // }82 // return data[(first + i) % capacity];83 //}84 85 //const T & at(size_t i) const {86 // if (i >= sz) {87 // throw std::runtime_error("ring buffer: index out of bounds");88 // }89 // return data[(first + i) % capacity];90 //}91 92 const T & rat(size_t i) const {93 if (i >= sz) {94 throw std::runtime_error("ring buffer: index out of bounds");95 }96 return data[(first + sz - i - 1) % capacity];97 }98 99 std::vector<T> to_vector() const {100 std::vector<T> result;101 result.reserve(sz);102 for (size_t i = 0; i < sz; i++) {103 result.push_back(data[(first + i) % capacity]);104 }105 return result;106 }107 108 void clear() {109 // here only reset the status of the buffer110 sz = 0;111 first = 0;112 pos = 0;113 }114 115 bool empty() const {116 return sz == 0;117 }118 119 size_t size() const {120 return sz;121 }122 123 size_t capacity = 0;124 size_t sz = 0;125 size_t first = 0;126 size_t pos = 0;127 128 std::vector<T> data;129};130 131static int llama_sample_dist(llama_token_data_array * cur_p, std::mt19937 & rng) {132 // iterator for the probabilities133#ifdef __GNUC__134 #pragma GCC diagnostic push135 #pragma GCC diagnostic ignored "-Wunused-local-typedefs"136#endif137 138 struct probs_iterator {139 typedef std::input_iterator_tag iterator_category;140 typedef float value_type;141 typedef float * pointer;142 typedef float & reference;143 typedef ptrdiff_t difference_type;144 145 const llama_token_data * data;146 147 bool operator==(const probs_iterator & other) const { return data == other.data; }148 bool operator!=(const probs_iterator & other) const { return data != other.data; }149 const float & operator*() const { return data->p; }150 probs_iterator & operator++() { ++data; return *this; }151 probs_iterator operator++(int) { probs_iterator tmp = *this; ++data; return tmp; }152 };153 154#ifdef __GNUC__155 #pragma GCC diagnostic pop156#endif157 158 std::discrete_distribution<int> dist(probs_iterator{cur_p->data}, probs_iterator{cur_p->data + cur_p->size});159 160 return dist(rng);161}162 163/*164static void llama_log_softmax(float * array, size_t size) {165 float max_l = *std::max_element(array, array + size);166 float sum = 0.f;167 for (size_t i = 0; i < size; ++i) {168 float p = expf(array[i] - max_l);169 sum += p;170 array[i] = p;171 }172 173 for (size_t i = 0; i < size; ++i) {174 array[i] = logf(array[i] / sum);175 }176}177*/178 179static void llama_sampler_temp_impl(llama_token_data_array * cur_p, float temp) {180 if (temp <= 0.0f) {181 // find the token with the highest logit and set the rest to -inf182 size_t max_i = 0;183 float max_l = cur_p->data[0].logit;184 185 for (size_t i = 1; i < cur_p->size; ++i) {186 if (cur_p->data[i ].logit > max_l) {187 cur_p->data[max_i].logit = -INFINITY;188 max_i = i;189 max_l = cur_p->data[i].logit;190 } else {191 cur_p->data[i].logit = -INFINITY;192 }193 }194 195 return;196 }197 198 for (size_t i = 0; i < cur_p->size; ++i) {199 cur_p->data[i].logit /= temp;200 }201}202 203static void llama_sampler_softmax_impl(llama_token_data_array * cur_p) {204 GGML_ASSERT(cur_p->size > 0);205 206 // Sort the logits in descending order207 if (!cur_p->sorted) {208 std::sort(cur_p->data, cur_p->data + cur_p->size, [](const llama_token_data & a, const llama_token_data & b) {209 return a.logit > b.logit;210 });211 cur_p->sorted = true;212 }213 214 float max_l = cur_p->data[0].logit;215 float cum_sum = 0.0f;216 217 for (size_t i = 0; i < cur_p->size; ++i) {218 float p = expf(cur_p->data[i].logit - max_l);219 cur_p->data[i].p = p;220 cum_sum += p;221 }222 223 for (size_t i = 0; i < cur_p->size; ++i) {224 cur_p->data[i].p /= cum_sum;225 }226}227 228static void llama_sampler_top_k_impl(llama_token_data_array * cur_p, int32_t k) {229 // TODO: move bucket sort to separate function so that top_p/typical/softmax first is equally fast230 // if (k >= (int32_t)cur_p->size) {231 // return;232 // }233 234 if (k <= 0) {235 k = cur_p->size;236 }237 238 k = std::min(k, (int) cur_p->size);239 240 // Sort scores in descending order241 if (!cur_p->sorted) {242 auto comp = [](const llama_token_data & a, const llama_token_data & b) {243 return a.logit > b.logit;244 };245 if (k <= 128) {246 std::partial_sort(cur_p->data, cur_p->data + k, cur_p->data + cur_p->size, comp);247 } else {248 constexpr int nbuckets = 128;249 constexpr float bucket_low = -10.0f;250 constexpr float bucket_high = 10.0f;251 constexpr float bucket_scale = nbuckets/(bucket_high - bucket_low);252 constexpr float bucket_inter = -bucket_low * bucket_scale;253 254 std::vector<int> bucket_idx(cur_p->size);255 std::vector<int> histo(nbuckets, 0);256 257 for (int i = 0; i < (int)cur_p->size; ++i) {258 const float val = cur_p->data[i].logit;259 int ib = int(bucket_scale * val + bucket_inter); //nbuckets * (val - bucket_low) / (bucket_high - bucket_low);260 ib = std::max(0, std::min(nbuckets - 1, ib));261 bucket_idx[i] = ib;262 ++histo[ib];263 }264 int nhave = 0;265 int ib = nbuckets - 1;266 for ( ; ib >= 0; --ib) {267 nhave += histo[ib];268 if (nhave >= k) {269 break;270 }271 }272 std::vector<llama_token_data> tmp_tokens(nhave);273 auto * ptr = tmp_tokens.data();274 std::vector<llama_token_data*> bucket_ptrs;275 bucket_ptrs.reserve(nbuckets - ib);276 for (int j = nbuckets - 1; j >= ib; --j) {277 bucket_ptrs.push_back(ptr);278 ptr += histo[j];279 }280 for (int i = 0; i < (int)cur_p->size; ++i) {281 int j = bucket_idx[i];282 if (j >= ib) {283 *bucket_ptrs[nbuckets - 1 - j]++ = cur_p->data[i];284 }285 }286 287 ptr = tmp_tokens.data();288 int ndone = 0;289 for (int j = nbuckets - 1; j > ib; --j) {290 std::sort(ptr, ptr + histo[j], comp);291 ptr += histo[j];292 ndone += histo[j];293 }294 std::partial_sort(ptr, ptr + k - ndone, ptr + histo[ib], comp);295 296 std::memcpy(cur_p->data, tmp_tokens.data(), k*sizeof(llama_token_data));297 298 }299 cur_p->sorted = true;300 }301 cur_p->size = k;302}303 304static uint32_t get_rng_seed(uint32_t seed) {305 if (seed == LLAMA_DEFAULT_SEED) {306 // use system clock if std::random_device is not a true RNG307 static bool is_rd_prng = std::random_device().entropy() == 0;308 if (is_rd_prng) {309 return (uint32_t) std::chrono::system_clock::now().time_since_epoch().count();310 }311 std::random_device rd;312 return rd();313 }314 return seed;315}316 317// llama_sampler API318 319struct llama_sampler * llama_sampler_init(const struct llama_sampler_i * iface, llama_sampler_context_t ctx) {320 return new llama_sampler {321 /* .iface = */ iface,322 /* .ctx = */ ctx,323 };324}325 326const char * llama_sampler_name(const struct llama_sampler * smpl) {327 if (!smpl->iface) {328 return "(null)";329 }330 331 return smpl->iface->name(smpl);332}333 334void llama_sampler_accept(struct llama_sampler * smpl, llama_token token) {335 if (smpl->iface->accept) {336 smpl->iface->accept(smpl, token);337 }338}339 340void llama_sampler_apply(struct llama_sampler * smpl, struct llama_token_data_array * cur_p) {341 GGML_ASSERT(smpl->iface->apply);342 smpl->iface->apply(smpl, cur_p);343}344 345void llama_sampler_reset(struct llama_sampler * smpl) {346 if (smpl->iface->reset) {347 smpl->iface->reset(smpl);348 }349}350 351struct llama_sampler * llama_sampler_clone(const struct llama_sampler * smpl) {352 if (smpl->iface->clone) {353 return smpl->iface->clone(smpl);354 }355 356 if (smpl->ctx == nullptr) {357 return llama_sampler_init(358 /* .iface = */ smpl->iface,359 /* .ctx = */ nullptr360 );361 }362 363 GGML_ABORT("the sampler does not support cloning");364}365 366void llama_sampler_free(struct llama_sampler * smpl) {367 if (smpl == nullptr) {368 return;369 }370 371 if (smpl->iface->free) {372 smpl->iface->free(smpl);373 }374 375 delete smpl;376}377 378llama_token llama_sampler_sample(struct llama_sampler * smpl, struct llama_context * ctx, int32_t idx) {379 const auto * logits = llama_get_logits_ith(ctx, idx);380 381 const llama_model * model = llama_get_model(ctx);382 const llama_vocab * vocab = llama_model_get_vocab(model);383 384 const int n_vocab = llama_vocab_n_tokens(vocab);385 386 // TODO: do not allocate each time387 std::vector<llama_token_data> cur;388 cur.reserve(n_vocab);389 for (llama_token token_id = 0; token_id < n_vocab; token_id++) {390 cur.emplace_back(llama_token_data{token_id, logits[token_id], 0.0f});391 }392 393 llama_token_data_array cur_p = {394 /* .data = */ cur.data(),395 /* .size = */ cur.size(),396 /* .selected = */ -1,397 /* .sorted = */ false,398 };399 400 llama_sampler_apply(smpl, &cur_p);401 402 GGML_ASSERT(cur_p.selected >= 0 && cur_p.selected < (int32_t) cur_p.size);403 404 auto token = cur_p.data[cur_p.selected].id;405 406 llama_sampler_accept(smpl, token);407 408 return token;409}410 411// sampler chain412 413static const char * llama_sampler_chain_name(const struct llama_sampler * /*smpl*/) {414 return "chain";415}416 417static void llama_sampler_chain_accept(struct llama_sampler * smpl, llama_token token) {418 auto * chain = (llama_sampler_chain *) smpl->ctx;419 420 time_meas tm(chain->t_sample_us, chain->params.no_perf);421 422 for (auto * smpl : chain->samplers) {423 llama_sampler_accept(smpl, token);424 }425 426 chain->n_sample++;427}428 429static void llama_sampler_chain_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {430 auto * chain = (llama_sampler_chain *) smpl->ctx;431 432 time_meas tm(chain->t_sample_us, chain->params.no_perf);433 434 for (auto * smpl : chain->samplers) {435 llama_sampler_apply(smpl, cur_p);436 }437}438 439static void llama_sampler_chain_reset(struct llama_sampler * smpl) {440 auto * chain = (llama_sampler_chain *) smpl->ctx;441 442 for (auto * smpl : chain->samplers) {443 llama_sampler_reset(smpl);444 }445 446 chain->t_sample_us = 0;447 chain->n_sample = 0;448}449 450static struct llama_sampler * llama_sampler_chain_clone(const struct llama_sampler * smpl) {451 const auto * chain_src = (const llama_sampler_chain *) smpl->ctx;452 453 auto * result = llama_sampler_chain_init(chain_src->params);454 455 for (auto * smpl : chain_src->samplers) {456 llama_sampler_chain_add(result, llama_sampler_clone(smpl));457 }458 459 return result;460}461 462static void llama_sampler_chain_free(struct llama_sampler * smpl) {463 auto * chain = (llama_sampler_chain *) smpl->ctx;464 465 for (auto * smpl : chain->samplers) {466 llama_sampler_free(smpl);467 }468 469 delete chain;470}471 472static struct llama_sampler_i llama_sampler_chain_i = {473 /* .name = */ llama_sampler_chain_name,474 /* .accept = */ llama_sampler_chain_accept,475 /* .apply = */ llama_sampler_chain_apply,476 /* .reset = */ llama_sampler_chain_reset,477 /* .clone = */ llama_sampler_chain_clone,478 /* .free = */ llama_sampler_chain_free,479};480 481struct llama_sampler * llama_sampler_chain_init(struct llama_sampler_chain_params params) {482 return llama_sampler_init(483 /* .iface = */ &llama_sampler_chain_i,484 /* .ctx = */ new llama_sampler_chain {485 /* .params = */ params,486 /* .samplers = */ {},487 /* .t_sample_us = */ 0,488 /* .n_sample = */ 0,489 }490 );491}492 493void llama_sampler_chain_add(struct llama_sampler * chain, struct llama_sampler * smpl) {494 auto * p = (llama_sampler_chain *) chain->ctx;495 p->samplers.push_back(smpl);496}497 498struct llama_sampler * llama_sampler_chain_get(const struct llama_sampler * chain, int32_t i) {499 const auto * p = (const llama_sampler_chain *) chain->ctx;500 501 if (i < 0 || (size_t) i >= p->samplers.size()) {502 return nullptr;503 }504 505 return p->samplers[i];506}507 508struct llama_sampler * llama_sampler_chain_remove(struct llama_sampler * chain, int32_t i) {509 auto * p = (llama_sampler_chain *) chain->ctx;510 511 if (i < 0 || (size_t) i >= p->samplers.size()) {512 return nullptr;513 }514 515 auto * result = p->samplers[i];516 p->samplers.erase(p->samplers.begin() + i);517 518 return result;519}520 521int llama_sampler_chain_n(const struct llama_sampler * chain) {522 const auto * p = (const llama_sampler_chain *) chain->ctx;523 524 return p->samplers.size();525}526 527//528// samplers529//530 531// greedy532 533static const char * llama_sampler_greedy_name(const struct llama_sampler * /*smpl*/) {534 return "greedy";535}536 537static void llama_sampler_greedy_apply(struct llama_sampler * /*smpl*/, llama_token_data_array * cur_p) {538 cur_p->selected = 0;539 for (size_t i = 1; i < cur_p->size; ++i) {540 if (cur_p->data[i].logit > cur_p->data[cur_p->selected].logit) {541 cur_p->selected = i;542 }543 }544}545 546static struct llama_sampler_i llama_sampler_greedy_i = {547 /* .name = */ llama_sampler_greedy_name,548 /* .accept = */ nullptr,549 /* .apply = */ llama_sampler_greedy_apply,550 /* .reset = */ nullptr,551 /* .clone = */ nullptr,552 /* .free = */ nullptr,553};554 555struct llama_sampler * llama_sampler_init_greedy() {556 return llama_sampler_init(557 /* .iface = */ &llama_sampler_greedy_i,558 /* .ctx = */ nullptr559 );560}561 562// dist563 564struct llama_sampler_dist {565 const uint32_t seed;566 uint32_t seed_cur;567 568 std::mt19937 rng;569};570 571static const char * llama_sampler_dist_name(const struct llama_sampler * /*smpl*/) {572 return "dist";573}574 575static void llama_sampler_dist_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {576 auto * ctx = (llama_sampler_dist *) smpl->ctx;577 578 llama_sampler_softmax_impl(cur_p);579 580 cur_p->selected = llama_sample_dist(cur_p, ctx->rng);581}582 583static struct llama_sampler * llama_sampler_dist_clone(const struct llama_sampler * smpl) {584 const auto * ctx = (const llama_sampler_dist *) smpl->ctx;585 auto * result = llama_sampler_init_dist(ctx->seed);586 587 // copy the state588 {589 auto * result_ctx = (llama_sampler_dist *) result->ctx;590 591 result_ctx->rng = ctx->rng;592 }593 594 return result;595}596 597static void llama_sampler_dist_reset(struct llama_sampler * smpl) {598 auto * ctx = (llama_sampler_dist *) smpl->ctx;599 ctx->seed_cur = get_rng_seed(ctx->seed);600 ctx->rng.seed(ctx->seed_cur);601}602 603static void llama_sampler_dist_free(struct llama_sampler * smpl) {604 delete (llama_sampler_dist *) smpl->ctx;605}606 607static struct llama_sampler_i llama_sampler_dist_i = {608 /* .name = */ llama_sampler_dist_name,609 /* .accept = */ nullptr,610 /* .apply = */ llama_sampler_dist_apply,611 /* .reset = */ llama_sampler_dist_reset,612 /* .clone = */ llama_sampler_dist_clone,613 /* .free = */ llama_sampler_dist_free,614};615 616struct llama_sampler * llama_sampler_init_dist(uint32_t seed) {617 auto seed_cur = get_rng_seed(seed);618 return llama_sampler_init(619 /* .iface = */ &llama_sampler_dist_i,620 /* .ctx = */ new llama_sampler_dist {621 /* .seed = */ seed,622 /* .seed_cur = */ seed_cur,623 /* .rng = */ std::mt19937(seed_cur),624 }625 );626}627 628// softmax629 630static const char * llama_sampler_softmax_name(const struct llama_sampler * /*smpl*/) {631 return "softmax";632}633 634static void llama_sampler_softmax_apply(struct llama_sampler * /*smpl*/, llama_token_data_array * cur_p) {635 llama_sampler_softmax_impl(cur_p);636}637 638static struct llama_sampler_i llama_sampler_softmax_i = {639 /* .name = */ llama_sampler_softmax_name,640 /* .accept = */ nullptr,641 /* .apply = */ llama_sampler_softmax_apply,642 /* .reset = */ nullptr,643 /* .clone = */ nullptr,644 /* .free = */ nullptr,645};646 647struct llama_sampler * llama_sampler_init_softmax() {648 return llama_sampler_init(649 /* .iface = */ &llama_sampler_softmax_i,650 /* .ctx = */ nullptr651 );652}653 654// top-k655 656struct llama_sampler_top_k {657 const int32_t k;658};659 660static const char * llama_sampler_top_k_name(const struct llama_sampler * /*smpl*/) {661 return "top-k";662}663 664static void llama_sampler_top_k_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {665 const auto * ctx = (llama_sampler_top_k *) smpl->ctx;666 llama_sampler_top_k_impl(cur_p, ctx->k);667}668 669static struct llama_sampler * llama_sampler_top_k_clone(const struct llama_sampler * smpl) {670 const auto * ctx = (const llama_sampler_top_k *) smpl->ctx;671 return llama_sampler_init_top_k(ctx->k);672}673 674static void llama_sampler_top_k_free(struct llama_sampler * smpl) {675 delete (llama_sampler_top_k *) smpl->ctx;676}677 678static struct llama_sampler_i llama_sampler_top_k_i = {679 /* .name = */ llama_sampler_top_k_name,680 /* .accept = */ nullptr,681 /* .apply = */ llama_sampler_top_k_apply,682 /* .reset = */ nullptr,683 /* .clone = */ llama_sampler_top_k_clone,684 /* .free = */ llama_sampler_top_k_free,685};686 687struct llama_sampler * llama_sampler_init_top_k(int32_t k) {688 return llama_sampler_init(689 /* .iface = */ &llama_sampler_top_k_i,690 /* .ctx = */ new llama_sampler_top_k {691 /* .k = */ k,692 }693 );694}695 696// top-p697 698struct llama_sampler_top_p {699 const float p;700 const size_t min_keep;701};702 703static const char * llama_sampler_top_p_name(const struct llama_sampler * /*smpl*/) {704 return "top-p";705}706 707static void llama_sampler_top_p_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {708 const auto * ctx = (llama_sampler_top_p *) smpl->ctx;709 710 if (ctx->p >= 1.0f) {711 return;712 }713 714 llama_sampler_softmax_impl(cur_p);715 716 // Compute the cumulative probabilities717 float cum_sum = 0.0f;718 size_t last_idx = cur_p->size;719 720 for (size_t i = 0; i < cur_p->size; ++i) {721 cum_sum += cur_p->data[i].p;722 723 // Check if the running sum is at least p or if we have kept at least min_keep tokens724 // we set the last index to i+1 to indicate that the current iterate should be included in the set725 if (cum_sum >= ctx->p && i + 1 >= ctx->min_keep) {726 last_idx = i + 1;727 break;728 }729 }730 731 // Resize the output vector to keep only the top-p tokens732 cur_p->size = last_idx;733}734 735static struct llama_sampler * llama_sampler_top_p_clone(const struct llama_sampler * smpl) {736 const auto * ctx = (const llama_sampler_top_p *) smpl->ctx;737 return llama_sampler_init_top_p(ctx->p, ctx->min_keep);738}739 740static void llama_sampler_top_p_free(struct llama_sampler * smpl) {741 delete (llama_sampler_top_p *) smpl->ctx;742}743 744static struct llama_sampler_i llama_sampler_top_p_i = {745 /* .name = */ llama_sampler_top_p_name,746 /* .accept = */ nullptr,747 /* .apply = */ llama_sampler_top_p_apply,748 /* .reset = */ nullptr,749 /* .clone = */ llama_sampler_top_p_clone,750 /* .free = */ llama_sampler_top_p_free,751};752 753struct llama_sampler * llama_sampler_init_top_p(float p, size_t min_keep) {754 return llama_sampler_init(755 /* .iface = */ &llama_sampler_top_p_i,756 /* .ctx = */ new llama_sampler_top_p {757 /* .p = */ p,758 /* .min_keep = */ min_keep,759 }760 );761}762 763// min-p764 765struct llama_sampler_min_p {766 const float p;767 const size_t min_keep;768};769 770static const char * llama_sampler_min_p_name(const struct llama_sampler * /*smpl*/) {771 return "min-p";772}773 774static void llama_sampler_min_p_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {775 const auto * ctx = (llama_sampler_min_p *) smpl->ctx;776 777 if (ctx->p <= 0.0f || !cur_p->size) {778 return;779 }780 781 bool min_p_applied = false;782 783 // if the cur_p aren't sorted, try the unsorted implementation first784 if (!cur_p->sorted) {785 std::vector<llama_token_data> filtered_tokens;786 787 float max_logit = -FLT_MAX;788 for (size_t i = 0; i < cur_p->size; ++i) {789 max_logit = std::max(max_logit, cur_p->data[i].logit);790 }791 const float min_logit = max_logit + logf(ctx->p); // min logit for p_i >= p * p_max792 793 for (size_t i = 0; i < cur_p->size; ++i) {794 if (cur_p->data[i].logit >= min_logit) {795 filtered_tokens.push_back(cur_p->data[i]);796 }797 }798 799 // if we have enough values the operation was a success800 if (filtered_tokens.size() >= ctx->min_keep) {801 memcpy(cur_p->data, filtered_tokens.data(), filtered_tokens.size()*sizeof(llama_token_data));802 cur_p->size = filtered_tokens.size();803 min_p_applied = true;804 }805 }806 807 // if the cur_p are sorted or the unsorted implementation failed, use this implementation808 if (!min_p_applied) {809 // Sort the logits in descending order810 if (!cur_p->sorted) {811 std::sort(cur_p->data, cur_p->data + cur_p->size, [](const llama_token_data & a, const llama_token_data & b) {812 return a.logit > b.logit;813 });814 cur_p->sorted = true;815 }816 817 const float min_logit = cur_p->data[0].logit + logf(ctx->p); // min logit for p_i >= p * p_max818 size_t i = 1; // first token always matches819 820 for (; i < cur_p->size; ++i) {821 if (cur_p->data[i].logit < min_logit && i >= ctx->min_keep) {822 break; // prob too small823 }824 }825 826 // Resize the output vector to keep only the matching tokens827 cur_p->size = i;828 }829}830 831static struct llama_sampler * llama_sampler_min_p_clone(const struct llama_sampler * smpl) {832 const auto * ctx = (const llama_sampler_min_p *) smpl->ctx;833 return llama_sampler_init_min_p(ctx->p, ctx->min_keep);834}835 836static void llama_sampler_min_p_free(struct llama_sampler * smpl) {837 delete (llama_sampler_min_p *) smpl->ctx;838}839 840static struct llama_sampler_i llama_sampler_min_p_i = {841 /* .name = */ llama_sampler_min_p_name,842 /* .accept = */ nullptr,843 /* .apply = */ llama_sampler_min_p_apply,844 /* .reset = */ nullptr,845 /* .clone = */ llama_sampler_min_p_clone,846 /* .free = */ llama_sampler_min_p_free,847};848 849struct llama_sampler * llama_sampler_init_min_p(float p, size_t min_keep) {850 return llama_sampler_init(851 /* .iface = */ &llama_sampler_min_p_i,852 /* .ctx = */ new llama_sampler_min_p {853 /* .p = */ p,854 /* .min_keep = */ min_keep,855 }856 );857}858 859// typical860 861struct llama_sampler_typical {862 const float p;863 const size_t min_keep;864};865 866static const char * llama_sampler_typical_name(const struct llama_sampler * /*smpl*/) {867 return "typical";868}869 870static void llama_sampler_typical_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {871 const auto * ctx = (llama_sampler_typical *) smpl->ctx;872 873 // Reference implementation:874 // https://github.com/huggingface/transformers/compare/main...cimeister:typical-sampling:typical-pr875 if (ctx->p >= 1.0f) {876 return;877 }878 879 // Compute the softmax of logits and calculate entropy880 llama_sampler_softmax_impl(cur_p);881 882 float entropy = 0.0f;883 for (size_t i = 0; i < cur_p->size; ++i) {884 entropy += -cur_p->data[i].p * logf(cur_p->data[i].p);885 }886 887 // Compute the absolute difference between negative log probability and entropy for each candidate888 std::vector<float> shifted_scores;889 for (size_t i = 0; i < cur_p->size; ++i) {890 float shifted_score = fabsf(-logf(cur_p->data[i].p) - entropy);891 shifted_scores.push_back(shifted_score);892 }893 894 // Sort tokens based on the shifted_scores and their corresponding indices895 std::vector<size_t> indices(cur_p->size);896 std::iota(indices.begin(), indices.end(), 0);897 898 std::sort(indices.begin(), indices.end(), [&](size_t a, size_t b) {899 return shifted_scores[a] < shifted_scores[b];900 });901 902 // Compute the cumulative probabilities903 float cum_sum = 0.0f;904 size_t last_idx = indices.size();905 906 for (size_t i = 0; i < indices.size(); ++i) {907 size_t idx = indices[i];908 cum_sum += cur_p->data[idx].p;909 910 // Check if the running sum is greater than typical or if we have kept at least min_keep tokens911 if (cum_sum > ctx->p && i >= ctx->min_keep - 1) {912 last_idx = i + 1;913 break;914 }915 }916 917 // Resize the output vector to keep only the locally typical tokens918 std::vector<llama_token_data> cur_p_new;919 for (size_t i = 0; i < last_idx; ++i) {920 size_t idx = indices[i];921 cur_p_new.push_back(cur_p->data[idx]);922 }923 924 // Replace the data in cur_p with the cur_p_new data925 std::copy(cur_p_new.begin(), cur_p_new.end(), cur_p->data);926 cur_p->size = cur_p_new.size();927 cur_p->sorted = false;928}929 930static struct llama_sampler * llama_sampler_typical_clone(const struct llama_sampler * smpl) {931 const auto * ctx = (const llama_sampler_typical *) smpl->ctx;932 return llama_sampler_init_typical(ctx->p, ctx->min_keep);933}934 935static void llama_sampler_typical_free(struct llama_sampler * smpl) {936 delete (llama_sampler_typical *) smpl->ctx;937}938 939static struct llama_sampler_i llama_sampler_typical_i = {940 /* .name = */ llama_sampler_typical_name,941 /* .accept = */ nullptr,942 /* .apply = */ llama_sampler_typical_apply,943 /* .reset = */ nullptr,944 /* .clone = */ llama_sampler_typical_clone,945 /* .free = */ llama_sampler_typical_free,946};947 948struct llama_sampler * llama_sampler_init_typical(float p, size_t min_keep) {949 return llama_sampler_init(950 /* .iface = */ &llama_sampler_typical_i,951 /* .ctx = */ new llama_sampler_typical {952 /* .p = */ p,953 /* .min_keep = */ min_keep,954 }955 );956}957 958// temp959 960struct llama_sampler_temp {961 const float temp;962};963 964static const char * llama_sampler_temp_name(const struct llama_sampler * /*smpl*/) {965 return "temp";966}967 968static void llama_sampler_temp_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {969 const auto * ctx = (llama_sampler_temp *) smpl->ctx;970 971 llama_sampler_temp_impl(cur_p, ctx->temp);972}973 974static struct llama_sampler * llama_sampler_temp_clone(const struct llama_sampler * smpl) {975 const auto * ctx = (const llama_sampler_temp *) smpl->ctx;976 return llama_sampler_init_temp(ctx->temp);977}978 979static void llama_sampler_temp_free(struct llama_sampler * smpl) {980 delete (llama_sampler_temp *) smpl->ctx;981}982 983static struct llama_sampler_i llama_sampler_temp_i = {984 /* .name = */ llama_sampler_temp_name,985 /* .accept = */ nullptr,986 /* .apply = */ llama_sampler_temp_apply,987 /* .reset = */ nullptr,988 /* .clone = */ llama_sampler_temp_clone,989 /* .free = */ llama_sampler_temp_free,990};991 992struct llama_sampler * llama_sampler_init_temp(float temp) {993 return llama_sampler_init(994 /* .iface = */ &llama_sampler_temp_i,995 /* .ctx = */ new llama_sampler_temp {996 /*.temp = */ temp,997 }998 );999}1000 1001// temp-ext1002 1003struct llama_sampler_temp_ext {1004 const float temp;1005 const float delta;1006 const float exponent;1007};1008 1009static const char * llama_sampler_temp_ext_name(const struct llama_sampler * /*smpl*/) {1010 return "temp-ext";1011}1012 1013static void llama_sampler_temp_ext_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {1014 const auto * ctx = (llama_sampler_temp_ext *) smpl->ctx;1015 if (ctx->delta > 0) {1016 const float min_temp = std::max(0.0f, ctx->temp - ctx->delta);1017 const float max_temp = ctx->temp + ctx->delta;1018 1019 float exponent_val = ctx->exponent;1020 1021 // no need to do anything if there is only one (or zero) candidates1022 if (cur_p->size <= 1) {1023 return;1024 }1025 1026 // Calculate maximum possible entropy1027 float max_entropy = -logf(1.0f / cur_p->size);1028 1029 llama_sampler_softmax_impl(cur_p);1030 1031 // Calculate entropy of the softmax probabilities1032 float entropy = 0.0f;1033 for (size_t i = 0; i < cur_p->size; ++i) {1034 float prob = cur_p->data[i].p;1035 if (prob > 0.0f) { // Ensure no log(0)1036 entropy -= prob * logf(prob);1037 }1038 }1039 1040 // Normalize the entropy (max_entropy cannot be 0 here because we checked cur_p->size != 1 above)1041 float normalized_entropy = entropy / max_entropy;1042 1043 // Map the normalized entropy to the desired temperature range using the power function1044 float dyn_temp = min_temp + (max_temp - min_temp) * powf(normalized_entropy, exponent_val);1045 1046 #ifdef DEBUG1047 LLAMA_LOG_INFO("Your text maxtemp value is: %f\n", max_temp);1048 LLAMA_LOG_INFO("Entropy: %f\n", entropy);1049 LLAMA_LOG_INFO("Max Possible Entropy: %f\n", max_entropy);1050 LLAMA_LOG_INFO("Normalized Entropy: %f\n", normalized_entropy);1051 LLAMA_LOG_INFO("Exponent: %f\n", exponent_val);1052 LLAMA_LOG_INFO("Dynamic Temperature (dyn_temp): %f\n", dyn_temp);1053 #endif1054 1055 // Apply the dynamically calculated temperature scaling1056 llama_sampler_temp_impl(cur_p, dyn_temp);1057 1058 // Re-compute softmax probabilities after scaling logits with dynamic temperature1059 const double max_l_double = cur_p->data[0].logit;1060 1061 double cum_sum_double = 0.0;1062 for (size_t i = 0; i < cur_p->size; ++i) {1063 double p = exp(cur_p->data[i].logit - max_l_double);1064 cur_p->data[i].p = p; // Store the scaled probability1065 cum_sum_double += p;1066 }1067 1068 for (size_t i = 0; i < cur_p->size; ++i) {1069 cur_p->data[i].p /= cum_sum_double; // Re-normalize the probabilities1070 }1071 1072 #ifdef DEBUG1073 // Print the updated top 25 probabilities after temperature scaling1074 LLAMA_LOG_INFO("\nUpdated Top 25 Probabilities After Dynamic Temperature Scaling (in percentages):\n");1075 for (size_t i = 0; i < 25 && i < cur_p->size; ++i) {1076 LLAMA_LOG_INFO("Token %zu: %f%%\n", i + 1, cur_p->data[i].p * 100.0f);1077 }1078 #endif1079 } else {1080 llama_sampler_temp_impl(cur_p, ctx->temp);1081 }1082}1083 1084static struct llama_sampler * llama_sampler_temp_ext_clone(const struct llama_sampler * smpl) {1085 const auto * ctx = (const llama_sampler_temp_ext *) smpl->ctx;1086 return llama_sampler_init_temp_ext(ctx->temp, ctx->delta, ctx->exponent);1087}1088 1089static void llama_sampler_temp_ext_free(struct llama_sampler * smpl) {1090 delete (llama_sampler_temp_ext *) smpl->ctx;1091}1092 1093static struct llama_sampler_i llama_sampler_temp_ext_i = {1094 /* .name = */ llama_sampler_temp_ext_name,1095 /* .accept = */ nullptr,1096 /* .apply = */ llama_sampler_temp_ext_apply,1097 /* .reset = */ nullptr,1098 /* .clone = */ llama_sampler_temp_ext_clone,1099 /* .free = */ llama_sampler_temp_ext_free,1100};1101 1102struct llama_sampler * llama_sampler_init_temp_ext(float temp, float delta, float exponent) {1103 return llama_sampler_init(1104 /* .iface = */ &llama_sampler_temp_ext_i,1105 /* .ctx = */ new llama_sampler_temp_ext {1106 /* .temp = */ temp,1107 /* .delta = */ delta,1108 /* .exponent = */ exponent,1109 }1110 );1111}1112 1113// xtc1114 1115struct llama_sampler_xtc {1116 const float probability;1117 const float threshold;1118 const size_t min_keep;1119 1120 const uint32_t seed;1121 uint32_t seed_cur;1122 1123 std::mt19937 rng;1124};1125 1126static const char * llama_sampler_xtc_name(const struct llama_sampler * /*smpl*/) {1127 return "xtc";1128}1129 1130static void llama_sample_xtc_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {1131 auto * ctx = (llama_sampler_xtc *) smpl->ctx;1132 1133 if (ctx->probability <= 0.0f1134 || ctx->threshold > 0.5f1135 || cur_p->size < 2) {1136 return;1137 }1138 1139 std::uniform_real_distribution<float> distribution(0.0f, 1.0f);1140 float chance = distribution(ctx->rng);1141 if (chance > ctx->probability) return;1142 1143 // in case it's not sorted/recalculated yet1144 llama_sampler_softmax_impl(cur_p);1145 1146 int pos_last = 0;1147 1148 for (size_t i = 0; i < cur_p->size; ++i) {1149 if (cur_p->data[i].p >= ctx->threshold) {1150 pos_last = i;1151 } else break;1152 }1153 1154 if (cur_p->size - pos_last >= ctx->min_keep && pos_last > 0) {1155 cur_p->data += pos_last;1156 cur_p->size -= pos_last;1157 }1158}1159 1160static struct llama_sampler * llama_sampler_xtc_clone(const struct llama_sampler * smpl) {1161 const auto * ctx = (const llama_sampler_xtc *) smpl->ctx;1162 auto * result = llama_sampler_init_xtc(ctx->probability, ctx->threshold, ctx->min_keep, ctx->seed);1163 1164 // copy the state1165 {1166 auto * result_ctx = (llama_sampler_xtc *) result->ctx;1167 1168 result_ctx->rng = ctx->rng;1169 }1170 1171 return result;1172}1173 1174static void llama_sampler_xtc_free(struct llama_sampler * smpl) {1175 delete (llama_sampler_xtc *) smpl->ctx;1176}1177 1178static void llama_sampler_xtc_reset(struct llama_sampler * smpl) {1179 auto * ctx = (llama_sampler_xtc *) smpl->ctx;1180 ctx->seed_cur = get_rng_seed(ctx->seed);1181 ctx->rng.seed(ctx->seed_cur);1182}1183 1184static struct llama_sampler_i llama_sampler_xtc_i = {1185 /* .name = */ llama_sampler_xtc_name,1186 /* .accept = */ nullptr,1187 /* .apply = */ llama_sample_xtc_apply,1188 /* .reset = */ llama_sampler_xtc_reset,1189 /* .clone = */ llama_sampler_xtc_clone,1190 /* .free = */ llama_sampler_xtc_free,1191};1192 1193struct llama_sampler * llama_sampler_init_xtc(float p, float t, size_t min_keep, uint32_t seed) {1194 auto seed_cur = get_rng_seed(seed);1195 return llama_sampler_init(1196 /* .iface = */ &llama_sampler_xtc_i,1197 /* .ctx = */ new llama_sampler_xtc {1198 /* .probability = */ p,1199 /* .threshold = */ t,1200 /* .min_keep = */ min_keep,