Kazeia-engine/dist/jni/sampler.h

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);