#include "sampler.h" #include #include #include void sampler_seed(Sampler& s, uint32_t seed) { s.rng_state = seed ? seed : 1u; s.recent.clear(); } void sampler_reset_history(Sampler& s) { s.recent.clear(); } static inline float xs_rand(uint32_t& st) { st ^= st << 13; st ^= st >> 17; st ^= st << 5; return (st & 0x00FFFFFF) / (float)0x01000000; } // HF transformers RepetitionPenaltyLogitsProcessor : // if logit > 0 : logit /= penalty // else : logit *= penalty static void apply_rep_penalty_(float* logits, int vocab, const int* tokens, int n_tokens, float penalty) { if (penalty == 1.0f || n_tokens == 0) return; for (int i = 0; i < n_tokens; i++) { int t = tokens[i]; if (t < 0 || t >= vocab) continue; float v = logits[t]; logits[t] = (v > 0.0f) ? (v / penalty) : (v * penalty); } } // Garde les top_k logits, met les autres à -INFINITY. // Si top_k <= 0 ou >= vocab, no-op. static void apply_top_k_(float* logits, int vocab, int top_k) { if (top_k <= 0 || top_k >= vocab) return; std::vector sorted(logits, logits + vocab); std::nth_element(sorted.begin(), sorted.begin() + (vocab - top_k), sorted.end()); float threshold = sorted[vocab - top_k]; for (int i = 0; i < vocab; ++i) if (logits[i] < threshold) logits[i] = -INFINITY; } // Nucleus sampling : sort logits desc, garde la masse cumsum jusqu'à top_p, met le reste à -inf. // Garde au moins 1 token. static void apply_top_p_(float* logits, int vocab, float top_p) { if (top_p >= 1.0f || top_p <= 0.0f) return; std::vector idx(vocab); for (int i = 0; i < vocab; ++i) idx[i] = i; std::sort(idx.begin(), idx.end(), [&](int a, int b){ return logits[a] > logits[b]; }); // softmax pour cumsum (stable : soustrait max) float mx = logits[idx[0]]; std::vector probs(vocab); double sum = 0.0; for (int i = 0; i < vocab; ++i) { float v = std::exp(logits[idx[i]] - mx); probs[i] = v; sum += v; } if (sum <= 0) return; double inv = 1.0 / sum; double cum = 0.0; int cutoff = 0; for (int i = 0; i < vocab; ++i) { cum += probs[i] * inv; if (cum >= top_p) { cutoff = i; break; } } // garde [0..cutoff] (inclus), met le reste à -inf for (int i = cutoff + 1; i < vocab; ++i) logits[idx[i]] = -INFINITY; } static int multinomial_(float* logits, int vocab, uint32_t& rng_state) { // softmax + sample float mx = -INFINITY; for (int i = 0; i < vocab; ++i) if (logits[i] > mx) mx = logits[i]; if (!std::isfinite(mx)) return 0; double sum = 0.0; std::vector w(vocab); for (int i = 0; i < vocab; ++i) { if (std::isfinite(logits[i])) { w[i] = std::exp(logits[i] - mx); sum += w[i]; } else { w[i] = 0.0; } } double r = xs_rand(rng_state) * sum; double acc = 0.0; for (int i = 0; i < vocab; ++i) { acc += w[i]; if (r <= acc) return i; } // fallback : argmax (au cas où numérique) int b = 0; for (int i = 1; i < vocab; ++i) if (logits[i] > logits[b]) b = i; return b; } int sampler_sample(Sampler& s, float* logits, int vocab) { // rep_penalty depuis l'historique (Talker = phrase entière) if (s.rep_penalty != 1.0f && !s.recent.empty()) { std::vector recent_v(s.recent.begin(), s.recent.end()); apply_rep_penalty_(logits, vocab, recent_v.data(), (int)recent_v.size(), s.rep_penalty); } // temperature if (s.temp > 0.0f && s.temp != 1.0f) for (int i = 0; i < vocab; ++i) logits[i] /= s.temp; apply_top_k_(logits, vocab, s.top_k); apply_top_p_(logits, vocab, s.top_p); int tok = multinomial_(logits, vocab, s.rng_state); s.recent.push_back(tok); while ((int)s.recent.size() > s.rep_window) s.recent.pop_front(); return tok; } int sampler_sample_local(Sampler& s, float* logits, int vocab, const int* recent_step, int n_recent) { if (s.rep_penalty != 1.0f && n_recent > 0) apply_rep_penalty_(logits, vocab, recent_step, n_recent, s.rep_penalty); if (s.temp > 0.0f && s.temp != 1.0f) for (int i = 0; i < vocab; ++i) logits[i] /= s.temp; apply_top_k_(logits, vocab, s.top_k); apply_top_p_(logits, vocab, s.top_p); return multinomial_(logits, vocab, s.rng_state); }