Ling-3.0-tiny-RKNN / include /ling3 /chat_sampling.h
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
1.7 kB
#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)];
}
}