chantier B TTS #7: sampler HF-style top_k+top_p+rep_penalty (générique)
sampler.{h,cpp} : module isolé reproduisant la chaine HF
LogitsProcessor :
1. RepetitionPenaltyLogitsProcessor (logit/penalty si >0, logit*penalty sinon)
2. TemperatureLogitsWarper (logit /= temp)
3. TopKLogitsWarper (garde top_k, autres -> -inf)
4. TopPLogitsWarper (nucleus : garde la masse cumulee jusqu'a top_p)
5. multinomial sample (xorshift32, deterministe via sampler_seed)
PRNG xorshift32 != torch.mt19937 donc meme seed != audio Python, mais stable
sur n'importe quel texte/voix (pas d'attracteur greedy).
Plumbing :
- CPState.sampler (defaut greedy top_k=1, temp=0 -> cp_validate continue a sortir
495/495 codes match). Pipeline override pour temp=0.9 top_k=50 rep_penalty=1.05
rep_window=16 (intra-frame, 15 codebooks max). CB1..15 utilisent
sampler_sample_local() qui prend les tokens deja samples CE step pour penaliser
repetition intra-frame.
- tts_pipeline : Sampler talker_sampler avec rep_window=64 (historique long sur
CB0). Override via env KZTTS_TEMP/TOPK/TOPP/REPP + KZTTS_CP_TEMP/TOPK/TOPP/REPP.
Validation sur 5 seeds (1, 42, 100, 777, 12345), phrase 'Bonjour je m'appelle Kazeia' :
- Tous EOS naturel (step 31-44, 2.48-3.67s audio @12Hz, proche du PY_REF 2.08s)
- silence_pct : 25-36% (vs > 60% avec greedy/sampler simple)
- cp_validate continue a sortir 495/495 codes match (greedy par defaut)
Vs WAV pipeline_v3 user-confirme 'parfait' (seed=123 temp=0.8 sans rep_penalty),
le sampler v4 est generique sur tout seed/texte, plus de besoin de chercher un
seed favorable. Trade-off : on n'a plus le PRNG torch donc l'audio diffère de
PyTorch a chaque run (different trajectoire), mais reste audio TTS coherent.
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
9ef5e11b53
commit
d20f2b297b
|
|
@ -165,30 +165,36 @@ static void cp_forward_lastpos(CPState& s, const std::vector<float>& embeds_flat
|
||||||
|
|
||||||
void cp_predict(CPState& s, const float* hidden, const float* cb0_emb, int32_t* out_codes) {
|
void cp_predict(CPState& s, const float* hidden, const float* cb0_emb, int32_t* out_codes) {
|
||||||
// embeds = [hidden, cb0_emb, codec_embs[0][cb1], codec_embs[1][cb2], ..., codec_embs[13][cb14]]
|
// embeds = [hidden, cb0_emb, codec_embs[0][cb1], codec_embs[1][cb2], ..., codec_embs[13][cb14]]
|
||||||
// step s (1..15) : forward sur les (s+1) tokens, sortir hidden à la position `s`, head[s-1] -> argmax.
|
// step s (1..15) : forward sur les (s+1) tokens, sortir hidden à la position `s`, head[s-1] -> sample.
|
||||||
std::vector<float> embeds; embeds.reserve((size_t)17 * N_EMBD);
|
std::vector<float> embeds; embeds.reserve((size_t)17 * N_EMBD);
|
||||||
embeds.insert(embeds.end(), hidden, hidden + N_EMBD);
|
embeds.insert(embeds.end(), hidden, hidden + N_EMBD);
|
||||||
embeds.insert(embeds.end(), cb0_emb, cb0_emb + N_EMBD);
|
embeds.insert(embeds.end(), cb0_emb, cb0_emb + N_EMBD);
|
||||||
|
|
||||||
std::vector<float> hn;
|
std::vector<float> hn;
|
||||||
|
std::vector<float> scores(N_VOCAB);
|
||||||
|
int local_history[N_CB]; // tokens samplés ce step (pour rep_penalty intra-frame)
|
||||||
|
int n_local = 0;
|
||||||
for (int step = 1; step <= N_CB; step++) {
|
for (int step = 1; step <= N_CB; step++) {
|
||||||
const int L = step + 1;
|
const int L = step + 1;
|
||||||
cp_forward_lastpos(s, embeds, L, step, hn);
|
cp_forward_lastpos(s, embeds, L, step, hn);
|
||||||
|
|
||||||
// tête[step-1] : argmax_j sum_k hn[k] * heads[step-1, j, k]
|
// tête[step-1] : scores[j] = sum_k hn[k] * heads[step-1, j, k]
|
||||||
const float* Wh = &s.heads[(size_t)(step - 1) * N_VOCAB * N_EMBD];
|
const float* Wh = &s.heads[(size_t)(step - 1) * N_VOCAB * N_EMBD];
|
||||||
int best = 0; float bv = -1e30f;
|
|
||||||
for (int j = 0; j < N_VOCAB; j++) {
|
for (int j = 0; j < N_VOCAB; j++) {
|
||||||
const float* wj = Wh + (size_t)j * N_EMBD;
|
const float* wj = Wh + (size_t)j * N_EMBD;
|
||||||
float dot = 0.f;
|
float dot = 0.f;
|
||||||
for (int k = 0; k < N_EMBD; k++) dot += hn[k] * wj[k];
|
for (int k = 0; k < N_EMBD; k++) dot += hn[k] * wj[k];
|
||||||
if (dot > bv) { bv = dot; best = j; }
|
scores[j] = dot;
|
||||||
}
|
}
|
||||||
out_codes[step - 1] = best;
|
// Sample via sampler partagé (HF-style). rep_penalty appliquée sur les tokens déjà
|
||||||
|
// émis CE step (intra-frame), pas l'historique long.
|
||||||
|
int chosen = sampler_sample_local(s.sampler, scores.data(), N_VOCAB,
|
||||||
|
local_history, n_local);
|
||||||
|
out_codes[step - 1] = chosen;
|
||||||
|
local_history[n_local++] = chosen;
|
||||||
|
|
||||||
// feedback : append codec_embs[step-1][best] aux embeds
|
|
||||||
if (step < N_CB) {
|
if (step < N_CB) {
|
||||||
const float* e = &s.codec_embs[(size_t)(step - 1) * N_VOCAB * N_EMBD + (size_t)best * N_EMBD];
|
const float* e = &s.codec_embs[(size_t)(step - 1) * N_VOCAB * N_EMBD + (size_t)chosen * N_EMBD];
|
||||||
embeds.insert(embeds.end(), e, e + N_EMBD);
|
embeds.insert(embeds.end(), e, e + N_EMBD);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -8,6 +8,7 @@
|
||||||
// (hidden_Talker[1024], cb0_emb[1024]). Approche "recompute" (T<=16 -> coût
|
// (hidden_Talker[1024], cb0_emb[1024]). Approche "recompute" (T<=16 -> coût
|
||||||
// quadratique négligeable, pas de KV cache manuel).
|
// quadratique négligeable, pas de KV cache manuel).
|
||||||
#pragma once
|
#pragma once
|
||||||
|
#include "sampler.h"
|
||||||
#include <cstdint>
|
#include <cstdint>
|
||||||
#include <string>
|
#include <string>
|
||||||
#include <vector>
|
#include <vector>
|
||||||
|
|
@ -22,6 +23,11 @@ struct CPState {
|
||||||
std::vector<float> heads; // [15, 2048, 1024] f32
|
std::vector<float> heads; // [15, 2048, 1024] f32
|
||||||
std::vector<float> codec_embs; // [15, 2048, 1024] f32
|
std::vector<float> codec_embs; // [15, 2048, 1024] f32
|
||||||
int n_threads = 4;
|
int n_threads = 4;
|
||||||
|
// Sampler HF-style (rep_penalty + top_k + top_p + temperature). Par défaut = greedy
|
||||||
|
// (top_k=1 -> argmax). Pour TTS prod, configurer en {temp=0.9, top_k=50, rep_penalty=1.05}
|
||||||
|
// ce qui matche Python subtalker_*.
|
||||||
|
Sampler sampler = { /*temp=*/0.0f, /*top_k=*/1, /*top_p=*/1.0f, /*rep_penalty=*/1.0f,
|
||||||
|
/*rep_window=*/16, /*rng_state=*/1u, {} };
|
||||||
};
|
};
|
||||||
|
|
||||||
// Charge cp_f16.gguf + cp_heads.bin + cp_codec_embs.bin. Retourne false en cas d'échec.
|
// Charge cp_f16.gguf + cp_heads.bin + cp_codec_embs.bin. Retourne false en cas d'échec.
|
||||||
|
|
@ -35,3 +41,4 @@ void cp_free(CPState& s);
|
||||||
// cb0_emb: float[1024]
|
// cb0_emb: float[1024]
|
||||||
// out_codes : int32_t[15] (CB1..CB15)
|
// out_codes : int32_t[15] (CB1..CB15)
|
||||||
void cp_predict(CPState& s, const float* hidden, const float* cb0_emb, int32_t* out_codes);
|
void cp_predict(CPState& s, const float* hidden, const float* cb0_emb, int32_t* out_codes);
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,111 @@
|
||||||
|
#include "sampler.h"
|
||||||
|
#include <algorithm>
|
||||||
|
#include <cmath>
|
||||||
|
#include <vector>
|
||||||
|
|
||||||
|
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<float> 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<int> 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<float> 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<double> 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<int> 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);
|
||||||
|
}
|
||||||
|
|
@ -0,0 +1,38 @@
|
||||||
|
// 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);
|
||||||
|
|
@ -23,6 +23,7 @@
|
||||||
#include "llama.h"
|
#include "llama.h"
|
||||||
#include "ggml-backend.h"
|
#include "ggml-backend.h"
|
||||||
#include "cp_inference.h"
|
#include "cp_inference.h"
|
||||||
|
#include "sampler.h"
|
||||||
#include "decoder.h" // Kazeia decoder ggml (linké via libqwen3tts-decoder.a)
|
#include "decoder.h" // Kazeia decoder ggml (linké via libqwen3tts-decoder.a)
|
||||||
|
|
||||||
using namespace kazeia::tts;
|
using namespace kazeia::tts;
|
||||||
|
|
@ -54,40 +55,6 @@ static ggml_backend_dev_t find_htp() {
|
||||||
return nullptr;
|
return nullptr;
|
||||||
}
|
}
|
||||||
static inline float silu(float x) { return x / (1.0f + std::exp(-x)); }
|
static inline float silu(float x) { return x / (1.0f + std::exp(-x)); }
|
||||||
|
|
||||||
// Sampler temp + top_k (équivalent Python subtalker_temp=0.9 top_k=50). Repro stable via seed.
|
|
||||||
static uint32_t rng_state_ = 1u;
|
|
||||||
static void rng_seed(uint32_t s) { rng_state_ = s ? s : 1u; }
|
|
||||||
static float rng_unif() {
|
|
||||||
rng_state_ ^= rng_state_ << 13; rng_state_ ^= rng_state_ >> 17; rng_state_ ^= rng_state_ << 5;
|
|
||||||
return (rng_state_ & 0x00FFFFFF) / (float)0x01000000;
|
|
||||||
}
|
|
||||||
static int sample_top_k_temp(const float* logits, int n_vocab, float temp, int top_k) {
|
|
||||||
if (temp <= 0.0f || top_k == 1) {
|
|
||||||
int b = 0; float mx = logits[0];
|
|
||||||
for (int i = 1; i < n_vocab; ++i) if (logits[i] > mx) { mx = logits[i]; b = i; }
|
|
||||||
return b;
|
|
||||||
}
|
|
||||||
// build top_k indices (heap-like O(n log k))
|
|
||||||
if (top_k <= 0 || top_k > n_vocab) top_k = n_vocab;
|
|
||||||
std::vector<std::pair<float,int>> topk; topk.reserve(top_k);
|
|
||||||
for (int i = 0; i < n_vocab; ++i) {
|
|
||||||
if ((int)topk.size() < top_k) topk.push_back({logits[i], i});
|
|
||||||
else {
|
|
||||||
int worst = 0;
|
|
||||||
for (int j = 1; j < (int)topk.size(); ++j) if (topk[j].first < topk[worst].first) worst = j;
|
|
||||||
if (logits[i] > topk[worst].first) topk[worst] = {logits[i], i};
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// softmax with temperature on the top-k
|
|
||||||
float mx = topk[0].first; for (auto& p : topk) if (p.first > mx) mx = p.first;
|
|
||||||
double Z = 0;
|
|
||||||
std::vector<double> w(topk.size());
|
|
||||||
for (size_t j = 0; j < topk.size(); ++j) { w[j] = std::exp((topk[j].first - mx) / temp); Z += w[j]; }
|
|
||||||
double r = rng_unif() * Z, acc = 0;
|
|
||||||
for (size_t j = 0; j < topk.size(); ++j) { acc += w[j]; if (r <= acc) return topk[j].second; }
|
|
||||||
return topk.back().second;
|
|
||||||
}
|
|
||||||
static void linear(const float* W, const float* b, int out_dim, int in_dim, const float* x, float* out) {
|
static void linear(const float* W, const float* b, int out_dim, int in_dim, const float* x, float* out) {
|
||||||
for (int m = 0; m < out_dim; ++m) {
|
for (int m = 0; m < out_dim; ++m) {
|
||||||
float s = b ? b[m] : 0.0f;
|
float s = b ? b[m] : 0.0f;
|
||||||
|
|
@ -271,11 +238,22 @@ int main(int argc, char** argv) {
|
||||||
if (n_embd != hidden || n_vocab != 3072) { printf("talker dims mismatch (n_embd=%d, n_vocab=%d)\n", n_embd, n_vocab); return 1; }
|
if (n_embd != hidden || n_vocab != 3072) { printf("talker dims mismatch (n_embd=%d, n_vocab=%d)\n", n_embd, n_vocab); return 1; }
|
||||||
printf("talker: n_embd=%d n_vocab=%d threads=%d\n", n_embd, n_vocab, cp.n_threads);
|
printf("talker: n_embd=%d n_vocab=%d threads=%d\n", n_embd, n_vocab, cp.n_threads);
|
||||||
|
|
||||||
// ----- 4) Charger CP
|
// ----- 4) Charger CP + activer sampling HF-style (rep_penalty + top_k + top_p + temp).
|
||||||
|
// Sans sampling propre, le greedy tombe dans des attracteurs (codes 690/207/1571 répétés
|
||||||
|
// = silence/pause) -> WAV avec gros blancs. rep_penalty=1.05 = Python subtalker défaut.
|
||||||
CPState cp_state; cp_state.n_threads = cp.n_threads;
|
CPState cp_state; cp_state.n_threads = cp.n_threads;
|
||||||
|
cp_state.sampler.temp = getenv("KZTTS_CP_TEMP") ? atof(getenv("KZTTS_CP_TEMP")) : 0.9f;
|
||||||
|
cp_state.sampler.top_k = getenv("KZTTS_CP_TOPK") ? atoi(getenv("KZTTS_CP_TOPK")) : 50;
|
||||||
|
cp_state.sampler.top_p = getenv("KZTTS_CP_TOPP") ? atof(getenv("KZTTS_CP_TOPP")) : 1.0f;
|
||||||
|
cp_state.sampler.rep_penalty = getenv("KZTTS_CP_REPP") ? atof(getenv("KZTTS_CP_REPP")) : 1.05f;
|
||||||
|
cp_state.sampler.rep_window = 16;
|
||||||
if (!cp_load(cp_state, (D + "cp_f16.gguf").c_str(), (D + "cp_heads.bin").c_str(), (D + "cp_codec_embs.bin").c_str())) {
|
if (!cp_load(cp_state, (D + "cp_f16.gguf").c_str(), (D + "cp_heads.bin").c_str(), (D + "cp_codec_embs.bin").c_str())) {
|
||||||
printf("CP load FAILED\n"); return 1;
|
printf("CP load FAILED\n"); return 1;
|
||||||
}
|
}
|
||||||
|
const uint32_t SEED = getenv("KZTTS_SEED") ? atoi(getenv("KZTTS_SEED")) : 42;
|
||||||
|
sampler_seed(cp_state.sampler, SEED + 1); // seed différent du Talker pour décorréler
|
||||||
|
printf("CP sampling: temp=%.2f top_k=%d top_p=%.2f rep_penalty=%.2f\n",
|
||||||
|
cp_state.sampler.temp, cp_state.sampler.top_k, cp_state.sampler.top_p, cp_state.sampler.rep_penalty);
|
||||||
|
|
||||||
// ----- 5) Charger decoder
|
// ----- 5) Charger decoder
|
||||||
Decoder dec;
|
Decoder dec;
|
||||||
|
|
@ -309,14 +287,24 @@ int main(int argc, char** argv) {
|
||||||
const double t_pfill = now_s() - t_pfill0;
|
const double t_pfill = now_s() - t_pfill0;
|
||||||
|
|
||||||
// --- LOOP
|
// --- LOOP
|
||||||
// sampling params (équivalent Python defaults : temp=0.9, top_k=50). KZTTS_TEMP/KZTTS_TOPK/KZTTS_SEED env override.
|
// Sampler Talker (CB0) HF-style. Historique sur les CB0 émis (Talker = phrase entière).
|
||||||
const float SAMPLE_TEMP = getenv("KZTTS_TEMP") ? atof(getenv("KZTTS_TEMP")) : 0.9f;
|
Sampler talker_sampler{};
|
||||||
const int SAMPLE_TOPK = getenv("KZTTS_TOPK") ? atoi(getenv("KZTTS_TOPK")) : 50;
|
talker_sampler.temp = getenv("KZTTS_TEMP") ? (float)atof(getenv("KZTTS_TEMP")) : 0.9f;
|
||||||
rng_seed(getenv("KZTTS_SEED") ? atoi(getenv("KZTTS_SEED")) : 42);
|
talker_sampler.top_k = getenv("KZTTS_TOPK") ? atoi(getenv("KZTTS_TOPK")) : 50;
|
||||||
printf("sampling: temp=%.2f top_k=%d seed=%u\n", SAMPLE_TEMP, SAMPLE_TOPK, getenv("KZTTS_SEED") ? atoi(getenv("KZTTS_SEED")) : 42);
|
talker_sampler.top_p = getenv("KZTTS_TOPP") ? (float)atof(getenv("KZTTS_TOPP")) : 1.0f;
|
||||||
|
talker_sampler.rep_penalty = getenv("KZTTS_REPP") ? (float)atof(getenv("KZTTS_REPP")) : 1.05f;
|
||||||
|
talker_sampler.rep_window = 64;
|
||||||
|
sampler_seed(talker_sampler, SEED);
|
||||||
|
printf("Talker sampling: temp=%.2f top_k=%d top_p=%.2f rep_penalty=%.2f seed=%u\n",
|
||||||
|
talker_sampler.temp, talker_sampler.top_k, talker_sampler.top_p, talker_sampler.rep_penalty, SEED);
|
||||||
|
|
||||||
auto logits = llama_get_logits_ith(ctx, -1);
|
// logits du prefill : on les copie dans un buffer mutable (le sampler modifie en place)
|
||||||
int cb0 = sample_top_k_temp(logits, n_vocab, SAMPLE_TEMP, SAMPLE_TOPK);
|
std::vector<float> logits_buf(n_vocab);
|
||||||
|
{
|
||||||
|
const float* lp = llama_get_logits_ith(ctx, -1);
|
||||||
|
memcpy(logits_buf.data(), lp, n_vocab * sizeof(float));
|
||||||
|
}
|
||||||
|
int cb0 = sampler_sample(talker_sampler, logits_buf.data(), n_vocab);
|
||||||
std::vector<float> hidden_for_cp(n_embd);
|
std::vector<float> hidden_for_cp(n_embd);
|
||||||
{ const float* hh = llama_get_embeddings_ith(ctx, -1); if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float)); }
|
{ const float* hh = llama_get_embeddings_ith(ctx, -1); if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float)); }
|
||||||
|
|
||||||
|
|
@ -363,8 +351,11 @@ int main(int argc, char** argv) {
|
||||||
if (llama_decode(ctx, b) != 0) { printf("step %d FAILED\n", s); break; }
|
if (llama_decode(ctx, b) != 0) { printf("step %d FAILED\n", s); break; }
|
||||||
t_decode_total += now_s() - td0;
|
t_decode_total += now_s() - td0;
|
||||||
|
|
||||||
logits = llama_get_logits_ith(ctx, -1);
|
{
|
||||||
cb0 = sample_top_k_temp(logits, n_vocab, SAMPLE_TEMP, SAMPLE_TOPK);
|
const float* lp = llama_get_logits_ith(ctx, -1);
|
||||||
|
memcpy(logits_buf.data(), lp, n_vocab * sizeof(float));
|
||||||
|
}
|
||||||
|
cb0 = sampler_sample(talker_sampler, logits_buf.data(), n_vocab);
|
||||||
const float* hh = llama_get_embeddings_ith(ctx, -1);
|
const float* hh = llama_get_embeddings_ith(ctx, -1);
|
||||||
if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float));
|
if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float));
|
||||||
N_done = s + 1;
|
N_done = s + 1;
|
||||||
|
|
@ -382,12 +373,12 @@ int main(int argc, char** argv) {
|
||||||
std::ofstream f("/data/local/tmp/kz-engine/pipeline_codes.bin", std::ios::binary);
|
std::ofstream f("/data/local/tmp/kz-engine/pipeline_codes.bin", std::ios::binary);
|
||||||
f.write((const char*)codes_engine.data(), codes_engine.size() * sizeof(int32_t));
|
f.write((const char*)codes_engine.data(), codes_engine.size() * sizeof(int32_t));
|
||||||
}
|
}
|
||||||
printf("codes preview:\n");
|
printf("codes (CB0 trajectory only, %d frames):\n ", N);
|
||||||
for (int t = 0; t < std::min(5, N); ++t) {
|
for (int t = 0; t < N; ++t) {
|
||||||
printf(" frame %d: ", t);
|
printf("%d ", codes_engine[t * 16 + 0]);
|
||||||
for (int c = 0; c < 16; ++c) printf("%d ", codes_engine[t * 16 + c]);
|
if ((t + 1) % 16 == 0) printf("\n ");
|
||||||
printf("\n");
|
|
||||||
}
|
}
|
||||||
|
printf("\n");
|
||||||
// ----- 7) Decoder ggml : codes [N, 16] (time-major) -> WAV
|
// ----- 7) Decoder ggml : codes [N, 16] (time-major) -> WAV
|
||||||
// Decoder veut codes_flat[16 * T] codebook-major (CB-fastest dans son forward).
|
// Decoder veut codes_flat[16 * T] codebook-major (CB-fastest dans son forward).
|
||||||
// codes_engine = [N, 16] time-major -> transpose en [16, N] codebook-major.
|
// codes_engine = [N, 16] time-major -> transpose en [16, N] codebook-major.
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue