Download include/ling3/chat_sampling.h from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 1.7 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/include/ling3/chat_sampling.h
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/include/ling3/chat_sampling.h
-
curl -L -o chat_sampling.h https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/include/ling3/chat_sampling.h
1.7 kB
| 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)]; | |
| } | |
| } | |