// Sampler logits compatible HuggingFace generation utilities : // apply_repetition_penalty (HF style : divise si logit>0, multiplie si logit<0) // apply_temperature (logit /= temp) // apply_top_k (garde les top_k logits) // apply_top_p (nucleus sampling : garde la masse cumulée jusqu'à top_p) // multinomial_sample (sample selon softmax des logits restants) // // PRNG xorshift32 déterministe (sample_seed). Le RNG diffère de torch.mt19937 // donc même seed ne reproduit PAS torch, mais évite les attracteurs du greedy // sur n'importe quel texte/voix. #pragma once #include #include #include struct Sampler { float temp = 0.9f; int top_k = 50; float top_p = 1.0f; float rep_penalty = 1.05f; int rep_window = 64; // nb de tokens récents à pénaliser uint32_t rng_state = 1u; // historique des tokens samplés (pour rep_penalty), ring-buffer de taille rep_window std::deque recent; }; void sampler_seed(Sampler& s, uint32_t seed); void sampler_reset_history(Sampler& s); // Sample un token depuis `logits` (modifié en place : on lui applique penalty/temp/top_k/top_p) // vocab = taille de logits. Met à jour s.recent. Retourne l'id sampled. int sampler_sample(Sampler& s, float* logits, int vocab); // Variante pour intra-step CP (rep_penalty sur les tokens déjà samplés CE STEP, pas historique long). // `recent_step` : pointeur vers les tokens déjà samplés ce step (NULL ou len=0 pour aucun). int sampler_sample_local(Sampler& s, float* logits, int vocab, const int* recent_step, int n_recent);