File size: 1,697 Bytes
3fd1a35
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
#pragma once
#include "ling3/chat_protocol.h"
#include <numeric>
#include <random>
#include <span>

namespace ling3::chat {
inline std::uint32_t SampleToken(std::span<const float> logits, const Request & r,
                                 const std::vector<bool> & seen, std::mt19937 & random) {
    if (logits.empty()) throw std::runtime_error("empty logits");
    std::vector<double> scores(logits.begin(), logits.end());
    for (std::size_t i = 0; i < scores.size(); ++i) {
        if (!std::isfinite(scores[i])) throw std::runtime_error("non-finite logits");
        if (i < seen.size() && seen[i])
            scores[i] = scores[i] < 0 ? scores[i] * r.repeat_penalty : scores[i] / r.repeat_penalty;
    }
    const auto best = std::max_element(scores.begin(), scores.end());
    if (r.temperature == 0 || r.top_k == 1) return best - scores.begin();
    std::vector<std::uint32_t> order(scores.size());
    std::iota(order.begin(), order.end(), 0);
    const auto n = r.top_k > 0 ? std::min<std::size_t>(r.top_k, order.size()) : order.size();
    std::partial_sort(order.begin(), order.begin() + n, order.end(), [&](auto a, auto b) {
        return scores[a] == scores[b] ? a < b : scores[a] > scores[b];
    });
    order.resize(n);
    std::vector<double> weights(n);
    double sum = 0;
    for (std::size_t i = 0; i < n; ++i) {
        weights[i] = std::exp((scores[order[i]] - *best) / r.temperature); sum += weights[i];
    }
    double cumulative = 0;
    std::size_t count = 0;
    do { cumulative += weights[count++]; } while (count < n && cumulative < sum * r.top_p);
    return order[std::discrete_distribution<std::size_t>(weights.begin(), weights.begin() + count)(random)];
}
}