diff options
author | Alexey Parfenov <zxed@alkatrazstudio.net> | 2023-12-23 09:31:49 +0000 |
---|---|---|
committer | GitHub <noreply@github.com> | 2023-12-23 11:31:49 +0200 |
commit | 6123979952385847d8348e295d77d6e01da8aa84 (patch) | |
tree | 2d536d31ef7e1b6f07468ff6b13710da1fe4f732 /common/sampling.cpp | |
parent | b9ec82d262cb20d7f0a8a1157bfa9aace40e2625 (diff) |
server : allow to specify custom prompt for penalty calculation (#3727)
Diffstat (limited to 'common/sampling.cpp')
-rw-r--r-- | common/sampling.cpp | 8 |
1 files changed, 5 insertions, 3 deletions
diff --git a/common/sampling.cpp b/common/sampling.cpp index 5b15204b..8e45909f 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -203,12 +203,14 @@ static llama_token llama_sampling_sample_impl( } // apply penalties - if (!prev.empty()) { + const auto& penalty_tokens = params.use_penalty_prompt_tokens ? params.penalty_prompt_tokens : prev; + const int penalty_tokens_used_size = std::min((int)penalty_tokens.size(), penalty_last_n); + if (penalty_tokens_used_size) { const float nl_logit = logits[llama_token_nl(llama_get_model(ctx_main))]; llama_sample_repetition_penalties(ctx_main, &cur_p, - prev.data() + prev.size() - penalty_last_n, - penalty_last_n, penalty_repeat, penalty_freq, penalty_present); + penalty_tokens.data() + penalty_tokens.size() - penalty_tokens_used_size, + penalty_tokens_used_size, penalty_repeat, penalty_freq, penalty_present); if (!penalize_nl) { for (size_t idx = 0; idx < cur_p.size; idx++) { |