chantier B TTS #1: API JNI embeds-only + M-RoPE qwen3 dense
Premières 2 pieces du chantier B (porter le Talker Qwen3-TTS sur l'engine pour TTS standalone tablette, sans dépendre du pipeline Python). P1 - API embeds-only dans le JNI (dist/jni/kazeia_engine_jni.cpp) Le Talker a vocab=3072 codes audio (pas de BPE), entrée = embeds pré-mélangés (text+x-vector au prefill, sum 16 codecs + tts_pad au decode), sortie = logits[3072] + hidden[1024] (pour le Code Predictor). Solution: pas un appel monolithique, 5 primitives qui laissent l'orchestration côté caller: - nEmbd(h), nVocab(h) - dimensionnement des buffers Kotlin - resetEmbeds(h) - KV clear + pos=0 - prefillEmbeds(h, embds[T*n_embd], T, outHidden[n_embd]) - decodeEmbed(h, embd[n_embd], outLogits[vocab], outHidden[n_embd]) Implementation: llama_set_embeddings(ctx, true), llama_batch.embd au lieu de .token, logits=1 sur la derniere position seulement (economie KV). KEngine porte un compteur pos_embd interne. Kotlin wrapper miroir dans dist/jni/EngineLlmEngine.kt. P2 - M-RoPE qwen3 dense (sous-module ql, commit c1609a1) Pointeur sous-module avance pour embarquer le support. Validation device: dist/jni/test_talker.cpp charge talker_f32.gguf, lance llama_decode embeds-only T=4 prefill puis 1 step, recupere hidden+logits sans crash. Log device: print_info rope_type=8, mrope sections=[24,20,20,0], prefill OK / step OK. argmax stable (entree bidon, juste smoke). Libs rebuiltes: dist/lib/libllama.so + libkazeia_engine.so + libggml-base.so. Reste P3 = cablage Talker->CP->Decoder cote orchestration. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
9af47bf388
commit
50814f3bca
|
|
@ -8,6 +8,17 @@ class EngineJni {
|
||||||
external fun generateRaw(h: Long, prompt: String, maxTok: Int): String // prompt complet déjà formaté
|
external fun generateRaw(h: Long, prompt: String, maxTok: Int): String // prompt complet déjà formaté
|
||||||
external fun reset(h: Long)
|
external fun reset(h: Long)
|
||||||
external fun free(h: Long)
|
external fun free(h: Long)
|
||||||
|
|
||||||
|
// -- TTS Talker (Qwen3-TTS) : I/O en embeddings, pas en tokens (vocab=3072 codes audio, pas de BPE).
|
||||||
|
// Le caller (cf. TalkerEngine.kt) construit les embeds de prefill (text+x-vector) et de step
|
||||||
|
// (sum 16 codecs + tts_pad + trailing_text_hidden), fait l'échantillonnage, et compose le
|
||||||
|
// pipeline Talker→CP→Decoder. Retours conventionnels : 0 = OK, négatif = erreur.
|
||||||
|
external fun nEmbd(h: Long): Int // taille d'un embed (1024 pour Talker-0.6B)
|
||||||
|
external fun nVocab(h: Long): Int // 3072 pour le Talker
|
||||||
|
external fun resetEmbeds(h: Long) // KV clear + pos=0, à appeler en début de génération TTS
|
||||||
|
external fun prefillEmbeds(h: Long, embdFlat: FloatArray, t: Int, outHidden: FloatArray): Int
|
||||||
|
external fun decodeEmbed(h: Long, embd: FloatArray, outLogits: FloatArray, outHidden: FloatArray): Int
|
||||||
|
|
||||||
companion object { init { System.loadLibrary("kazeia_engine") } }
|
companion object { init { System.loadLibrary("kazeia_engine") } }
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -32,6 +32,7 @@ struct KEngine {
|
||||||
llama_model* m_h; llama_context* c_h; // prefill HTP (nullptr si pas de HTP)
|
llama_model* m_h; llama_context* c_h; // prefill HTP (nullptr si pas de HTP)
|
||||||
llama_model* m_c; llama_context* c_c; // decode CPU (nullptr si HTP contexte-unique)
|
llama_model* m_c; llama_context* c_c; // decode CPU (nullptr si HTP contexte-unique)
|
||||||
const llama_vocab* v; llama_sampler* s;
|
const llama_vocab* v; llama_sampler* s;
|
||||||
|
int pos_embd; // position courante en mode talker embeds-only (TTS)
|
||||||
};
|
};
|
||||||
|
|
||||||
static llama_context* make_ctx(llama_model* m, int nctx, int nthreads, enum llama_flash_attn_type fa) {
|
static llama_context* make_ctx(llama_model* m, int nctx, int nthreads, enum llama_flash_attn_type fa) {
|
||||||
|
|
@ -132,7 +133,8 @@ Java_com_kazeia_llm_EngineJni_load(JNIEnv* e, jobject, jstring path, jint nctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
auto* k = new KEngine{ m_h, c_h, m_c, c_c,
|
auto* k = new KEngine{ m_h, c_h, m_c, c_c,
|
||||||
llama_model_get_vocab(m_h ? m_h : m_c), llama_sampler_init_greedy() };
|
llama_model_get_vocab(m_h ? m_h : m_c), llama_sampler_init_greedy(),
|
||||||
|
/*pos_embd=*/0 };
|
||||||
return (jlong) k;
|
return (jlong) k;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -174,3 +176,124 @@ Java_com_kazeia_llm_EngineJni_free(JNIEnv*, jobject, jlong h){
|
||||||
if (k->m_c) llama_model_free(k->m_c);
|
if (k->m_c) llama_model_free(k->m_c);
|
||||||
delete k;
|
delete k;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ============================================================================
|
||||||
|
// TTS Talker API : I/O en embeddings, pas en tokens.
|
||||||
|
// Le Talker (Qwen3-TTS) a vocab = 3072 codes audio (pas de BPE), entrée = embeds
|
||||||
|
// pré-mélangés (text+x-vector au prefill, sum 16 codecs + pad au decode), sortie
|
||||||
|
// = logits sur les 3072 codes audio + hidden state (1024) pour le Code Predictor.
|
||||||
|
//
|
||||||
|
// Contrainte : le contexte est partagé avec generate/generateRaw (mêmes KV-cache).
|
||||||
|
// Appeler resetEmbeds() avant de switcher entre mode texte et mode TTS, idem en
|
||||||
|
// début de génération TTS pour repartir d'un KV propre.
|
||||||
|
//
|
||||||
|
// Routing : le talker est qwen3 dense -> charge via le chemin DENSE (HTP HVX,
|
||||||
|
// HMX off), donc c_h non-null et c_c null. pf == dec == c_h, exactement comme
|
||||||
|
// kengine_run.
|
||||||
|
// ============================================================================
|
||||||
|
|
||||||
|
extern "C" JNIEXPORT jint JNICALL
|
||||||
|
Java_com_kazeia_llm_EngineJni_nEmbd(JNIEnv*, jobject, jlong h) {
|
||||||
|
auto* k = (KEngine*) h;
|
||||||
|
return llama_model_n_embd(k->m_h ? k->m_h : k->m_c);
|
||||||
|
}
|
||||||
|
|
||||||
|
extern "C" JNIEXPORT jint JNICALL
|
||||||
|
Java_com_kazeia_llm_EngineJni_nVocab(JNIEnv*, jobject, jlong h) {
|
||||||
|
auto* k = (KEngine*) h;
|
||||||
|
return llama_vocab_n_tokens(k->v);
|
||||||
|
}
|
||||||
|
|
||||||
|
extern "C" JNIEXPORT void JNICALL
|
||||||
|
Java_com_kazeia_llm_EngineJni_resetEmbeds(JNIEnv*, jobject, jlong h) {
|
||||||
|
auto* k = (KEngine*) h;
|
||||||
|
if (k->c_h) llama_memory_clear(llama_get_memory(k->c_h), true);
|
||||||
|
if (k->c_c) llama_memory_clear(llama_get_memory(k->c_c), true);
|
||||||
|
k->pos_embd = 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper interne : décode un batch d'embeddings, ne demande la sortie que pour la
|
||||||
|
// dernière position (économie KV + logits), récupère hidden[n_embd] dans out_hidden
|
||||||
|
// si fourni. Avance k->pos_embd de T. Renvoie 0 si OK.
|
||||||
|
static int decode_embd_batch(KEngine* k, const float* embd, int T, float* out_hidden) {
|
||||||
|
llama_context* ctx = k->c_h ? k->c_h : k->c_c;
|
||||||
|
llama_set_embeddings(ctx, true);
|
||||||
|
|
||||||
|
std::vector<llama_pos> pos (T);
|
||||||
|
std::vector<int32_t> nsd (T, 1);
|
||||||
|
std::vector<llama_seq_id> sid0(T, 0);
|
||||||
|
std::vector<llama_seq_id*> sids(T);
|
||||||
|
std::vector<int8_t> lg (T, 0);
|
||||||
|
for (int i = 0; i < T; ++i) { pos[i] = k->pos_embd + i; sids[i] = &sid0[i]; }
|
||||||
|
lg[T-1] = 1; // n'output logits/embeddings que pour la dernière position
|
||||||
|
|
||||||
|
llama_batch b{};
|
||||||
|
b.n_tokens = T;
|
||||||
|
b.token = nullptr;
|
||||||
|
b.embd = const_cast<float*>(embd);
|
||||||
|
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) return -1;
|
||||||
|
k->pos_embd += T;
|
||||||
|
|
||||||
|
if (out_hidden) {
|
||||||
|
const float* eh = llama_get_embeddings_ith(ctx, -1);
|
||||||
|
if (!eh) return -2;
|
||||||
|
int n_embd = llama_model_n_embd(k->m_h ? k->m_h : k->m_c);
|
||||||
|
memcpy(out_hidden, eh, sizeof(float) * n_embd);
|
||||||
|
}
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prefill embeds : T positions, renvoie le hidden state de la dernière dans out_hidden[n_embd].
|
||||||
|
// embd_flat doit être de longueur T * n_embd (float, row-major : position 0 d'abord).
|
||||||
|
extern "C" JNIEXPORT jint JNICALL
|
||||||
|
Java_com_kazeia_llm_EngineJni_prefillEmbeds(JNIEnv* e, jobject, jlong h,
|
||||||
|
jfloatArray embd_flat, jint T,
|
||||||
|
jfloatArray out_hidden) {
|
||||||
|
auto* k = (KEngine*) h;
|
||||||
|
if (T <= 0) return -10;
|
||||||
|
int n_embd = llama_model_n_embd(k->m_h ? k->m_h : k->m_c);
|
||||||
|
if (e->GetArrayLength(embd_flat) != T * n_embd) return -11;
|
||||||
|
if (e->GetArrayLength(out_hidden) != n_embd) return -12;
|
||||||
|
|
||||||
|
jfloat* embd = e->GetFloatArrayElements(embd_flat, nullptr);
|
||||||
|
std::vector<float> hidden(n_embd);
|
||||||
|
int rc = decode_embd_batch(k, embd, T, hidden.data());
|
||||||
|
e->ReleaseFloatArrayElements(embd_flat, embd, JNI_ABORT);
|
||||||
|
if (rc != 0) return rc;
|
||||||
|
e->SetFloatArrayRegion(out_hidden, 0, n_embd, hidden.data());
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Decode one embed step : avance une position. Renvoie logits[vocab] + hidden[n_embd].
|
||||||
|
// Le caller fait l'échantillonnage (greedy / temp / top_k / rep_penalty) et compose
|
||||||
|
// l'embed suivant (sum des 16 codec_embs + tts_pad_embed + trailing_text_hidden).
|
||||||
|
extern "C" JNIEXPORT jint JNICALL
|
||||||
|
Java_com_kazeia_llm_EngineJni_decodeEmbed(JNIEnv* e, jobject, jlong h,
|
||||||
|
jfloatArray embd,
|
||||||
|
jfloatArray out_logits,
|
||||||
|
jfloatArray out_hidden) {
|
||||||
|
auto* k = (KEngine*) h;
|
||||||
|
int n_embd = llama_model_n_embd(k->m_h ? k->m_h : k->m_c);
|
||||||
|
int n_vocab = llama_vocab_n_tokens(k->v);
|
||||||
|
if (e->GetArrayLength(embd) != n_embd) return -11;
|
||||||
|
if (e->GetArrayLength(out_logits) != n_vocab) return -12;
|
||||||
|
if (e->GetArrayLength(out_hidden) != n_embd) return -13;
|
||||||
|
|
||||||
|
jfloat* ev = e->GetFloatArrayElements(embd, nullptr);
|
||||||
|
std::vector<float> hidden(n_embd);
|
||||||
|
int rc = decode_embd_batch(k, ev, 1, hidden.data());
|
||||||
|
e->ReleaseFloatArrayElements(embd, ev, JNI_ABORT);
|
||||||
|
if (rc != 0) return rc;
|
||||||
|
|
||||||
|
llama_context* ctx = k->c_h ? k->c_h : k->c_c;
|
||||||
|
const float* lg = llama_get_logits_ith(ctx, -1);
|
||||||
|
if (!lg) return -14;
|
||||||
|
e->SetFloatArrayRegion(out_logits, 0, n_vocab, lg);
|
||||||
|
e->SetFloatArrayRegion(out_hidden, 0, n_embd, hidden.data());
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,118 @@
|
||||||
|
// Smoke test : charge le Talker (qwen3 dense + M-RoPE) via le DENSE path engine,
|
||||||
|
// puis exerce les primitives embeds (prefillEmbeds/decodeEmbed) avec un input bidon.
|
||||||
|
//
|
||||||
|
// But : prouver que (a) Patch 2 active M-RoPE (log "mrope sections = [24,20,20,0]"),
|
||||||
|
// (b) Patch 1 fait tourner llama_decode embeds-only + ressort hidden state + logits.
|
||||||
|
// Pas de vérification de qualité ici — c'est le job du câblage CP+Decoder (Patch 3).
|
||||||
|
//
|
||||||
|
// Usage : ./test_talker <gguf>
|
||||||
|
//
|
||||||
|
// Build : copie dans dist/jni/, link avec libkazeia_engine.so dépendances.
|
||||||
|
#include <cstdio>
|
||||||
|
#include <cstring>
|
||||||
|
#include <cstdlib>
|
||||||
|
#include <vector>
|
||||||
|
#include "llama.h"
|
||||||
|
#include "ggml-backend.h"
|
||||||
|
|
||||||
|
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;
|
||||||
|
}
|
||||||
|
|
||||||
|
int main(int argc, char** argv) {
|
||||||
|
if (argc < 2) { printf("usage: %s <gguf>\n", argv[0]); return 1; }
|
||||||
|
|
||||||
|
// dense path : HMX off (cf. kazeia_engine_jni)
|
||||||
|
setenv("GGML_HEXAGON_USE_HMX", "0", 1);
|
||||||
|
llama_backend_init();
|
||||||
|
|
||||||
|
auto mp = llama_model_default_params();
|
||||||
|
// 2e argument optionnel : "cpu" force le path CPU même si HTP dispo (utile pour itérer
|
||||||
|
// sans payer l'upload HTP de 1.66GB f32 à chaque smoke test).
|
||||||
|
const bool force_cpu = (argc >= 3 && !strcmp(argv[2], "cpu"));
|
||||||
|
ggml_backend_dev_t devs[2] = { force_cpu ? nullptr : find_htp(), nullptr };
|
||||||
|
if (devs[0]) { mp.n_gpu_layers = 99; mp.devices = devs; printf("HTP0 -> ngl=99\n"); }
|
||||||
|
else { mp.n_gpu_layers = 0; printf("CPU only%s\n", force_cpu ? " (force)" : " (pas de HTP)"); }
|
||||||
|
|
||||||
|
auto m = llama_model_load_from_file(argv[1], mp);
|
||||||
|
if (!m) { printf("model load FAILED\n"); return 1; }
|
||||||
|
printf("model loaded. n_embd=%d, n_vocab=%d\n",
|
||||||
|
llama_model_n_embd(m), llama_vocab_n_tokens(llama_model_get_vocab(m)));
|
||||||
|
|
||||||
|
auto cp = llama_context_default_params();
|
||||||
|
cp.n_ctx = 1024; cp.n_batch = 1024; cp.n_threads = 4;
|
||||||
|
cp.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED;
|
||||||
|
cp.embeddings = true; // équivaut à llama_set_embeddings(ctx, true)
|
||||||
|
|
||||||
|
auto ctx = llama_init_from_model(m, cp);
|
||||||
|
if (!ctx) { printf("ctx init 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));
|
||||||
|
|
||||||
|
// --- prefill : 4 positions d'embeds aléatoires (juste pour valider le pipeline,
|
||||||
|
// PAS pour produire du français sensé) ---
|
||||||
|
const int T = 4;
|
||||||
|
std::vector<float> embd(T * n_embd);
|
||||||
|
srand(42);
|
||||||
|
for (auto& x : embd) x = (rand() / (float)RAND_MAX - 0.5f) * 0.01f;
|
||||||
|
|
||||||
|
std::vector<llama_pos> pos(T);
|
||||||
|
std::vector<int32_t> nsd(T, 1);
|
||||||
|
std::vector<llama_seq_id> sid0(T, 0);
|
||||||
|
std::vector<llama_seq_id*> sids(T);
|
||||||
|
std::vector<int8_t> lg(T, 0);
|
||||||
|
for (int i = 0; i < T; ++i) { pos[i] = i; sids[i] = &sid0[i]; }
|
||||||
|
lg[T-1] = 1;
|
||||||
|
|
||||||
|
llama_batch b{};
|
||||||
|
b.n_tokens = T; b.token = nullptr; b.embd = embd.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 embeds FAILED\n"); return 1; }
|
||||||
|
printf("prefill OK (T=%d)\n", T);
|
||||||
|
|
||||||
|
const float* h_prefill = llama_get_embeddings_ith(ctx, -1);
|
||||||
|
if (!h_prefill) { printf("get_embeddings_ith FAILED\n"); return 1; }
|
||||||
|
float h_sum = 0, h_min = h_prefill[0], h_max = h_prefill[0];
|
||||||
|
for (int i = 0; i < n_embd; ++i) {
|
||||||
|
h_sum += h_prefill[i];
|
||||||
|
if (h_prefill[i] < h_min) h_min = h_prefill[i];
|
||||||
|
if (h_prefill[i] > h_max) h_max = h_prefill[i];
|
||||||
|
}
|
||||||
|
printf("hidden[prefill] : mean=%.6f min=%.6f max=%.6f\n", h_sum / n_embd, h_min, h_max);
|
||||||
|
|
||||||
|
const float* logits = llama_get_logits_ith(ctx, -1);
|
||||||
|
if (!logits) { printf("get_logits_ith FAILED\n"); return 1; }
|
||||||
|
int argmax = 0; float lmax = logits[0];
|
||||||
|
for (int i = 1; i < n_vocab; ++i) if (logits[i] > lmax) { lmax = logits[i]; argmax = i; }
|
||||||
|
printf("logits[prefill] : argmax=%d (val=%.3f)\n", argmax, lmax);
|
||||||
|
|
||||||
|
// --- decode step : un seul embed de plus ---
|
||||||
|
std::vector<float> embd1(n_embd);
|
||||||
|
for (auto& x : embd1) x = (rand() / (float)RAND_MAX - 0.5f) * 0.01f;
|
||||||
|
|
||||||
|
llama_pos p1 = T;
|
||||||
|
int32_t n1 = 1;
|
||||||
|
llama_seq_id s1 = 0; llama_seq_id* sp1 = &s1;
|
||||||
|
int8_t l1 = 1;
|
||||||
|
llama_batch b1{};
|
||||||
|
b1.n_tokens = 1; b1.token = nullptr; b1.embd = embd1.data();
|
||||||
|
b1.pos = &p1; b1.n_seq_id = &n1; b1.seq_id = &sp1; b1.logits = &l1;
|
||||||
|
|
||||||
|
if (llama_decode(ctx, b1) != 0) { printf("decode embeds FAILED\n"); return 1; }
|
||||||
|
const float* h_step = llama_get_embeddings_ith(ctx, -1);
|
||||||
|
const float* l_step = llama_get_logits_ith(ctx, -1);
|
||||||
|
int argmax_s = 0; float lmax_s = l_step[0];
|
||||||
|
for (int i = 1; i < n_vocab; ++i) if (l_step[i] > lmax_s) { lmax_s = l_step[i]; argmax_s = i; }
|
||||||
|
printf("step OK : argmax=%d (val=%.3f), hidden[0..2]=%.4f %.4f %.4f\n",
|
||||||
|
argmax_s, lmax_s, h_step[0], h_step[1], h_step[2]);
|
||||||
|
|
||||||
|
llama_free(ctx);
|
||||||
|
llama_model_free(m);
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
2
ql
2
ql
|
|
@ -1 +1 @@
|
||||||
Subproject commit 956f72ffe4fe2758e7c5236ac8ea84752a48f7e3
|
Subproject commit c1609a141d1c372b6f08a00fc272c56ce0fb5e6b
|
||||||
Loading…
Reference in New Issue