| 100 | } |
| 101 | |
| 102 | llama_token llama_sampling_sample( |
| 103 | struct llama_sampling_context * ctx_sampling, |
| 104 | struct llama_context * ctx_main, |
| 105 | struct llama_context * ctx_cfg, |
| 106 | const int idx) { |
| 107 | const llama_sampling_params & params = ctx_sampling->params; |
| 108 | |
| 109 | const int n_vocab = llama_n_vocab(llama_get_model(ctx_main)); |
| 110 | |
| 111 | const float temp = params.temp; |
| 112 | const int32_t top_k = params.top_k <= 0 ? n_vocab : params.top_k; |
| 113 | const float top_p = params.top_p; |
| 114 | const float min_p = params.min_p; |
| 115 | const float tfs_z = params.tfs_z; |
| 116 | const float typical_p = params.typical_p; |
| 117 | const int32_t penalty_last_n = params.penalty_last_n < 0 ? params.n_prev : params.penalty_last_n; |
| 118 | const float penalty_repeat = params.penalty_repeat; |
| 119 | const float penalty_freq = params.penalty_freq; |
| 120 | const float penalty_present = params.penalty_present; |
| 121 | const int mirostat = params.mirostat; |
| 122 | const float mirostat_tau = params.mirostat_tau; |
| 123 | const float mirostat_eta = params.mirostat_eta; |
| 124 | const bool penalize_nl = params.penalize_nl; |
| 125 | |
| 126 | auto & prev = ctx_sampling->prev; |
| 127 | auto & cur = ctx_sampling->cur; |
| 128 | |
| 129 | llama_token id = 0; |
| 130 | |
| 131 | float * logits = llama_get_logits_ith(ctx_main, idx); |
| 132 | |
| 133 | // apply params.logit_bias map |
| 134 | for (auto it = params.logit_bias.begin(); it != params.logit_bias.end(); it++) { |
| 135 | logits[it->first] += it->second; |
| 136 | } |
| 137 | |
| 138 | cur.clear(); |
| 139 | |
| 140 | for (llama_token token_id = 0; token_id < n_vocab; token_id++) { |
| 141 | cur.emplace_back(llama_token_data{token_id, logits[token_id], 0.0f}); |
| 142 | } |
| 143 | |
| 144 | llama_token_data_array cur_p = { cur.data(), cur.size(), false }; |
| 145 | |
| 146 | if (ctx_cfg) { |
| 147 | llama_sample_classifier_free_guidance(ctx_main, &cur_p, ctx_cfg, params.cfg_scale); |
| 148 | } |
| 149 | |
| 150 | // apply penalties |
| 151 | if (!prev.empty()) { |
| 152 | const float nl_logit = logits[llama_token_nl(llama_get_model(ctx_main))]; |
| 153 | |
| 154 | llama_sample_repetition_penalties(ctx_main, &cur_p, |
| 155 | prev.data() + prev.size() - penalty_last_n, |
| 156 | penalty_last_n, penalty_repeat, penalty_freq, penalty_present); |
| 157 | |
| 158 | if (!penalize_nl) { |
| 159 | for (size_t idx = 0; idx < cur_p.size; idx++) { |