#pragma once #include "ling3/chat_protocol.h" #include #include #include namespace ling3::chat { inline std::uint32_t SampleToken(std::span logits, const Request & r, const std::vector & seen, std::mt19937 & random) { if (logits.empty()) throw std::runtime_error("empty logits"); std::vector 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 order(scores.size()); std::iota(order.begin(), order.end(), 0); const auto n = r.top_k > 0 ? std::min(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 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(weights.begin(), weights.begin() + count)(random)]; } }