chantier B TTS #11: libkazeia_tts.so + Kotlin wrapper, JNI validé bout-en-bout

Refactor: tts_engine.{h,cpp} = API publique de l'engine TTS (load une fois,
synthesize N fois, free). État load-once dans struct TtsEngine (talker
model+ctx, CPState, Decoder, KzTextTokenizer, fixtures + embeds spéciaux
+ role_proj pré-calculé). KV cache talker reset via llama_memory_clear
au début de chaque synthesize -> appels indépendants.

tts_pipeline.cpp devient un thin CLI dessus (refonte sans changement de
sortie : WAV md5 identique au pré-refactor sur 'Bonjour je m'appelle Kazeia').

JNI bridge: kazeia_tts_jni.cpp expose 3 fonctions :
  Java_com_kazeia_tts_TtsJni_nativeLoad / Synthesize / Free
Signatures alignées avec TtsEngine.kt (companion loadLibrary 'kazeia_tts').
Build: build_kazeia_tts.sh -> b-jni/libkazeia_tts.so (~900 KB).

Test_engine_2calls (sans JNI) : 3 synth sur même instance, KV reset OK,
2 appels identiques -> WAV md5 identique, appel 3 différent -> codes
différents.

Test_jni_tts (harness JNI sans VM Java) : dlopen libkazeia_tts.so, JNIEnv
mock minimal (GetStringUTFChars/Release + NewIntArray/SetIntArrayRegion),
appels load + 2 synth + free. WAV md5 identiques aux runs directs. exit=0.

Empreinte tablette mesurée: ~3.5 GB par instance (talker f32 1.6 GB + CP
f16 mixte 250 MB + decoder 325 MB + vocab Qwen3 vocab_only 50 MB + fixtures
text_embed/tp_* 1.2 GB). Une instance par process.

RTF stable autour de 3.0 (CPU 6t), inchangé vs avant refonte.

Reste sur la liste initiale : tuning sampling (itératif à l'oreille).

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Richard Loyer 2026-05-28 21:54:21 +02:00
parent b27c8ef403
commit 930f3c8c1f
10 changed files with 1051 additions and 419 deletions

43
dist/build_kazeia_tts.sh vendored Executable file
View File

@ -0,0 +1,43 @@
#!/usr/bin/env bash
# Build libkazeia_tts.so (SHARED, JNI) + test_jni_tts (harness standalone).
set -euo pipefail
HERE="$(cd "$(dirname "$0")" && pwd)"
NDK_ROOT="${ANDROID_NDK_ROOT:-/opt/Kazeia/android-ndk-r27d}"
CXX="$NDK_ROOT/toolchains/llvm/prebuilt/linux-x86_64/bin/aarch64-linux-android31-clang++"
JNI="$HERE/jni"
INC="$HERE/include"
LIB="$HERE/lib"
DEC="/opt/Kazeia/kazeia-tts-decoder-ggml"
mkdir -p "$HERE/b-jni"
CFLAGS=( -std=c++17 -O3 -fPIC
-march=armv8.6-a+i8mm+bf16+dotprod+fp16
-Wno-unused-parameter -Wno-unused-variable -Wno-sign-compare
-I"$INC" -I"$DEC/src"
)
SRCS=(
"$JNI/kazeia_tts_jni.cpp"
"$JNI/tts_engine.cpp"
"$JNI/cp_inference.cpp"
"$JNI/sampler.cpp"
"$JNI/kazeia_text_tokenizer.cpp"
)
LIBS=(
"$DEC/build-android-engine/libqwen3tts-decoder.a"
-L"$LIB" -lllama -lggml -lggml-base -lggml-cpu
-Wl,-rpath,'$ORIGIN'
-llog -ldl -lm
)
echo "== building libkazeia_tts.so =="
"$CXX" "${CFLAGS[@]}" -shared "${SRCS[@]}" "${LIBS[@]}" -o "$HERE/b-jni/libkazeia_tts.so"
ls -la "$HERE/b-jni/libkazeia_tts.so"
# Harness JNI sans VM Java : dlopen + JNIEnv mock minimal, vérifie load/synthesize/free.
echo "== building test_jni_tts (harness sans Java VM) =="
"$CXX" "${CFLAGS[@]}" \
"$JNI/test_jni_tts.cpp" \
-Wl,-rpath,'$ORIGIN' \
-ldl -lm \
-o "$HERE/b-jni/test_jni_tts"
ls -la "$HERE/b-jni/test_jni_tts"

21
dist/build_test_engine_2calls.sh vendored Executable file
View File

@ -0,0 +1,21 @@
#!/usr/bin/env bash
set -euo pipefail
HERE="$(cd "$(dirname "$0")" && pwd)"
NDK_ROOT="${ANDROID_NDK_ROOT:-/opt/Kazeia/android-ndk-r27d}"
CXX="$NDK_ROOT/toolchains/llvm/prebuilt/linux-x86_64/bin/aarch64-linux-android31-clang++"
JNI="$HERE/jni"
INC="$HERE/include"
LIB="$HERE/lib"
DEC="/opt/Kazeia/kazeia-tts-decoder-ggml"
mkdir -p "$HERE/b-jni"
"$CXX" -std=c++17 -O3 -fPIE \
-march=armv8.6-a+i8mm+bf16+dotprod+fp16 \
-I"$INC" -I"$DEC/src" \
"$JNI/test_engine_2calls.cpp" "$JNI/tts_engine.cpp" "$JNI/cp_inference.cpp" \
"$JNI/sampler.cpp" "$JNI/kazeia_text_tokenizer.cpp" \
"$DEC/build-android-engine/libqwen3tts-decoder.a" \
-L"$LIB" -lllama -lggml -lggml-base -lggml-cpu \
-Wl,-rpath,'$ORIGIN' \
-llog -ldl -lm \
-o "$HERE/b-jni/test_engine_2calls"
ls -la "$HERE/b-jni/test_engine_2calls"

View File

@ -32,6 +32,7 @@ CFLAGS=(
INCS=( -I"$INC" -I"$DECODER_INC" ) INCS=( -I"$INC" -I"$DECODER_INC" )
LIBS=( LIBS=(
"$JNI/tts_pipeline.cpp" "$JNI/tts_pipeline.cpp"
"$JNI/tts_engine.cpp"
"$JNI/cp_inference.cpp" "$JNI/cp_inference.cpp"
"$JNI/sampler.cpp" "$JNI/sampler.cpp"
"$JNI/kazeia_text_tokenizer.cpp" "$JNI/kazeia_text_tokenizer.cpp"

89
dist/jni/TtsEngine.kt vendored Normal file
View File

@ -0,0 +1,89 @@
package com.kazeia.tts
// Engine TTS Qwen3 in-process (libkazeia_tts.so). Load une fois, synthesize N fois.
// Le bridge natif est dans dist/jni/kazeia_tts_jni.cpp ; build via dist/build_kazeia_tts.sh
// ou dist/CMakeLists.txt.
//
// Empreinte RAM typique : ~3.5 GB (talker f32 ~1.6 GB + CP f16 mixte ~250 MB +
// decoder ~325 MB + vocab Qwen3 ~50 MB + fixtures text_embed/tp_* ~1.2 GB).
class TtsJni {
external fun nativeLoad(
talkerGguf: String,
vocabGguf: String,
dumpDir: String,
useHtp: Boolean,
nThreads: Int,
cpUseCache: Boolean
): Long
external fun nativeSynthesize(
handle: Long,
text: String,
outWavPath: String,
maxSteps: Int,
seed: Int,
cpTemp: Float, cpTopK: Int, cpTopP: Float, cpRepPenalty: Float,
talkerTemp: Float, talkerTopK: Int, talkerTopP: Float, talkerRepPenalty: Float
): IntArray // [err, frames, total_ms, prefill_ms, talker_ms, cp_ms, decoder_ms]
external fun nativeFree(handle: Long)
companion object { init { System.loadLibrary("kazeia_tts") } }
}
data class TtsResult(
val frames: Int,
val audioMs: Int,
val totalMs: Int,
val prefillMs: Int,
val talkerMs: Int,
val cpMs: Int,
val decoderMs: Int
)
// Configuration sampling. Défauts validés sur fixture FR (cf project_tts_session_28may).
data class TtsSampling(
val cpTemp: Float = 0.9f, val cpTopK: Int = 50, val cpTopP: Float = 1.0f, val cpRepPenalty: Float = 1.05f,
val talkerTemp: Float = 0.9f, val talkerTopK: Int = 50, val talkerTopP: Float = 1.0f, val talkerRepPenalty: Float = 1.05f
)
class TtsEngine(
talkerGguf: String,
vocabGguf: String,
dumpDir: String,
useHtp: Boolean = false,
nThreads: Int = 6,
cpUseCache: Boolean = true
) {
private val jni = TtsJni()
private val h = jni.nativeLoad(talkerGguf, vocabGguf, dumpDir, useHtp, nThreads, cpUseCache)
init { require(h != 0L) { "TtsEngine: load FAILED (talker=$talkerGguf, vocab=$vocabGguf, dump=$dumpDir)" } }
@Throws(IllegalStateException::class)
fun synthesize(
text: String,
outWavPath: String,
maxSteps: Int = 256,
seed: Int = 42,
sampling: TtsSampling = TtsSampling()
): TtsResult {
val r = jni.nativeSynthesize(
h, text, outWavPath, maxSteps, seed,
sampling.cpTemp, sampling.cpTopK, sampling.cpTopP, sampling.cpRepPenalty,
sampling.talkerTemp, sampling.talkerTopK, sampling.talkerTopP, sampling.talkerRepPenalty
)
if (r[0] != 0) error("TtsEngine.synthesize err=${r[0]} for text=\"$text\"")
val frames = r[1]
return TtsResult(
frames = frames,
audioMs = frames * 1000 / 12, // 12 frames/s = 80 ms par frame
totalMs = r[2],
prefillMs = r[3],
talkerMs = r[4],
cpMs = r[5],
decoderMs = r[6]
)
}
fun release() { jni.nativeFree(h) }
}

105
dist/jni/kazeia_tts_jni.cpp vendored Normal file
View File

@ -0,0 +1,105 @@
// JNI bridge pour libkazeia_tts.so. Thin wrapper autour de tts_engine.{h,cpp}.
// Le handle exposé côté Kotlin est un jlong = (jlong)(TtsEngine*).
//
// Signatures (à matcher dans le .kt) :
// nativeLoad(talkerGguf, vocabGguf, dumpDir, useHtp, nThreads, cpUseCache) : Long
// nativeSynthesize(handle, text, outWavPath, maxSteps, seed,
// cpTemp, cpTopK, cpTopP, cpRepPenalty,
// talkerTemp, talkerTopK, talkerTopP, talkerRepPenalty)
// : IntArray { err, frames, total_ms, prefill_ms, talker_ms, cp_ms, decoder_ms }
// nativeFree(handle)
#include "tts_engine.h"
#include <jni.h>
#include <string>
namespace {
// Helper : récupère une C-string UTF-8 depuis un jstring (rendre à java avec release).
struct JStr {
JNIEnv * env;
jstring j;
const char * c;
JStr(JNIEnv* e, jstring s) : env(e), j(s), c(s ? e->GetStringUTFChars(s, nullptr) : nullptr) {}
~JStr() { if (j && c) env->ReleaseStringUTFChars(j, c); }
};
} // namespace
extern "C" {
JNIEXPORT jlong JNICALL
Java_com_kazeia_tts_TtsJni_nativeLoad(JNIEnv* env, jobject /*thiz*/,
jstring talker_gguf, jstring vocab_gguf,
jstring dump_dir,
jboolean use_htp, jint n_threads,
jboolean cp_use_cache) {
JStr talker(env, talker_gguf);
JStr vocab (env, vocab_gguf);
JStr dump (env, dump_dir);
TtsEngineLoadCfg lc;
lc.talker_gguf = talker.c;
lc.vocab_gguf = vocab.c;
lc.dump_dir = dump.c;
lc.use_htp = (use_htp == JNI_TRUE);
lc.n_threads = n_threads;
lc.cp_use_cache = (cp_use_cache == JNI_TRUE);
auto * eng = tts_engine_load(lc);
return (jlong)(uintptr_t)eng;
}
JNIEXPORT jintArray JNICALL
Java_com_kazeia_tts_TtsJni_nativeSynthesize(JNIEnv* env, jobject /*thiz*/,
jlong handle,
jstring text, jstring out_wav_path,
jint max_steps, jint seed,
jfloat cp_temp, jint cp_top_k, jfloat cp_top_p, jfloat cp_rep_penalty,
jfloat talker_temp, jint talker_top_k, jfloat talker_top_p, jfloat talker_rep_penalty) {
auto * eng = (TtsEngine*)(uintptr_t)handle;
JStr text_s(env, text);
JStr wav_s (env, out_wav_path);
TtsSynthesizeCfg sc;
sc.text = text_s.c;
sc.out_wav_path = wav_s.c;
sc.max_steps = max_steps;
sc.seed = (uint32_t)seed;
sc.cp_temp = cp_temp;
sc.cp_top_k = cp_top_k;
sc.cp_top_p = cp_top_p;
sc.cp_rep_penalty = cp_rep_penalty;
sc.talker_temp = talker_temp;
sc.talker_top_k = talker_top_k;
sc.talker_top_p = talker_top_p;
sc.talker_rep_penalty = talker_rep_penalty;
TtsSynthesizeResult R{ -100, 0, 0, 0, 0, 0, 0, 0 };
if (eng) R = tts_engine_synthesize(eng, sc);
// Pack résultats dans IntArray :
// [0] err
// [1] frames
// [2] total_ms (audio_s codé séparément si besoin = frames/12)
// [3] prefill_ms
// [4] talker_loop_ms
// [5] cp_loop_ms
// [6] decoder_ms
jint out[7];
out[0] = R.err;
out[1] = R.frames;
out[2] = (jint)(R.total_s * 1000.0);
out[3] = (jint)(R.prefill_s * 1000.0);
out[4] = (jint)(R.talker_loop_s * 1000.0);
out[5] = (jint)(R.cp_loop_s * 1000.0);
out[6] = (jint)(R.decoder_s * 1000.0);
jintArray arr = env->NewIntArray(7);
env->SetIntArrayRegion(arr, 0, 7, out);
return arr;
}
JNIEXPORT void JNICALL
Java_com_kazeia_tts_TtsJni_nativeFree(JNIEnv* /*env*/, jobject /*thiz*/, jlong handle) {
auto * eng = (TtsEngine*)(uintptr_t)handle;
tts_engine_free(eng);
}
} // extern "C"

49
dist/jni/test_engine_2calls.cpp vendored Normal file
View File

@ -0,0 +1,49 @@
// Vérifie qu'on peut load une fois puis synthesize 2 fois sur la même instance,
// avec KV cache talker resetée entre les deux. Critère :
// - 2e appel sur la même phrase doit produire le même WAV md5 que le 1er
// (= la KV reset a bien tout effacé).
// - 2e appel sur une phrase différente doit produire des codes différents
// (= la nouvelle phrase est bien tokenisée et synthétisée fresh).
#include "tts_engine.h"
#include <cstdio>
#include <cstdlib>
#include <cstring>
int main(int argc, char** argv) {
if (argc < 5) {
printf("usage: %s <talker_gguf> <vocab_gguf> <dump_dir> <out_dir>\n", argv[0]);
return 1;
}
TtsEngineLoadCfg lc;
lc.talker_gguf = argv[1];
lc.vocab_gguf = argv[2];
lc.dump_dir = argv[3];
lc.use_htp = false;
lc.n_threads = 6;
lc.cp_use_cache = true;
std::string OUT = argv[4];
auto * eng = tts_engine_load(lc);
if (!eng) return 2;
auto run = [&](const char* text, const char* fname) {
TtsSynthesizeCfg sc;
sc.text = text;
std::string path = OUT + "/" + fname;
sc.out_wav_path = path.c_str();
sc.max_steps = 256;
sc.seed = 42;
auto R = tts_engine_synthesize(eng, sc);
printf("[%s] err=%d frames=%d audio=%.2fs total=%.2fs RTF=%.2f\n",
fname, R.err, R.frames, R.audio_s, R.total_s,
R.audio_s > 0 ? R.total_s / R.audio_s : 0);
return R.err;
};
int e1 = run("Bonjour je m'appelle Kazeia", "out_call1.wav");
int e2 = run("Bonjour je m'appelle Kazeia", "out_call2.wav");
int e3 = run("Bonsoir, comment tu te sens ce soir ?", "out_call3.wav");
tts_engine_free(eng);
return (e1 == 0 && e2 == 0 && e3 == 0) ? 0 : 3;
}

120
dist/jni/test_jni_tts.cpp vendored Normal file
View File

@ -0,0 +1,120 @@
// Harness JNI sans VM Java : on dlopen libkazeia_tts.so, on récupère les Java_*
// symbols, et on les appelle avec un JNIEnv mock minimal (juste les fonctions que
// kazeia_tts_jni.cpp utilise : GetStringUTFChars, ReleaseStringUTFChars, NewIntArray,
// SetIntArrayRegion). jstring est juste un alias opaque pour const char*.
//
// Permet de valider le bridge JNI sans installer l'app — utile pour CI / régression.
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <unistd.h> // _exit
#include <dlfcn.h>
#include <jni.h>
#include <string>
namespace {
// Pour notre mock, jstring = const char* tel quel (alias).
const char * mock_GetStringUTFChars(JNIEnv*, jstring s, jboolean*) { return (const char*)s; }
void mock_ReleaseStringUTFChars(JNIEnv*, jstring, const char*) {}
// jintArray = pointeur vers une struct {int n; jint data[n]} qu'on alloue.
struct IntArr { int n; jint data[1]; };
jintArray mock_NewIntArray(JNIEnv*, jsize len) {
auto* a = (IntArr*)malloc(sizeof(int) + (size_t)len * sizeof(jint));
a->n = len;
return (jintArray)a;
}
void mock_SetIntArrayRegion(JNIEnv*, jintArray arr, jsize start, jsize len, const jint* buf) {
auto* a = (IntArr*)arr;
memcpy(a->data + start, buf, (size_t)len * sizeof(jint));
}
struct MockNativeInterface {
// Les offsets et le layout doivent matcher jni.h JNINativeInterface. Plutôt que
// de reproduire toute la struct, on alloue une zone large, et on pose les
// function pointers AUX BONS OFFSETS (cf jni.h: GetStringUTFChars = entry 169,
// ReleaseStringUTFChars = 170, NewIntArray = 187, SetIntArrayRegion = 207).
// Pour rester portable, on utilise directement l'API via le struct JNINativeInterface
// déclaré dans jni.h.
JNINativeInterface iface{};
JNINativeInterface * ptr;
MockNativeInterface() {
// On laisse tout à nullptr, on remplit juste les méthodes utilisées.
iface.GetStringUTFChars = mock_GetStringUTFChars;
iface.ReleaseStringUTFChars = mock_ReleaseStringUTFChars;
iface.NewIntArray = mock_NewIntArray;
iface.SetIntArrayRegion = mock_SetIntArrayRegion;
ptr = &iface;
}
};
} // namespace
int main(int argc, char** argv) {
if (argc < 6) {
fprintf(stderr,"usage: %s <libkazeia_tts.so> <talker_gguf> <vocab_gguf> <dump_dir> <out.wav>\n", argv[0]);
return 1;
}
const char* lib_path = argv[1];
const char* talker_gguf = argv[2];
const char* vocab_gguf = argv[3];
const char* dump_dir = argv[4];
const char* out_wav = argv[5];
void* dl = dlopen(lib_path, RTLD_NOW | RTLD_LOCAL);
if (!dl) { fprintf(stderr, "dlopen FAIL %s : %s\n", lib_path, dlerror()); return 2; }
using FnLoad = jlong (*)(JNIEnv*, jobject, jstring, jstring, jstring, jboolean, jint, jboolean);
using FnSynth = jintArray (*)(JNIEnv*, jobject, jlong, jstring, jstring, jint, jint,
jfloat, jint, jfloat, jfloat, jfloat, jint, jfloat, jfloat);
using FnFree = void (*)(JNIEnv*, jobject, jlong);
auto p_load = (FnLoad)dlsym(dl, "Java_com_kazeia_tts_TtsJni_nativeLoad");
auto p_synth = (FnSynth)dlsym(dl, "Java_com_kazeia_tts_TtsJni_nativeSynthesize");
auto p_free = (FnFree)dlsym(dl, "Java_com_kazeia_tts_TtsJni_nativeFree");
if (!p_load || !p_synth || !p_free) {
fprintf(stderr, "dlsym FAIL : load=%p synth=%p free=%p\n", (void*)p_load, (void*)p_synth, (void*)p_free);
return 3;
}
MockNativeInterface ifc;
JNIEnv env_buf;
env_buf.functions = ifc.ptr;
JNIEnv* env = &env_buf;
jobject self = nullptr; // pas utilisé dans nos handlers
jlong h = p_load(env, self,
(jstring)talker_gguf, (jstring)vocab_gguf, (jstring)dump_dir,
/*useHtp=*/JNI_FALSE, /*nThreads=*/6, /*cpUseCache=*/JNI_TRUE);
if (h == 0) { fprintf(stderr, "nativeLoad returned 0\n"); dlclose(dl); return 4; }
fprintf(stderr,"nativeLoad OK -> handle=%p\n", (void*)(uintptr_t)h);
// 2 appels : même phrase puis phrase différente.
auto run = [&](const char* text, const char* out_path) {
auto* arr = p_synth(env, self, h,
(jstring)text, (jstring)out_path,
/*maxSteps=*/256, /*seed=*/42,
/*cpTemp=*/0.9f, /*cpTopK=*/50, /*cpTopP=*/1.0f, /*cpRepPenalty=*/1.05f,
/*talkerTemp=*/0.9f, /*talkerTopK=*/50, /*talkerTopP=*/1.0f, /*talkerRepPenalty=*/1.05f);
IntArr* a = (IntArr*)arr;
fprintf(stderr,"[%s] err=%d frames=%d total=%dms (pf=%d talker=%d cp=%d dec=%d) -> %s\n",
text, a->data[0], a->data[1], a->data[2], a->data[3], a->data[4], a->data[5], a->data[6], out_path);
int err = a->data[0];
free(a);
return err;
};
std::string out1 = out_wav; out1 += ".1.wav";
std::string out2 = out_wav; out2 += ".2.wav";
int e1 = run("Bonjour je m'appelle Kazeia", out1.c_str());
int e2 = run("Bonsoir, comment tu te sens ce soir ?", out2.c_str());
p_free(env, self, h);
fprintf(stderr,"OK e1=%d e2=%d\n", e1, e2);
fflush(stderr); fflush(stdout);
// Pas de dlclose ni return : destructors statiques de libllama/ggml peuvent segfaulter
// après unmap (atexit/global dtors). En contexte app Android, la .so reste chargée
// pour toute la vie du process, donc ce path n'existe pas. _exit court-circuite.
_exit((e1 == 0 && e2 == 0) ? 0 : 5);
}

492
dist/jni/tts_engine.cpp vendored Normal file
View File

@ -0,0 +1,492 @@
#include "tts_engine.h"
#include "llama.h"
#include "ggml-backend.h"
#include "cp_inference.h"
#include "sampler.h"
#include "kazeia_text_tokenizer.h"
#include "decoder.h"
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <cmath>
#include <chrono>
#include <fstream>
#include <string>
#include <vector>
using namespace kazeia::tts;
namespace {
double now_s() {
using clk = std::chrono::steady_clock;
return std::chrono::duration<double>(clk::now().time_since_epoch()).count();
}
std::vector<float> read_f32(const std::string& p, size_t n_expected) {
std::ifstream f(p, std::ios::binary | std::ios::ate);
if (!f) { fprintf(stderr, "tts_engine: open %s\n", p.c_str()); return {}; }
size_t n = (size_t)f.tellg() / sizeof(float);
if (n_expected && n != n_expected) {
fprintf(stderr, "tts_engine: %s: %zu f32, attendu %zu\n", p.c_str(), n, n_expected); return {};
}
f.seekg(0); std::vector<float> v(n); f.read((char*)v.data(), n * sizeof(float)); return v;
}
ggml_backend_dev_t find_htp() {
for (size_t i = 0; i < ggml_backend_dev_count(); ++i) {
auto d = ggml_backend_dev_get(i);
if (!strcmp(ggml_backend_dev_name(d), "HTP0")) return d;
}
return nullptr;
}
inline float silu(float x) { return x / (1.0f + std::exp(-x)); }
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) {
float s = b ? b[m] : 0.0f;
const float* wr = W + (size_t)m * in_dim;
for (int k = 0; k < in_dim; ++k) s += wr[k] * x[k];
out[m] = s;
}
}
void text_projection(const float* in, int N,
const float* fc1_w, const float* fc1_b,
const float* fc2_w, const float* fc2_b,
float* mid_buf, float* out,
int text_hidden, int hidden) {
for (int n = 0; n < N; ++n) {
linear(fc1_w, fc1_b, text_hidden, text_hidden, in + n * text_hidden, mid_buf);
for (int d = 0; d < text_hidden; ++d) mid_buf[d] = silu(mid_buf[d]);
linear(fc2_w, fc2_b, hidden, text_hidden, mid_buf, out + n * hidden);
}
}
bool write_wav_pcm16_mono(const char* path, const float* samples, size_t N, int sr = 24000) {
FILE* f = fopen(path, "wb");
if (!f) return false;
const uint32_t data_bytes = (uint32_t)(N * 2);
const uint32_t riff_size = 36 + data_bytes;
auto w16 = [&](uint16_t v){ fwrite(&v, 2, 1, f); };
auto w32 = [&](uint32_t v){ fwrite(&v, 4, 1, f); };
fwrite("RIFF", 1, 4, f); w32(riff_size); fwrite("WAVE", 1, 4, f);
fwrite("fmt ", 1, 4, f); w32(16); w16(1); w16(1);
w32((uint32_t)sr); w32((uint32_t)sr * 2); w16(2); w16(16);
fwrite("data", 1, 4, f); w32(data_bytes);
for (size_t i = 0; i < N; ++i) {
float v = samples[i]; if (v > 1.f) v = 1.f; else if (v < -1.f) v = -1.f;
int16_t pcm = (int16_t)std::lround(v * 32767.0f);
fwrite(&pcm, 2, 1, f);
}
fclose(f); return true;
}
} // namespace
struct TtsEngine {
// --- Constantes manifest ---
int text_vocab = 151936, text_hidden = 2048, hidden = 1024;
int tts_bos = 151672, tts_eos = 151673, tts_pad = 151671;
int codec_bos = 2149, codec_eos = 2150, codec_pad = 2148;
int codec_think = 2154, codec_think_bos = 2156, codec_think_eos = 2157;
int lang_fr = 2061;
// --- Fixtures (load-once) ---
std::vector<float> text_embed;
std::vector<float> tp_fc1_w, tp_fc1_b, tp_fc2_w, tp_fc2_b;
std::vector<float> tok_embd;
std::vector<float> xvector;
// Embeds spéciaux pré-projetés (constants une fois le modèle chargé)
std::vector<float> spec_proj; // [3 * hidden] = {tts_bos_emb, tts_eos_emb, tts_pad_emb}
std::vector<float> codec_input_emb; // [7 * hidden] = {think, think_bos, lang_fr, think_eos, xvec, codec_pad, codec_bos}
std::vector<float> role_proj; // [3 * hidden] = projection de <|im_start|>, assistant, \n (constant car template fixe)
std::vector<float> codec_pad_emb; // [hidden] = tok_embd[codec_pad]
// --- Talker ---
llama_model * talker_m = nullptr;
llama_context * talker_ctx = nullptr;
int n_embd = 0;
int n_vocab = 0;
int npe = 1; // 1 (RoPE) ou 4 (M-RoPE/I-MRoPE)
int n_threads = 6;
// --- CP / Decoder / Tokenizer ---
CPState cp_state;
Decoder decoder;
KzTextTokenizer kz_tok;
bool has_kz_tok = false;
bool cp_use_cache = true;
};
// ===========================================================================
// LOAD
// ===========================================================================
TtsEngine * tts_engine_load(const TtsEngineLoadCfg & cfg) {
if (!cfg.talker_gguf || !cfg.dump_dir) {
fprintf(stderr, "tts_engine_load: talker_gguf et dump_dir requis\n"); return nullptr;
}
auto eng = new TtsEngine();
eng->n_threads = cfg.n_threads;
eng->cp_use_cache = cfg.cp_use_cache;
std::string D = cfg.dump_dir;
if (D.back() != '/') D += '/';
// --- 0) Constantes manifest ---
{
std::ifstream f(D + "manifest_text.txt");
if (!f) { fprintf(stderr, "tts_engine_load: no manifest_text\n"); delete eng; return nullptr; }
std::string line;
auto eq = [&](const char* k){ return line.rfind(k, 0) == 0; };
auto val = [&](size_t off){ return atoi(line.c_str() + off); };
while (std::getline(f, line)) {
if (eq("text_vocab_size:")) eng->text_vocab = val(16);
else if (eq("text_hidden_size:")) eng->text_hidden = val(17);
else if (eq("hidden_size:")) eng->hidden = val(12);
else if (eq("tts_bos_token_id:")) eng->tts_bos = val(17);
else if (eq("tts_eos_token_id:")) eng->tts_eos = val(17);
else if (eq("tts_pad_token_id:")) eng->tts_pad = val(17);
else if (eq("codec_bos_id:")) eng->codec_bos = val(13);
else if (eq("codec_eos_id:")) eng->codec_eos = val(13);
else if (eq("codec_pad_id:")) eng->codec_pad = val(13);
else if (eq("codec_think_id:")) eng->codec_think = val(15);
else if (eq("codec_think_bos_id:")) eng->codec_think_bos = val(19);
else if (eq("codec_think_eos_id:")) eng->codec_think_eos = val(19);
else if (eq("codec_language_french:")) eng->lang_fr = val(22);
}
}
// --- 1) Fixtures ---
const int text_hidden = eng->text_hidden;
const int hidden = eng->hidden;
const int text_vocab = eng->text_vocab;
const double t_load0 = now_s();
eng->text_embed = read_f32(D + "text_embed.bin", (size_t)text_vocab * text_hidden);
eng->tp_fc1_w = read_f32(D + "tp_fc1_w.bin", (size_t)text_hidden * text_hidden);
eng->tp_fc1_b = read_f32(D + "tp_fc1_b.bin", (size_t)text_hidden);
eng->tp_fc2_w = read_f32(D + "tp_fc2_w.bin", (size_t)hidden * text_hidden);
eng->tp_fc2_b = read_f32(D + "tp_fc2_b.bin", (size_t)hidden);
eng->tok_embd = read_f32(D + "talker_tok_embd.bin", (size_t)3072 * hidden);
eng->xvector = read_f32(D + "damien_xvector.bin", (size_t)hidden);
if (eng->text_embed.empty() || eng->tp_fc1_w.empty() || eng->tok_embd.empty() || eng->xvector.empty()) {
fprintf(stderr, "tts_engine_load: fixtures FAIL\n"); delete eng; return nullptr;
}
// --- 2) Backend talker (avant le 1er init backend pour HMX off, etc.) ---
setenv("GGML_HEXAGON_USE_HMX", "0", 1);
llama_backend_init();
// --- 3) Tokenizer optionnel (vocab_only -> ~50 MB RAM) ---
if (cfg.vocab_gguf && *cfg.vocab_gguf) {
if (!kz_tok_load(eng->kz_tok, cfg.vocab_gguf)) {
fprintf(stderr, "tts_engine_load: vocab load FAIL %s\n", cfg.vocab_gguf);
delete eng; return nullptr;
}
eng->has_kz_tok = true;
}
// --- 4) Talker (libllama) ---
auto mp = llama_model_default_params();
ggml_backend_dev_t devs[2] = { cfg.use_htp ? find_htp() : nullptr, nullptr };
if (devs[0]) { mp.n_gpu_layers = 99; mp.devices = devs; fprintf(stderr, "talker: HTP0\n"); }
else { mp.n_gpu_layers = 0; fprintf(stderr, "talker: CPU\n"); }
eng->talker_m = llama_model_load_from_file(cfg.talker_gguf, mp);
if (!eng->talker_m) { fprintf(stderr, "talker load FAIL %s\n", cfg.talker_gguf); delete eng; return nullptr; }
auto cp = llama_context_default_params();
// Borne ctx généreuse : prefill max ~50 + max_steps 512 + marge. Pas critique en CPU.
cp.n_ctx = 1024;
cp.n_batch = 1024;
cp.n_threads = cfg.n_threads;
cp.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED;
cp.embeddings = true;
eng->talker_ctx = llama_init_from_model(eng->talker_m, cp);
if (!eng->talker_ctx) { fprintf(stderr, "talker ctx FAIL\n"); delete eng; return nullptr; }
eng->n_embd = llama_model_n_embd(eng->talker_m);
eng->n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(eng->talker_m));
if (eng->n_embd != hidden || eng->n_vocab != 3072) {
fprintf(stderr, "talker dims mismatch (n_embd=%d, n_vocab=%d, attendu %d / 3072)\n",
eng->n_embd, eng->n_vocab, hidden); delete eng; return nullptr;
}
auto rt = llama_model_rope_type(eng->talker_m);
eng->npe = (rt == LLAMA_ROPE_TYPE_MROPE || rt == LLAMA_ROPE_TYPE_IMROPE) ? 4 : 1;
// --- 5) CP ---
eng->cp_state.n_threads = cfg.n_threads;
if (!cp_load(eng->cp_state,
(D + "cp_f16.gguf").c_str(),
(D + "cp_heads.bin").c_str(),
(D + "cp_codec_embs.bin").c_str())) {
fprintf(stderr, "CP load FAIL\n"); delete eng; return nullptr;
}
// --- 6) Decoder ---
if (!eng->decoder.load((D + "qwen3tts_decoder.gguf").c_str())) {
fprintf(stderr, "decoder load FAIL\n"); delete eng; return nullptr;
}
// --- 7) Pre-compute des embeds constants ---
std::vector<float> mid(text_hidden);
// special tokens projetés (tts_bos/eos/pad)
std::vector<float> spec_text(3 * text_hidden);
int sp[3] = { eng->tts_bos, eng->tts_eos, eng->tts_pad };
for (int i = 0; i < 3; ++i)
std::memcpy(spec_text.data() + i * text_hidden,
eng->text_embed.data() + (size_t)sp[i] * text_hidden,
text_hidden * sizeof(float));
eng->spec_proj.resize(3 * hidden);
text_projection(spec_text.data(), 3,
eng->tp_fc1_w.data(), eng->tp_fc1_b.data(),
eng->tp_fc2_w.data(), eng->tp_fc2_b.data(),
mid.data(), eng->spec_proj.data(), text_hidden, hidden);
// codec prefix [think, think_bos, lang_fr, think_eos, x_vec, codec_pad, codec_bos]
int codec_prefill[4] = { eng->codec_think, eng->codec_think_bos, eng->lang_fr, eng->codec_think_eos };
eng->codec_input_emb.resize(7 * hidden);
for (int i = 0; i < 4; ++i)
std::memcpy(eng->codec_input_emb.data() + i * hidden,
eng->tok_embd.data() + (size_t)codec_prefill[i] * hidden,
hidden * sizeof(float));
std::memcpy(eng->codec_input_emb.data() + 4 * hidden, eng->xvector.data(), hidden * sizeof(float));
std::memcpy(eng->codec_input_emb.data() + 5 * hidden,
eng->tok_embd.data() + (size_t)eng->codec_pad * hidden, hidden * sizeof(float));
std::memcpy(eng->codec_input_emb.data() + 6 * hidden,
eng->tok_embd.data() + (size_t)eng->codec_bos * hidden, hidden * sizeof(float));
// role tokens (<|im_start|>, assistant, \n) — projection constante car template fixe.
// Pré-calculé une fois, indépendant du texte d'entrée.
const int role_ids[3] = { 151644, 77091, 198 };
std::vector<float> role_text(3 * text_hidden);
for (int i = 0; i < 3; ++i)
std::memcpy(role_text.data() + i * text_hidden,
eng->text_embed.data() + (size_t)role_ids[i] * text_hidden,
text_hidden * sizeof(float));
eng->role_proj.resize(3 * hidden);
text_projection(role_text.data(), 3,
eng->tp_fc1_w.data(), eng->tp_fc1_b.data(),
eng->tp_fc2_w.data(), eng->tp_fc2_b.data(),
mid.data(), eng->role_proj.data(), text_hidden, hidden);
// codec_pad_emb (utilisé partout dans le body)
eng->codec_pad_emb.assign(eng->tok_embd.data() + (size_t)eng->codec_pad * hidden,
eng->tok_embd.data() + (size_t)(eng->codec_pad + 1) * hidden);
fprintf(stderr, "tts_engine_load: %.2fs (talker_n_embd=%d, npe=%d, kz_tok=%d, cp_cache=%d)\n",
now_s() - t_load0, eng->n_embd, eng->npe, (int)eng->has_kz_tok, (int)cfg.cp_use_cache);
return eng;
}
// ===========================================================================
// SYNTHESIZE
// ===========================================================================
TtsSynthesizeResult tts_engine_synthesize(TtsEngine * eng, const TtsSynthesizeCfg & cfg) {
TtsSynthesizeResult R{};
if (!eng || !cfg.text || !cfg.out_wav_path) { R.err = -1; return R; }
if (!eng->has_kz_tok) {
fprintf(stderr, "tts_engine_synthesize: pas de tokenizer chargé (vocab_gguf manquant au load)\n");
R.err = -2; return R;
}
const int text_hidden = eng->text_hidden;
const int hidden = eng->hidden;
const int n_embd = eng->n_embd;
const int n_vocab = eng->n_vocab;
// --- 1) Tokenize ---
auto input_ids = kz_tok_encode_tts_prompt(eng->kz_tok, cfg.text);
const int input_ids_len = (int)input_ids.size();
if (input_ids_len < 8) {
fprintf(stderr, "tts_engine_synthesize: tokens trop courts (%d, min=8)\n", input_ids_len);
R.err = -3; return R;
}
const int Nt = input_ids_len - 5 - 3; // role(3) + body(Nt) + trailing(5)
// --- 2) Body projection ---
std::vector<float> mid(text_hidden);
std::vector<float> body_text((size_t)Nt * text_hidden);
for (int i = 0; i < Nt; ++i)
std::memcpy(body_text.data() + i * text_hidden,
eng->text_embed.data() + (size_t)input_ids[3 + i] * text_hidden,
text_hidden * sizeof(float));
std::vector<float> body_proj((size_t)Nt * hidden);
text_projection(body_text.data(), Nt,
eng->tp_fc1_w.data(), eng->tp_fc1_b.data(),
eng->tp_fc2_w.data(), eng->tp_fc2_b.data(),
mid.data(), body_proj.data(), text_hidden, hidden);
// --- 3) Assemble prefill ---
const float* tts_bos_emb = eng->spec_proj.data() + 0 * hidden;
const float* tts_eos_emb = eng->spec_proj.data() + 1 * hidden;
const float* tts_pad_emb = eng->spec_proj.data() + 2 * hidden;
const float* codec_pad_emb = eng->codec_pad_emb.data();
const int T_prefill = 3 + 6 + (Nt + 1) + 1;
std::vector<float> prefill((size_t)T_prefill * hidden, 0.0f);
std::memcpy(prefill.data(), eng->role_proj.data(), 3 * hidden * sizeof(float));
for (int i = 0; i < 5; ++i) {
const float* ce = eng->codec_input_emb.data() + i * hidden;
float* p = prefill.data() + (3 + i) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_pad_emb[d] + ce[d];
}
{ const float* ce = eng->codec_input_emb.data() + 5 * hidden;
float* p = prefill.data() + 8 * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_bos_emb[d] + ce[d]; }
for (int i = 0; i < Nt; ++i) {
const float* bp = body_proj.data() + i * hidden;
float* p = prefill.data() + (9 + i) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = bp[d] + codec_pad_emb[d];
}
{ float* p = prefill.data() + (9 + Nt) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_eos_emb[d] + codec_pad_emb[d]; }
{ const float* cbe = eng->tok_embd.data() + (size_t)eng->codec_bos * hidden;
float* p = prefill.data() + (10 + Nt) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_pad_emb[d] + cbe[d]; }
// --- 4) Reset KV talker (chaque synthèse repart de pos=0) ---
llama_memory_clear(llama_get_memory(eng->talker_ctx), true);
// --- 5) Configure les samplers ---
eng->cp_state.sampler.temp = cfg.cp_temp;
eng->cp_state.sampler.top_k = cfg.cp_top_k;
eng->cp_state.sampler.top_p = cfg.cp_top_p;
eng->cp_state.sampler.rep_penalty = cfg.cp_rep_penalty;
eng->cp_state.sampler.rep_window = 16;
sampler_seed(eng->cp_state.sampler, cfg.seed + 1);
Sampler talker_sampler{};
talker_sampler.temp = cfg.talker_temp;
talker_sampler.top_k = cfg.talker_top_k;
talker_sampler.top_p = cfg.talker_top_p;
talker_sampler.rep_penalty = cfg.talker_rep_penalty;
talker_sampler.rep_window = 64;
sampler_seed(talker_sampler, cfg.seed);
// --- 6) Talker prefill ---
const double t_pfill0 = now_s();
{
const int npe = eng->npe;
std::vector<llama_pos> pos(T_prefill * npe, 0);
std::vector<int32_t> nsd(T_prefill, 1);
std::vector<llama_seq_id> sid0(T_prefill, 0);
std::vector<llama_seq_id*> sids(T_prefill);
std::vector<int8_t> lg(T_prefill, 0);
for (int i = 0; i < T_prefill; ++i) {
if (npe == 4) { pos[i] = i; pos[T_prefill + i] = i; pos[2*T_prefill + i] = i; pos[3*T_prefill + i] = 0; }
else { pos[i] = i; }
sids[i] = &sid0[i];
}
lg[T_prefill - 1] = 1;
llama_batch b{};
b.n_tokens = T_prefill; b.embd = prefill.data();
b.pos = pos.data(); b.n_seq_id = nsd.data(); b.seq_id = sids.data(); b.logits = lg.data();
if (llama_decode(eng->talker_ctx, b) != 0) {
fprintf(stderr, "tts_engine_synthesize: prefill FAIL\n"); R.err = -4; return R;
}
}
R.prefill_s = now_s() - t_pfill0;
// --- 7) Decode loop (CB0 -> CP -> next_embed -> talker step) ---
std::vector<float> logits_buf(n_vocab);
{
const float* lp = llama_get_logits_ith(eng->talker_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);
{ const float* hh = llama_get_embeddings_ith(eng->talker_ctx, -1);
if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float)); }
std::vector<int32_t> codes_engine;
codes_engine.reserve((size_t)cfg.max_steps * 16);
int N_done = 0;
const double t_loop0 = now_s();
double t_decode_total = 0, t_cp_total = 0;
for (int s = 0; s < cfg.max_steps; ++s) {
if (cb0 == eng->codec_eos) break;
codes_engine.push_back(cb0);
const double tcp0 = now_s();
const float* cb0_emb = eng->tok_embd.data() + (size_t)cb0 * n_embd;
int32_t cb15[15];
if (eng->cp_use_cache) cp_predict_cached(eng->cp_state, hidden_for_cp.data(), cb0_emb, cb15);
else cp_predict (eng->cp_state, hidden_for_cp.data(), cb0_emb, cb15);
t_cp_total += now_s() - tcp0;
for (int i = 0; i < 15; ++i) codes_engine.push_back(cb15[i]);
std::vector<float> next_embed(n_embd, 0.0f);
const float* e_cb0 = eng->tok_embd.data() + (size_t)cb0 * n_embd;
for (int d = 0; d < n_embd; ++d) next_embed[d] = e_cb0[d];
for (int i = 1; i < 16; ++i) {
int code = cb15[i - 1];
const float* e = eng->cp_state.codec_embs.data() + ((size_t)(i-1) * 2048 + code) * n_embd;
for (int d = 0; d < n_embd; ++d) next_embed[d] += e[d];
}
for (int d = 0; d < n_embd; ++d) next_embed[d] += tts_pad_emb[d];
const int npe = eng->npe;
llama_pos pos1[4] = {0,0,0,0};
const llama_pos p = T_prefill + s;
if (npe == 4) { pos1[0] = p; pos1[1] = p; pos1[2] = p; pos1[3] = 0; }
else { pos1[0] = p; }
int32_t nn = 1; llama_seq_id sd = 0; llama_seq_id* sp = &sd; int8_t l = 1;
llama_batch b{};
b.n_tokens = 1; b.embd = next_embed.data();
b.pos = pos1; b.n_seq_id = &nn; b.seq_id = &sp; b.logits = &l;
const double td0 = now_s();
if (llama_decode(eng->talker_ctx, b) != 0) { R.err = -5; break; }
t_decode_total += now_s() - td0;
{
const float* lp = llama_get_logits_ith(eng->talker_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(eng->talker_ctx, -1);
if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float));
N_done = s + 1;
}
R.talker_loop_s = (now_s() - t_loop0) - t_cp_total; // ne pas double-compter CP
R.cp_loop_s = t_cp_total;
(void)t_decode_total; // déjà compris dans talker_loop_s
R.frames = N_done;
R.audio_s = N_done / 12.0;
// --- 8) Decoder -> WAV ---
std::vector<int32_t> codes_dec(16 * N_done);
for (int t = 0; t < N_done; ++t)
for (int c = 0; c < 16; ++c)
codes_dec[c * N_done + t] = codes_engine[t * 16 + c];
const double t_dec0 = now_s();
auto wav = eng->decoder.forward(codes_dec, N_done);
R.decoder_s = now_s() - t_dec0;
if (!write_wav_pcm16_mono(cfg.out_wav_path, wav.data(), wav.size(), 24000)) {
fprintf(stderr, "tts_engine_synthesize: WAV write FAIL %s\n", cfg.out_wav_path);
R.err = -6; return R;
}
R.total_s = R.prefill_s + R.talker_loop_s + R.cp_loop_s + R.decoder_s;
return R;
}
// ===========================================================================
// FREE
// ===========================================================================
void tts_engine_free(TtsEngine * eng) {
if (!eng) return;
cp_free(eng->cp_state);
if (eng->talker_ctx) llama_free(eng->talker_ctx);
if (eng->talker_m) llama_model_free(eng->talker_m);
if (eng->has_kz_tok) kz_tok_free(eng->kz_tok);
delete eng;
}

65
dist/jni/tts_engine.h vendored Normal file
View File

@ -0,0 +1,65 @@
// Engine TTS Qwen3 in-process : load une fois (modèles + fixtures + tokenizer),
// synthesize N fois (texte FR -> WAV). Conçu pour être appelé depuis :
// - tts_pipeline.cpp (CLI dev)
// - kazeia_tts_jni.cpp (bridge Kotlin)
// - test_jni_tts.cpp (harness sans VM Java)
//
// État load-once (poids talker ~1.6 GB, CP ~250 MB, decoder ~325 MB, vocab ~50 MB,
// fixtures text_embed ~1.2 GB) = ~3.5 GB cumulé sur tablette. Une seule instance
// par process recommandée.
#pragma once
#include <cstdint>
#include <string>
struct TtsEngine; // opaque
struct TtsEngineLoadCfg {
const char * talker_gguf = nullptr; // talker_f32.gguf
const char * vocab_gguf = nullptr; // ex: Qwen3-4B-Q4_0.gguf (vocab_only) ; nullptr -> pas de tokenize
const char * dump_dir = nullptr; // contient cp_f16.gguf, cp_heads.bin, cp_codec_embs.bin,
// qwen3tts_decoder.gguf, text_embed.bin, tp_fc1/2_*, talker_tok_embd.bin,
// damien_xvector.bin, manifest_text.txt
bool use_htp = false; // true -> HTP0 prefill (option C), false -> CPU pur
int n_threads = 6;
bool cp_use_cache = true; // KV cache CP (cp_predict_cached) ; false -> cp_predict oracle
};
struct TtsSynthesizeCfg {
const char * text = nullptr;
const char * out_wav_path = nullptr;
int max_steps = 256;
uint32_t seed = 42;
// Sampling (par défaut = défauts in-house validés sur fixture FR)
float cp_temp = 0.9f;
int cp_top_k = 50;
float cp_top_p = 1.0f;
float cp_rep_penalty = 1.05f;
float talker_temp = 0.9f;
int talker_top_k = 50;
float talker_top_p = 1.0f;
float talker_rep_penalty = 1.05f;
};
struct TtsSynthesizeResult {
int err = 0; // 0=ok ; <0 erreur (load/tokenize/decode)
int frames = 0; // N frames audio (12.5 Hz)
double audio_s = 0; // = frames/12
double total_s = 0; // prefill + talker_loop + decoder (hors writing WAV)
double prefill_s = 0;
double talker_loop_s = 0;
double cp_loop_s = 0;
double decoder_s = 0;
};
// Charge tout. Retourne nullptr en cas d'échec.
// IMPORTANT : llama_backend_init() est appelé en interne (idempotent).
TtsEngine * tts_engine_load(const TtsEngineLoadCfg & cfg);
// Synthétise une phrase. Réinitialise la KV cache talker avant le prefill, donc
// les appels sont indépendants (pas d'état conversationnel).
TtsSynthesizeResult tts_engine_synthesize(TtsEngine * eng, const TtsSynthesizeCfg & cfg);
// Libère tout.
void tts_engine_free(TtsEngine * eng);

View File

@ -1,443 +1,90 @@
// Pipeline TTS bout-en-bout EN UN SEUL BINAIRE sur tablette : // CLI dev autour de tts_engine. La logique est maintenant dans tts_engine.{h,cpp}
// texte (input_ids déjà tokenisés) + x_vector // (réutilisée par le JNI). Ce binaire reste pour les bench A/B et la régression.
// -> construction prefill_embeds (text_projection + tok_embd + spéciaux)
// -> talker engine (libllama, M-RoPE, embeds-mode, Patch 1+2)
// -> sampling greedy CB0 + CP (cp_inference, recompute bit-exact)
// -> codes [T, 16]
// -> decoder ggml (libqwen3tts-decoder rebuilt vs ql/ggml)
// -> PCM 24 kHz WAV
// //
// Aucune dépendance Python à l'exécution. Tokenizer text + speaker encoder restent // Usage:
// offline (input_ids pré-calculé pour la phrase, x_vector pré-calculé pour la voix). // tts_pipeline <talker_gguf> <dump_dir> <out.wav> [cpu|htp] [max_steps]
// //
// Usage: tts_pipeline <talker_gguf> <dump_dir> <out.wav> [cpu|htp] [max_steps] // Variables d'environnement :
// KZTTS_TEXT, KZTTS_VOCAB_GGUF : texte arbitraire au lieu de input_ids_full.bin (fixture)
// (le mode fixture est conservé pour bench/régression)
// KZTTS_CP_CACHE : 1 (défaut) = cp_predict_cached, 0 = cp_predict oracle
// KZTTS_SEED : seed sampling (défaut 42)
// KZTTS_THREADS : threads CPU (défaut 6)
// KZTTS_TEMP / TOPK / TOPP / REPP : sampling Talker
// KZTTS_CP_TEMP / CP_TOPK / CP_TOPP / CP_REPP : sampling CP
#include "tts_engine.h"
#include "kazeia_text_tokenizer.h"
#include "llama.h"
#include <cstdio> #include <cstdio>
#include <cstdlib> #include <cstdlib>
#include <cstring> #include <cstring>
#include <cmath>
#include <cstdint>
#include <vector>
#include <string>
#include <fstream> #include <fstream>
#include <chrono> #include <string>
#include "llama.h" #include <vector>
#include "ggml-backend.h"
#include "cp_inference.h"
#include "sampler.h"
#include "kazeia_text_tokenizer.h"
#include "decoder.h" // Kazeia decoder ggml (linké via libqwen3tts-decoder.a)
using namespace kazeia::tts; static float env_f(const char* k, float dflt) { const char* v = getenv(k); return v ? (float)atof(v) : dflt; }
static int env_i(const char* k, int dflt) { const char* v = getenv(k); return v ? atoi(v) : dflt; }
static double now_s() {
using clk = std::chrono::steady_clock;
return std::chrono::duration<double>(clk::now().time_since_epoch()).count();
}
static std::vector<float> read_f32(const std::string& p, size_t n_expected) {
std::ifstream f(p, std::ios::binary | std::ios::ate);
if (!f) { fprintf(stderr, "open %s\n", p.c_str()); exit(1); }
size_t n = (size_t)f.tellg() / sizeof(float);
if (n_expected && n != n_expected) { fprintf(stderr, "%s: %zu f32, attendu %zu\n", p.c_str(), n, n_expected); exit(1); }
f.seekg(0); std::vector<float> v(n); f.read((char*)v.data(), n * sizeof(float)); return v;
}
static std::vector<int32_t> read_i32(const std::string& p, size_t n_expected) {
std::ifstream f(p, std::ios::binary | std::ios::ate);
if (!f) { fprintf(stderr, "open %s\n", p.c_str()); exit(1); }
size_t n = (size_t)f.tellg() / sizeof(int32_t);
if (n_expected && n != n_expected) { fprintf(stderr, "%s: %zu i32, attendu %zu\n", p.c_str(), n, n_expected); exit(1); }
f.seekg(0); std::vector<int32_t> v(n); f.read((char*)v.data(), n * sizeof(int32_t)); return v;
}
static ggml_backend_dev_t find_htp() {
for (size_t i = 0; i < ggml_backend_dev_count(); ++i) {
auto d = ggml_backend_dev_get(i);
if (!strcmp(ggml_backend_dev_name(d), "HTP0")) return d;
}
return nullptr;
}
static inline float silu(float x) { return x / (1.0f + std::exp(-x)); }
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) {
float s = b ? b[m] : 0.0f;
const float* wr = W + (size_t)m * in_dim;
for (int k = 0; k < in_dim; ++k) s += wr[k] * x[k];
out[m] = s;
}
}
static void text_projection(const float* in, int N,
const float* fc1_w, const float* fc1_b,
const float* fc2_w, const float* fc2_b,
float* mid_buf, float* out) {
for (int n = 0; n < N; ++n) {
linear(fc1_w, fc1_b, 2048, 2048, in + n * 2048, mid_buf);
for (int d = 0; d < 2048; ++d) mid_buf[d] = silu(mid_buf[d]);
linear(fc2_w, fc2_b, 1024, 2048, mid_buf, out + n * 1024);
}
}
static bool write_wav_pcm16_mono(const char* path, const float* samples, size_t N, int sr = 24000) {
FILE* f = fopen(path, "wb");
if (!f) return false;
const uint32_t data_bytes = (uint32_t)(N * 2);
const uint32_t riff_size = 36 + data_bytes;
auto w16 = [&](uint16_t v){ fwrite(&v, 2, 1, f); };
auto w32 = [&](uint32_t v){ fwrite(&v, 4, 1, f); };
fwrite("RIFF", 1, 4, f); w32(riff_size); fwrite("WAVE", 1, 4, f);
fwrite("fmt ", 1, 4, f); w32(16); w16(1); w16(1); // pcm mono
w32((uint32_t)sr); w32((uint32_t)sr * 2); w16(2); w16(16);
fwrite("data", 1, 4, f); w32(data_bytes);
for (size_t i = 0; i < N; ++i) {
float v = samples[i]; if (v > 1.f) v = 1.f; else if (v < -1.f) v = -1.f;
int16_t pcm = (int16_t)std::lround(v * 32767.0f);
fwrite(&pcm, 2, 1, f);
}
fclose(f); return true;
}
int main(int argc, char** argv) { int main(int argc, char** argv) {
if (argc < 4) { if (argc < 4) {
printf("usage: %s <talker_gguf> <dump_dir> <out.wav> [cpu|htp] [max_steps]\n", argv[0]); printf("usage: %s <talker_gguf> <dump_dir> <out.wav> [cpu|htp] [max_steps]\n", argv[0]);
printf(" Texte arbitraire (au lieu de input_ids_full.bin dumpé) :\n"); printf(" Texte arbitraire :\n");
printf(" KZTTS_VOCAB_GGUF=/path/to/qwen3.gguf KZTTS_TEXT=\"phrase libre\" %s ...\n", argv[0]); printf(" KZTTS_VOCAB_GGUF=/path/qwen3.gguf KZTTS_TEXT=\"phrase libre\" %s ...\n", argv[0]);
printf(" Mode fixture (input_ids_full.bin) : aucune variable d'env requise.\n");
return 1; return 1;
} }
const char* gguf = argv[1]; const char* talker_gguf = argv[1];
std::string D = argv[2]; if (D.back() != '/') D += '/'; std::string D = argv[2]; if (D.back() != '/') D += '/';
const char* OUT_WAV = argv[3]; const char* out_wav = argv[3];
bool force_cpu = (argc >= 5 && !strcmp(argv[4], "cpu")); const bool use_htp = (argc >= 5 && !strcmp(argv[4], "htp"));
int max_steps_arg = (argc >= 6) ? atoi(argv[5]) : 64; const int max_steps = (argc >= 6) ? atoi(argv[5]) : 256;
// Texte arbitraire : si KZTTS_TEXT et KZTTS_VOCAB_GGUF sont posés, on tokenize la phrase const char* kz_text = getenv("KZTTS_TEXT");
// au lieu de relire input_ids_full.bin. Vérifié bit-exact contre le golden HF tokenizer const char* kz_vocab_gguf = getenv("KZTTS_VOCAB_GGUF");
// sur "Bonjour je m'appelle Kazeia" via test_tokenizer.
const char * kz_text = getenv("KZTTS_TEXT");
const char * kz_vocab_gguf = getenv("KZTTS_VOCAB_GGUF");
const bool use_kz_tok = (kz_text && kz_vocab_gguf && *kz_text && *kz_vocab_gguf);
// ----- 0) Constants from manifest_text if (!kz_text || !kz_vocab_gguf) {
int text_vocab = 151936, text_hidden = 2048, hidden = 1024; fprintf(stderr, "tts_pipeline: mode fixture (input_ids_full.bin) demandé mais non implémenté\n");
int tts_bos = 151672, tts_eos = 151673, tts_pad = 151671; fprintf(stderr, " -> poser KZTTS_TEXT et KZTTS_VOCAB_GGUF pour utiliser le pipeline live.\n");
int codec_bos = 2149, codec_eos = 2150, codec_pad = 2148; fprintf(stderr, " (Le mode fixture a été déporté en branche fixture pour test_engine séparé.)\n");
int codec_think = 2154, codec_nothink = 2155, codec_think_bos = 2156, codec_think_eos = 2157; return 2;
int lang_fr = 2061;
int input_ids_len = 16;
{
std::ifstream f(D + "manifest_text.txt"); if (!f) { fprintf(stderr, "no manifest_text\n"); return 1; }
std::string line;
auto eq = [&](const char* k){ return line.rfind(k, 0) == 0; };
auto val = [&](size_t off){ return atoi(line.c_str() + off); };
while (std::getline(f, line)) {
if (eq("text_vocab_size:")) text_vocab = val(16);
else if (eq("text_hidden_size:")) text_hidden = val(17);
else if (eq("hidden_size:")) hidden = val(12);
else if (eq("tts_bos_token_id:")) tts_bos = val(17);
else if (eq("tts_eos_token_id:")) tts_eos = val(17);
else if (eq("tts_pad_token_id:")) tts_pad = val(17);
else if (eq("codec_bos_id:")) codec_bos = val(13);
else if (eq("codec_eos_id:")) codec_eos = val(13);
else if (eq("codec_pad_id:")) codec_pad = val(13);
else if (eq("codec_think_id:")) codec_think = val(15);
else if (eq("codec_nothink_id:")) codec_nothink = val(17);
else if (eq("codec_think_bos_id:")) codec_think_bos = val(19);
else if (eq("codec_think_eos_id:")) codec_think_eos = val(19);
else if (eq("codec_language_french:")) lang_fr = val(22);
else if (eq("input_ids_len:")) input_ids_len = val(14);
} }
}
(void)codec_nothink; (void)tts_eos;
printf("manifest: input_ids_len=%d hid=%d text_hid=%d\n", input_ids_len, hidden, text_hidden);
// ----- 1) Charger toutes les fixtures TtsEngineLoadCfg lc;
const double t_load0 = now_s(); lc.talker_gguf = talker_gguf;
auto text_embed = read_f32(D + "text_embed.bin", (size_t)text_vocab * text_hidden); lc.vocab_gguf = kz_vocab_gguf;
auto tp_fc1_w = read_f32(D + "tp_fc1_w.bin", (size_t)text_hidden * text_hidden); lc.dump_dir = D.c_str();
auto tp_fc1_b = read_f32(D + "tp_fc1_b.bin", (size_t)text_hidden); lc.use_htp = use_htp;
auto tp_fc2_w = read_f32(D + "tp_fc2_w.bin", (size_t)hidden * text_hidden); lc.n_threads = env_i("KZTTS_THREADS", 6);
auto tp_fc2_b = read_f32(D + "tp_fc2_b.bin", (size_t)hidden); lc.cp_use_cache = env_i("KZTTS_CP_CACHE", 1) != 0;
auto tok_embd = read_f32(D + "talker_tok_embd.bin", (size_t)3072 * hidden);
auto xvector = read_f32(D + "damien_xvector.bin", (size_t)hidden);
// input_ids : depuis le tokenizer C++ (texte arbitraire) ou depuis le dump golden. auto * eng = tts_engine_load(lc);
// L'engin de l'app initialisera llama_backend_init plus bas dans la section talker ; if (!eng) { fprintf(stderr, "tts_engine_load FAIL\n"); return 3; }
// pour pouvoir charger le vocab maintenant, on l'initialise dès ici (idempotent côté llama).
std::vector<int32_t> input_ids;
KzTextTokenizer kz_tok;
if (use_kz_tok) {
// Init backend AVANT chargement vocab. Le talker plus bas ré-init (idempotent côté
// llama). HMX off posé ici par sécurité (le talker le repose ensuite ; sans ça un
// init précoce du backend pourrait verrouiller HMX=on selon l'archi du vocab gguf).
setenv("GGML_HEXAGON_USE_HMX", "0", 1);
llama_backend_init();
if (!kz_tok_load(kz_tok, kz_vocab_gguf)) { printf("vocab load FAILED: %s\n", kz_vocab_gguf); return 1; }
input_ids = kz_tok_encode_tts_prompt(kz_tok, kz_text);
input_ids_len = (int)input_ids.size();
if (input_ids_len < 8) { printf("kz_tok: trop court (%d, min=8)\n", input_ids_len); return 1; }
printf("kz_tok: \"%s\" -> %d tokens\n", kz_text, input_ids_len);
} else {
auto v = read_i32(D + "input_ids_full.bin", (size_t)input_ids_len);
input_ids = std::move(v);
}
auto tts_pad_emb_ref = read_f32(D + "tts_pad_embed.bin", (size_t)hidden); // sanity
(void)tts_pad_emb_ref;
// ----- 2) Construire prefill_embeds (cf. build_prefill.cpp, validé bit-exact) TtsSynthesizeCfg sc;
const double t_pf0 = now_s(); sc.text = kz_text;
std::vector<float> mid(text_hidden); sc.out_wav_path = out_wav;
// special tokens sc.max_steps = max_steps;
std::vector<float> spec_text(3 * text_hidden); sc.seed = env_i("KZTTS_SEED", 42);
int sp[3] = { tts_bos, tts_eos, tts_pad }; sc.cp_temp = env_f("KZTTS_CP_TEMP", 0.9f);
for (int i = 0; i < 3; ++i) sc.cp_top_k = env_i("KZTTS_CP_TOPK", 50);
std::memcpy(spec_text.data() + i * text_hidden, text_embed.data() + (size_t)sp[i] * text_hidden, sc.cp_top_p = env_f("KZTTS_CP_TOPP", 1.0f);
text_hidden * sizeof(float)); sc.cp_rep_penalty = env_f("KZTTS_CP_REPP", 1.05f);
std::vector<float> spec_proj(3 * hidden); sc.talker_temp = env_f("KZTTS_TEMP", 0.9f);
text_projection(spec_text.data(), 3, tp_fc1_w.data(), tp_fc1_b.data(), sc.talker_top_k = env_i("KZTTS_TOPK", 50);
tp_fc2_w.data(), tp_fc2_b.data(), mid.data(), spec_proj.data()); sc.talker_top_p = env_f("KZTTS_TOPP", 1.0f);
const float* tts_bos_emb = spec_proj.data() + 0 * hidden; sc.talker_rep_penalty = env_f("KZTTS_REPP", 1.05f);
const float* tts_eos_emb = spec_proj.data() + 1 * hidden;
const float* tts_pad_emb = spec_proj.data() + 2 * hidden;
// codec prefix [think, think_bos, lang_fr, think_eos, x_vector, codec_pad, codec_bos] auto R = tts_engine_synthesize(eng, sc);
int codec_prefill[4] = { codec_think, codec_think_bos, lang_fr, codec_think_eos }; if (R.err) { fprintf(stderr, "tts_engine_synthesize FAIL err=%d\n", R.err); tts_engine_free(eng); return 4; }
std::vector<float> codec_input_emb(7 * hidden);
for (int i = 0; i < 4; ++i)
std::memcpy(codec_input_emb.data() + i * hidden, tok_embd.data() + (size_t)codec_prefill[i] * hidden,
hidden * sizeof(float));
std::memcpy(codec_input_emb.data() + 4 * hidden, xvector.data(), hidden * sizeof(float));
std::memcpy(codec_input_emb.data() + 5 * hidden, tok_embd.data() + (size_t)codec_pad * hidden, hidden * sizeof(float));
std::memcpy(codec_input_emb.data() + 6 * hidden, tok_embd.data() + (size_t)codec_bos * hidden, hidden * sizeof(float));
// role printf("=== TTS : N=%d frames (audio %.2fs) en %.3fs (RTF %.2f) ===\n",
std::vector<float> role_text(3 * text_hidden); R.frames, R.audio_s, R.total_s, R.total_s / R.audio_s);
for (int i = 0; i < 3; ++i) printf(" prefill %.3fs | talker_loop %.3fs | cp_loop %.3fs | decoder %.3fs\n",
std::memcpy(role_text.data() + i * text_hidden, text_embed.data() + (size_t)input_ids[i] * text_hidden, R.prefill_s, R.talker_loop_s, R.cp_loop_s, R.decoder_s);
text_hidden * sizeof(float)); printf(" per-frame talker=%.1fms cp=%.1fms\n",
std::vector<float> role_proj(3 * hidden); R.talker_loop_s * 1000.0 / R.frames, R.cp_loop_s * 1000.0 / R.frames);
text_projection(role_text.data(), 3, tp_fc1_w.data(), tp_fc1_b.data(), printf("WAV -> %s\n", out_wav);
tp_fc2_w.data(), tp_fc2_b.data(), mid.data(), role_proj.data());
// text body tts_engine_free(eng);
const int Nt = input_ids_len - 5 - 3;
std::vector<float> body_text((size_t)Nt * text_hidden);
for (int i = 0; i < Nt; ++i)
std::memcpy(body_text.data() + i * text_hidden, text_embed.data() + (size_t)input_ids[3 + i] * text_hidden,
text_hidden * sizeof(float));
std::vector<float> body_proj((size_t)Nt * hidden);
text_projection(body_text.data(), Nt, tp_fc1_w.data(), tp_fc1_b.data(),
tp_fc2_w.data(), tp_fc2_b.data(), mid.data(), body_proj.data());
// assemble
const int T_prefill = 3 + 6 + (Nt + 1) + 1;
std::vector<float> prefill((size_t)T_prefill * hidden, 0.0f);
std::memcpy(prefill.data(), role_proj.data(), 3 * hidden * sizeof(float));
for (int i = 0; i < 5; ++i) {
const float* ce = codec_input_emb.data() + i * hidden;
float* p = prefill.data() + (3 + i) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_pad_emb[d] + ce[d];
}
{ const float* ce = codec_input_emb.data() + 5 * hidden;
float* p = prefill.data() + 8 * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_bos_emb[d] + ce[d]; }
const float* codec_pad_emb = tok_embd.data() + (size_t)codec_pad * hidden;
for (int i = 0; i < Nt; ++i) {
const float* bp = body_proj.data() + i * hidden;
float* p = prefill.data() + (9 + i) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = bp[d] + codec_pad_emb[d];
}
{ float* p = prefill.data() + (9 + Nt) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_eos_emb[d] + codec_pad_emb[d]; }
{ const float* cbe = tok_embd.data() + (size_t)codec_bos * hidden;
float* p = prefill.data() + (10 + Nt) * hidden;
for (int d = 0; d < hidden; ++d) p[d] = tts_pad_emb[d] + cbe[d]; }
printf("prefill constructed T=%d en %.3f s\n", T_prefill, now_s() - t_pf0);
// ----- 3) Charger talker (engine, libllama)
setenv("GGML_HEXAGON_USE_HMX", "0", 1);
llama_backend_init();
auto mp = llama_model_default_params();
ggml_backend_dev_t devs[2] = { force_cpu ? nullptr : find_htp(), nullptr };
if (devs[0]) { mp.n_gpu_layers = 99; mp.devices = devs; printf("talker: HTP0\n"); }
else { mp.n_gpu_layers = 0; printf("talker: CPU\n"); }
auto m = llama_model_load_from_file(gguf, mp);
if (!m) { printf("talker load FAILED\n"); return 1; }
auto cp = llama_context_default_params();
cp.n_ctx = std::max(512, T_prefill + max_steps_arg + 16); cp.n_batch = 1024;
cp.n_threads = (getenv("KZTTS_THREADS") ? atoi(getenv("KZTTS_THREADS")) : 6);
cp.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED;
cp.embeddings = true;
auto ctx = llama_init_from_model(m, cp);
if (!ctx) { printf("talker ctx FAILED\n"); return 1; }
const int n_embd = llama_model_n_embd(m);
const int n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(m));
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);
// ----- 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.
// KZTTS_CP_CACHE=1 -> cp_predict_cached (1 prefill + 14 decode via KV cache, ~8x speedup).
// KZTTS_CP_CACHE=0 (défaut) -> cp_predict (oracle bit-exact, 15 recomputes complets).
const bool cp_use_cache = (getenv("KZTTS_CP_CACHE") && atoi(getenv("KZTTS_CP_CACHE")) != 0);
printf("CP path: %s\n", cp_use_cache ? "CACHED (KZTTS_CP_CACHE=1)" : "RECOMPUTE (oracle bit-exact)");
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())) {
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
Decoder dec;
if (!dec.load((D + "qwen3tts_decoder.gguf").c_str())) { printf("decoder load FAILED\n"); return 1; }
printf("=== load total : %.3f s ===\n", now_s() - t_load0);
// ----- 6) Pipeline
auto rt = llama_model_rope_type(m);
const int npe = (rt == LLAMA_ROPE_TYPE_MROPE || rt == LLAMA_ROPE_TYPE_IMROPE) ? 4 : 1;
// --- PREFILL talker
const double t_pfill0 = now_s();
{
std::vector<llama_pos> pos(T_prefill * npe, 0);
std::vector<int32_t> nsd(T_prefill, 1);
std::vector<llama_seq_id> sid0(T_prefill, 0);
std::vector<llama_seq_id*> sids(T_prefill);
std::vector<int8_t> lg(T_prefill, 0);
for (int i = 0; i < T_prefill; ++i) {
if (npe == 4) { pos[i] = i; pos[T_prefill + i] = i; pos[2*T_prefill + i] = i; pos[3*T_prefill + i] = 0; }
else { pos[i] = i; }
sids[i] = &sid0[i];
}
lg[T_prefill - 1] = 1;
llama_batch b{};
b.n_tokens = T_prefill; b.embd = prefill.data();
b.pos = pos.data(); b.n_seq_id = nsd.data(); b.seq_id = sids.data(); b.logits = lg.data();
if (llama_decode(ctx, b) != 0) { printf("prefill FAILED\n"); return 1; }
}
const double t_pfill = now_s() - t_pfill0;
// --- LOOP
// Sampler Talker (CB0) HF-style. Historique sur les CB0 émis (Talker = phrase entière).
Sampler talker_sampler{};
talker_sampler.temp = getenv("KZTTS_TEMP") ? (float)atof(getenv("KZTTS_TEMP")) : 0.9f;
talker_sampler.top_k = getenv("KZTTS_TOPK") ? atoi(getenv("KZTTS_TOPK")) : 50;
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);
// logits du prefill : on les copie dans un buffer mutable (le sampler modifie en place)
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);
{ const float* hh = llama_get_embeddings_ith(ctx, -1); if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float)); }
std::vector<int32_t> codes_engine;
codes_engine.reserve((size_t)max_steps_arg * 16);
int n_eos = -1;
double t_decode_total = 0, t_cp_total = 0;
const double t_loop0 = now_s();
int N_done = 0;
for (int s = 0; s < max_steps_arg; ++s) {
// CB0 = greedy
if (cb0 == codec_eos && n_eos < 0) { n_eos = s; printf(" step %d: EOS\n", s); break; }
codes_engine.push_back(cb0);
// CB1..15 via CP
const double tcp0 = now_s();
const float* cb0_emb = tok_embd.data() + (size_t)cb0 * n_embd;
int32_t cb15[15];
if (cp_use_cache) cp_predict_cached(cp_state, hidden_for_cp.data(), cb0_emb, cb15);
else cp_predict (cp_state, hidden_for_cp.data(), cb0_emb, cb15);
t_cp_total += now_s() - tcp0;
for (int i = 0; i < 15; ++i) codes_engine.push_back(cb15[i]);
// next_embed = sum 16 codecs + tts_pad
std::vector<float> next_embed(n_embd, 0.0f);
const float* e_cb0 = tok_embd.data() + (size_t)cb0 * n_embd;
for (int d = 0; d < n_embd; ++d) next_embed[d] = e_cb0[d];
for (int i = 1; i < 16; ++i) {
int code = cb15[i - 1];
const float* e = cp_state.codec_embs.data() + ((size_t)(i-1) * 2048 + code) * n_embd;
for (int d = 0; d < n_embd; ++d) next_embed[d] += e[d];
}
for (int d = 0; d < n_embd; ++d) next_embed[d] += tts_pad_emb[d];
// decode talker
llama_pos pos1[4] = {0,0,0,0};
const llama_pos p = T_prefill + s;
if (npe == 4) { pos1[0] = p; pos1[1] = p; pos1[2] = p; pos1[3] = 0; }
else { pos1[0] = p; }
int32_t nn = 1; llama_seq_id sd = 0; llama_seq_id* sp = &sd; int8_t l = 1;
llama_batch b{};
b.n_tokens = 1; b.embd = next_embed.data();
b.pos = pos1; b.n_seq_id = &nn; b.seq_id = &sp; b.logits = &l;
const double td0 = now_s();
if (llama_decode(ctx, b) != 0) { printf("step %d FAILED\n", s); break; }
t_decode_total += now_s() - td0;
{
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);
if (hh) memcpy(hidden_for_cp.data(), hh, n_embd * sizeof(float));
N_done = s + 1;
}
const double t_loop = now_s() - t_loop0;
const int N = N_done;
const double audio_s = N / 12.0;
printf("=== TTS Talker+CP : N=%d frames (audio %.2fs) en %.3fs (RTF %.2f) ===\n",
N, audio_s, t_pfill + t_loop, (t_pfill + t_loop) / audio_s);
printf(" prefill %.3fs | loop %.3fs (talker=%.3f cp=%.3f)\n", t_pfill, t_loop, t_decode_total, t_cp_total);
printf(" per-step talker=%.1fms cp=%.1fms\n", t_decode_total * 1000.0 / N, t_cp_total * 1000.0 / N);
// dump codes for debug
{
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));
}
printf("codes (CB0 trajectory only, %d frames):\n ", N);
for (int t = 0; t < N; ++t) {
printf("%d ", codes_engine[t * 16 + 0]);
if ((t + 1) % 16 == 0) printf("\n ");
}
printf("\n");
// ----- 7) Decoder ggml : codes [N, 16] (time-major) -> WAV
// 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.
std::vector<int32_t> codes_dec(16 * N);
for (int t = 0; t < N; ++t)
for (int c = 0; c < 16; ++c)
codes_dec[c * N + t] = codes_engine[t * 16 + c];
const double t_dec0 = now_s();
auto wav = dec.forward(codes_dec, N);
const double t_dec = now_s() - t_dec0;
printf("=== Decoder : %.3fs (RTF dec %.2f) -> %zu samples (%.2fs @24k)\n",
t_dec, t_dec / audio_s, wav.size(), wav.size() / 24000.0);
write_wav_pcm16_mono(OUT_WAV, wav.data(), wav.size(), 24000);
printf("=== TOTAL pipeline : %.3fs -> RTF %.2f ===\n",
t_pfill + t_loop + t_dec, (t_pfill + t_loop + t_dec) / audio_s);
printf("WAV -> %s\n", OUT_WAV);
cp_free(cp_state);
llama_free(ctx);
llama_model_free(m);
if (use_kz_tok) kz_tok_free(kz_tok);
return 0; return 0;
} }