39 lines
1.6 KiB
C++
39 lines
1.6 KiB
C++
// 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 <cstdint>
|
|
#include <vector>
|
|
#include <deque>
|
|
|
|
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<int> 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);
|