Kazeia-engine/dist/jni/kazeia_engine_jni.cpp

626 lines
30 KiB
C++
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// kazeia_engine_jni.cpp — bridge LLM Kazeia-Engine (llama.cpp fork ql + Hexagon).
// GÉNÉRIQUE : fait tourner N'IMPORTE QUEL LLM. Au load, détecte la STRUCTURE du modèle
// (llama_model_is_hybrid/is_recurrent) et route :
// - HYBRIDE GDN (qwen3.5 / qwen3next) -> OPTION C : prefill HTP -> transfert KV -> decode CPU
// (le decode GDN sur HTP est lent ; le split le garde sur CPU).
// - DENSE (qwen3, llama, ...) -> HTP contexte-unique : prefill + decode sur HTP
// (decode dense HTP ~= CPU, pas de pénalité GDN ; prefill ~98 vs 14 CPU ; pas de dual-load => pas de crash 0x2e).
// - pas de device HTP -> CPU pur (repli universel).
// API : load -> generate(sys,usr) | generateRaw(prompt) -> reset/free. Sans état entre appels.
#include <jni.h>
#include <cstdio>
#include <cstdlib>
#include <string>
#include <vector>
#include <cstring>
#include <cmath>
#include <chrono>
#include "llama.h"
#include "ggml-backend.h"
#include "gguf.h"
// Lit general.architecture dans l'en-tête GGUF SANS init backend (pour décider USE_HMX avant).
static std::string read_arch(const char* path) {
struct gguf_init_params gp = { /*no_alloc=*/true, /*ctx=*/nullptr };
gguf_context* g = gguf_init_from_file(path, gp);
if (!g) return "";
int64_t kid = gguf_find_key(g, "general.architecture");
std::string arch = (kid >= 0) ? gguf_get_val_str(g, kid) : "";
gguf_free(g);
return arch;
}
struct KEngine {
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)
const llama_vocab* v; llama_sampler* s;
int pos_embd; // position courante en mode talker embeds-only (TTS)
// --- session conversationnelle + cache de prefixe KV (API sessionStart/Ask/Reset) ---
int n_sys = -1; // checkpoint = fin du system prompt (-1 = pas de session)
int n_past = 0; // position KV courante (system + tours accumulés)
std::vector<llama_token> sys_toks; // tokens du system (pour re-prefill au reset des archis récurrentes)
// --- stats du dernier appel (prefill/decode), exposées via getLastStats ---
double last_prefill_ms = 0, last_decode_ms = 0;
int last_n_gen = 0;
};
static inline double now_ms() {
return std::chrono::duration<double, std::milli>(std::chrono::steady_clock::now().time_since_epoch()).count();
}
// Prefille `toks` (n) au pos de depart `start` (positions explicites), logits sur le dernier seulement.
// Renvoie le pos final (start+n), ou -1 si echec.
static int kengine_prefill_at(llama_context* ctx, const llama_token* toks, int n, int start) {
if (n <= 0) return start;
llama_seq_id zero = 0;
std::vector<llama_pos> pos(n);
std::vector<int32_t> nsid(n, 1);
std::vector<llama_seq_id*> sidp(n, &zero);
std::vector<int8_t> lg(n, 0);
for (int i = 0; i < n; ++i) pos[i] = start + i;
lg[n - 1] = 1;
llama_batch b{};
b.n_tokens = n; b.token = const_cast<llama_token*>(toks); b.pos = pos.data();
b.n_seq_id = nsid.data(); b.seq_id = sidp.data(); b.logits = lg.data();
if (llama_decode(ctx, b) != 0) return -1;
return start + n;
}
// Decode en streaming depuis les logits courants de `ctx`. Emet chaque piece via le callback Kotlin
// `onTok(String):Boolean` (false = stop). Accumule le KV. Renvoie le pos final ; *n_gen = tokens generes.
static int kengine_decode_stream(llama_context* ctx, KEngine* k, int start, int maxTok,
JNIEnv* e, jobject cb, jmethodID onTok, int* n_gen) {
llama_token id = llama_sampler_sample(k->s, ctx, -1);
int pos = start; char zbuf[256]; *n_gen = 0;
llama_seq_id zero = 0;
for (int i = 0; i < maxTok; ++i) {
if (llama_vocab_is_eog(k->v, id)) break;
int l = llama_token_to_piece(k->v, id, zbuf, sizeof zbuf, 0, true);
if (l > 0 && e && cb && onTok) {
jstring js = e->NewStringUTF(std::string(zbuf, l).c_str());
jboolean cont = e->CallBooleanMethod(cb, onTok, js);
e->DeleteLocalRef(js);
if (!cont) break;
}
(*n_gen)++;
llama_token tok = id; llama_pos pp = pos; int32_t ns = 1; llama_seq_id* sp = &zero; int8_t lgf = 1;
llama_batch sb{};
sb.n_tokens = 1; sb.token = &tok; sb.pos = &pp; sb.n_seq_id = &ns; sb.seq_id = &sp; sb.logits = &lgf;
if (llama_decode(ctx, sb) != 0) break;
pos++; id = llama_sampler_sample(k->s, ctx, -1);
}
return pos;
}
static std::vector<llama_token> kengine_tokenize(KEngine* k, const std::string& p, bool add_bos) {
int n = -llama_tokenize(k->v, p.c_str(), p.size(), nullptr, 0, add_bos, true);
std::vector<llama_token> t(n);
llama_tokenize(k->v, p.c_str(), p.size(), t.data(), n, add_bos, true);
return t;
}
static llama_context* make_ctx(llama_model* m, int nctx, int nthreads, enum llama_flash_attn_type fa) {
auto cp = llama_context_default_params();
cp.n_ctx = nctx; cp.n_batch = 2048; cp.n_threads = nthreads;
cp.flash_attn_type = fa; // ENABLED (decode CPU) / DISABLED (prefill HTP dense)
cp.type_k = GGML_TYPE_F16; cp.type_v = GGML_TYPE_F16; // KV f16 (q8_0 = -40% decode, mesuré)
return llama_init_from_model(m, cp);
}
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;
}
// Cœur : prefill (HTP si dispo) -> [transfert KV si dual-ctx] -> decode (CPU en option C, sinon HTP).
static std::string kengine_run(KEngine* k, const std::string& p, int maxTok) {
int n = -llama_tokenize(k->v, p.c_str(), p.size(), nullptr, 0, true, true);
std::vector<llama_token> t(n);
llama_tokenize(k->v, p.c_str(), p.size(), t.data(), n, true, true);
llama_context* pf = k->c_h ? k->c_h : k->c_c; // contexte de prefill
llama_context* dec = k->c_c ? k->c_c : k->c_h; // contexte de decode
llama_memory_clear(llama_get_memory(pf), true);
llama_batch b = llama_batch_get_one(t.data(), n);
if (llama_decode(pf, b) != 0) return std::string();
if (k->c_h && k->c_c) { // option C : transfert KV HTP -> CPU
size_t sz = llama_state_seq_get_size(k->c_h, 0);
std::vector<uint8_t> buf(sz);
llama_state_seq_get_data(k->c_h, buf.data(), sz, 0);
llama_memory_clear(llama_get_memory(k->c_c), true);
llama_state_seq_set_data(k->c_c, buf.data(), sz, 0);
}
llama_token id = llama_sampler_sample(k->s, pf, -1); // 1er token depuis les logits de prefill
std::string out; char zbuf[256]; int pos = n;
for (int i = 0; i < maxTok; ++i) {
if (llama_vocab_is_eog(k->v, id)) break;
int l = llama_token_to_piece(k->v, id, zbuf, sizeof zbuf, 0, true);
if (l > 0) out.append(zbuf, l);
llama_token tok = id; llama_pos pp = pos; int32_t ns = 1; llama_seq_id sd = 0, *spd = &sd; int8_t lg = 1;
llama_batch sb; memset(&sb, 0, sizeof sb);
sb.n_tokens = 1; sb.token = &tok; sb.pos = &pp; sb.n_seq_id = &ns; sb.seq_id = &spd; sb.logits = &lg;
if (llama_decode(dec, sb) != 0) break;
pos++; id = llama_sampler_sample(k->s, dec, -1);
}
return out;
}
// LOAD : nThreads contrôle le n_threads du decode CPU.
// Sur Snapdragon 8 Elite (SM8750), l'optimum mesuré sur r3 est t=6 pour Qwen3.5-4B
// (decode 9.8 tok/s) et t=6 pour 9B (5.2 tok/s). t=8 réduit à cause de contention
// thermique + interaction avec le scheduler Android (in-app souvent dégradé).
// Param exposé pour permettre le tuning par voix (Speaker vs Thinker).
extern "C" JNIEXPORT jlong JNICALL
Java_com_kazeia_llm_EngineJni_load(JNIEnv* e, jobject, jstring path, jint nctx, jint nThreads) {
const char* path_c = e->GetStringUTFChars(path, 0);
std::string p(path_c);
e->ReleaseStringUTFChars(path, path_c);
const int N_DECODE = (nThreads > 0) ? nThreads : 6; // défaut sweet spot r3
const int N_PREFILL = 8; // prefill HTP : 8 OK (HTP gère)
// Détection d'archi AVANT init backend (USE_HMX est latché à l'init).
std::string arch = read_arch(p.c_str());
const bool hybrid = (arch.find("qwen3next") != std::string::npos) || (arch == "qwen35");
// Le matmul HMX faute (0x2e) sur les activations réelles du DENSE (massive activations > plage fp16).
// -> HMX OFF pour le dense (matmul HVX, stable, ~3x CPU). L'hybride (3.5) garde HMX (prefill rapide, OK).
if (!hybrid) setenv("GGML_HEXAGON_USE_HMX", "0", 1);
setenv("GGML_HEXAGON_GDN_PREFILL", "1", 1); // sans effet sur les modèles sans GDN
llama_backend_init();
llama_model* m_h = nullptr; llama_context* c_h = nullptr;
llama_model* m_c = nullptr; llama_context* c_c = nullptr;
ggml_backend_dev_t htp = find_htp();
if (htp) {
ggml_backend_dev_t devs[2] = { htp, nullptr };
auto mp_h = llama_model_default_params(); mp_h.n_gpu_layers = 99; mp_h.devices = devs;
m_h = llama_model_load_from_file(p.c_str(), mp_h);
}
if (m_h && hybrid) {
// qwen3.5-like (GDN) -> OPTION C : prefill HTP (HMX) + instance CPU pour le decode (decode GDN HTP = lent).
c_h = make_ctx(m_h, nctx, N_PREFILL, LLAMA_FLASH_ATTN_TYPE_ENABLED);
auto mp_c = llama_model_default_params(); mp_c.n_gpu_layers = 0;
m_c = llama_model_load_from_file(p.c_str(), mp_c);
if (m_c) c_c = make_ctx(m_c, nctx, N_DECODE, LLAMA_FLASH_ATTN_TYPE_ENABLED);
fprintf(stderr, "kazeia-engine: HYBRIDE GDN (%s) -> option C, prefill HTP+HMX (t=%d) / decode CPU (t=%d)\n",
arch.c_str(), N_PREFILL, N_DECODE);
} else if (m_h) {
// dense -> HTP contexte-unique, matmul HVX (HMX off) : prefill+decode HTP, stable, ~3x CPU.
c_h = make_ctx(m_h, nctx, N_DECODE, LLAMA_FLASH_ATTN_TYPE_ENABLED);
fprintf(stderr, "kazeia-engine: DENSE (%s) -> HTP contexte-unique, matmul HVX (HMX off) (t=%d)\n",
arch.c_str(), N_DECODE);
} else {
// pas de HTP (ou échec du load HTP) -> CPU pur universel.
auto mp_c = llama_model_default_params(); mp_c.n_gpu_layers = 0;
m_c = llama_model_load_from_file(p.c_str(), mp_c);
if (!m_c) return 0;
c_c = make_ctx(m_c, nctx, N_DECODE, LLAMA_FLASH_ATTN_TYPE_ENABLED);
fprintf(stderr, "kazeia-engine: CPU pur (%s, pas de HTP) (t=%d)\n", arch.c_str(), N_DECODE);
}
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(),
/*pos_embd=*/0 };
return (jlong) k;
}
// Mono-tour : construit le ChatML (system + 1 tour user) + thinking-off, puis infère.
extern "C" JNIEXPORT jstring JNICALL
Java_com_kazeia_llm_EngineJni_generate(JNIEnv* e, jobject, jlong h, jstring sys, jstring usr, jint maxTok) {
auto* k = (KEngine*) h;
const char* sp = e->GetStringUTFChars(sys, 0); const char* up = e->GetStringUTFChars(usr, 0);
std::string p = "<|im_start|>system\n"; p += sp; p += "<|im_end|>\n<|im_start|>user\n"; p += up;
p += "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n";
e->ReleaseStringUTFChars(sys, sp); e->ReleaseStringUTFChars(usr, up);
std::string out = kengine_run(k, p, maxTok);
return e->NewStringUTF(out.c_str());
}
// Multi-tour : l'app/Kotlin fournit le prompt complet déjà formaté (voir ChatSession).
extern "C" JNIEXPORT jstring JNICALL
Java_com_kazeia_llm_EngineJni_generateRaw(JNIEnv* e, jobject, jlong h, jstring prompt, jint maxTok) {
auto* k = (KEngine*) h;
const char* pp = e->GetStringUTFChars(prompt, 0);
std::string out = kengine_run(k, std::string(pp), maxTok);
e->ReleaseStringUTFChars(prompt, pp);
return e->NewStringUTF(out.c_str());
}
// ============================================================================
// Streaming + session + cache de prefixe KV (latence in-app).
// Contexte unique = decode (c_c sinon c_h) : adapte au build CPU-only in-app.
// generateStream = mono-tour streamé. sessionStart/Ask/Reset = system prefillé 1×, réutilisé.
// ============================================================================
static jmethodID get_ontoken(JNIEnv* e, jobject cb) {
if (!cb) return nullptr;
jclass cls = e->GetObjectClass(cb);
return e->GetMethodID(cls, "onToken", "(Ljava/lang/String;)Z"); // Boolean : false = stop
}
// generateStream(h, sys, usr, maxTok, callback) : mono-tour, KV remis à zéro, réponse streamée token par token.
extern "C" JNIEXPORT void JNICALL
Java_com_kazeia_llm_EngineJni_generateStream(JNIEnv* e, jobject, jlong h, jstring sys, jstring usr, jint maxTok, jobject cb) {
auto* k = (KEngine*) h;
llama_context* ctx = k->c_c ? k->c_c : k->c_h;
const char* sp = e->GetStringUTFChars(sys, 0); const char* up = e->GetStringUTFChars(usr, 0);
std::string p = "<|im_start|>system\n"; p += sp; p += "<|im_end|>\n<|im_start|>user\n"; p += up;
p += "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n";
e->ReleaseStringUTFChars(sys, sp); e->ReleaseStringUTFChars(usr, up);
jmethodID onTok = get_ontoken(e, cb);
auto toks = kengine_tokenize(k, p, true);
llama_memory_clear(llama_get_memory(ctx), true);
k->n_sys = -1; // generateStream est sans session
double t0 = now_ms();
int pos = kengine_prefill_at(ctx, toks.data(), (int)toks.size(), 0);
k->last_prefill_ms = now_ms() - t0;
if (pos < 0) { k->last_decode_ms = 0; k->last_n_gen = 0; return; }
double t1 = now_ms();
int ng = 0;
kengine_decode_stream(ctx, k, pos, maxTok, e, cb, onTok, &ng);
k->last_decode_ms = now_ms() - t1;
k->last_n_gen = ng;
}
// sessionStart(h, sys) : prefille le system prompt UNE fois et pose le checkpoint KV. Pas de génération.
extern "C" JNIEXPORT void JNICALL
Java_com_kazeia_llm_EngineJni_sessionStart(JNIEnv* e, jobject, jlong h, jstring sys) {
auto* k = (KEngine*) h;
llama_context* ctx = k->c_c ? k->c_c : k->c_h;
const char* sp = e->GetStringUTFChars(sys, 0);
std::string p = "<|im_start|>system\n"; p += sp; p += "<|im_end|>\n";
e->ReleaseStringUTFChars(sys, sp);
auto toks = kengine_tokenize(k, p, true);
llama_memory_clear(llama_get_memory(ctx), true);
if (k->s) llama_sampler_reset(k->s);
double t0 = now_ms();
int pos = kengine_prefill_at(ctx, toks.data(), (int)toks.size(), 0);
k->last_prefill_ms = now_ms() - t0; k->last_decode_ms = 0; k->last_n_gen = 0;
k->n_sys = (pos < 0) ? -1 : pos;
k->n_past = k->n_sys;
k->sys_toks = std::move(toks); // mémorisés pour sessionReset (archis récurrentes)
}
// sessionAsk(h, usr, maxTok, callback) : ne prefille QUE le nouveau tour user (KV système+historique conservé).
// Le KV accumule (mémoire conversationnelle). Appeler sessionReset pour repartir du system.
extern "C" JNIEXPORT void JNICALL
Java_com_kazeia_llm_EngineJni_sessionAsk(JNIEnv* e, jobject, jlong h, jstring usr, jint maxTok, jobject cb) {
auto* k = (KEngine*) h;
if (k->n_sys < 0) { k->last_prefill_ms = 0; k->last_decode_ms = 0; k->last_n_gen = 0; return; }
llama_context* ctx = k->c_c ? k->c_c : k->c_h;
const char* up = e->GetStringUTFChars(usr, 0);
std::string p;
if (k->n_past > k->n_sys) p = "<|im_end|>\n"; // ferme le tour assistant précédent (multi-tour)
p += "<|im_start|>user\n"; p += up; p += "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n";
e->ReleaseStringUTFChars(usr, up);
jmethodID onTok = get_ontoken(e, cb);
auto toks = kengine_tokenize(k, p, false); // continuation : pas de BOS
double t0 = now_ms();
int pos = kengine_prefill_at(ctx, toks.data(), (int)toks.size(), k->n_past); // ne prefille que le user
k->last_prefill_ms = now_ms() - t0;
if (pos < 0) { k->last_decode_ms = 0; k->last_n_gen = 0; return; }
double t1 = now_ms();
int ng = 0;
int end = kengine_decode_stream(ctx, k, pos, maxTok, e, cb, onTok, &ng);
k->last_decode_ms = now_ms() - t1; k->last_n_gen = ng;
k->n_past = end; // l'historique (réponse incluse) reste dans le KV
}
// sessionReset(h) : vide l'historique des tours, repart du checkpoint système.
// Sur archi récurrente/hybride (Qwen3.5 DeltaNet : état SSM), le rewind PARTIEL
// (seq_rm d'un suffixe) est impossible -> llama_memory_seq_rm renvoie false et ne
// fait rien : le KV reste incohérent et le tour suivant sort vide. On teste donc
// le retour : si le rewind échoue, reset COMPLET (clear + re-prefill du system).
extern "C" JNIEXPORT void JNICALL
Java_com_kazeia_llm_EngineJni_sessionReset(JNIEnv*, jobject, jlong h) {
auto* k = (KEngine*) h;
if (k->n_sys < 0) return;
llama_context* ctx = k->c_c ? k->c_c : k->c_h;
llama_memory_t mem = llama_get_memory(ctx);
if (llama_memory_seq_rm(mem, 0, k->n_sys, -1)) {
// rewind partiel OK (archi attention pure) : cheap, on garde le system en KV.
if (k->s) llama_sampler_reset(k->s);
k->n_past = k->n_sys;
} else {
// rewind impossible (récurrent/hybride) : reset complet + re-prefill system.
llama_memory_clear(mem, true);
if (k->s) llama_sampler_reset(k->s);
int pos = (k->sys_toks.empty())
? -1
: kengine_prefill_at(ctx, k->sys_toks.data(), (int)k->sys_toks.size(), 0);
k->n_sys = (pos < 0) ? -1 : pos;
k->n_past = (pos < 0) ? 0 : pos;
}
}
// getLastStats(h) : [prefill_ms, decode_ms, n_tokens_generes] du dernier appel.
extern "C" JNIEXPORT jlongArray JNICALL
Java_com_kazeia_llm_EngineJni_getLastStats(JNIEnv* e, jobject, jlong h) {
auto* k = (KEngine*) h;
jlong vals[3] = { (jlong)(k->last_prefill_ms + 0.5), (jlong)(k->last_decode_ms + 0.5), (jlong)k->last_n_gen };
jlongArray arr = e->NewLongArray(3);
e->SetLongArrayRegion(arr, 0, 3, vals);
return arr;
}
extern "C" JNIEXPORT void JNICALL
Java_com_kazeia_llm_EngineJni_reset(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->n_sys = -1; k->n_past = 0;
}
extern "C" JNIEXPORT void JNICALL
Java_com_kazeia_llm_EngineJni_free(JNIEnv*, jobject, jlong h){
auto* k = (KEngine*) h;
llama_sampler_free(k->s);
if (k->c_h) llama_free(k->c_h);
if (k->c_c) llama_free(k->c_c);
if (k->m_h) llama_model_free(k->m_h);
if (k->m_c) llama_model_free(k->m_c);
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.
//
// ⚠ M-RoPE (talker Qwen3-TTS) : llm_graph_input_pos::set_input attend des positions
// SOIT comme tokens et fait la conversion 1D→4D, SOIT directement T*4 entiers en
// embeds-mode (la branche else copie raw). Si on passe T positions seulement, ggml
// lit out-of-bounds. -> ici on alloue T*4 et duplique pos sur les 3 premiers axes
// (4ème = 0) comme le framework le fait pour les tokens.
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);
auto rt = llama_model_rope_type(k->m_h ? k->m_h : k->m_c);
const int npe = (rt == LLAMA_ROPE_TYPE_MROPE || rt == LLAMA_ROPE_TYPE_IMROPE) ? 4 : 1;
std::vector<llama_pos> pos (T * npe, 0);
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) {
const llama_pos p = k->pos_embd + i;
if (npe == 4) {
pos[ i] = p; // axis t
pos[ T + i] = p; // axis h
pos[2 * T + i] = p; // axis w
pos[3 * T + i] = 0; // axis e (zéro pour text-only)
} else {
pos[i] = p;
}
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;
}
// ============================================================================
// EMBEDDINGS (RAG) : modèle dédié BERT-like (e5/bge/gemma-embedding), CPU pur,
// pooling de séquence. Handle SÉPARÉ du KEngine Speaker/TTS — aucun HTP/HMX.
//
// ⚠ Bloqueur SELinux connu (cf REBUILD_CPU_ONLY.md / project_engine_selinux_fastrpc_blocker) :
// libllama linke libggml-hexagon qui ouvre /dev/fastrpc-cdsp à l'init backend ->
// SIGABRT dans un APK untrusted_app. L'embedder DOIT rouler sur le libllama
// CPU-only (GGML_HEXAGON=OFF), comme LLM/TTS in-app. n_gpu_layers=0 ne suffit pas
// à éviter le crash : c'est l'ÉNUMÉRATION des devices au backend_init qui plante.
//
// Chemin reproduit à l'identique de examples/embedding/embedding.cpp (source de
// vérité du fork) : llama_decode (PAS llama_encode — BERT encoder-only y est routé
// via decode), pooling != NONE -> llama_get_embeddings_seq, normalisation L2.
// ============================================================================
struct KEmbedder {
llama_model* m;
llama_context* c;
const llama_vocab* v;
int n_embd;
};
// pooling: -1 = défaut du modèle (recommandé) ; 1 = MEAN (e5) ; 2 = CLS (bge)
extern "C" JNIEXPORT jlong JNICALL
Java_com_kazeia_llm_EngineJni_loadEmbedder(JNIEnv* e, jobject, jstring path, jint nThreads, jint pooling) {
const char* pc = e->GetStringUTFChars(path, 0);
std::string p(pc); e->ReleaseStringUTFChars(path, pc);
llama_backend_init(); // idempotent
auto mp = llama_model_default_params(); mp.n_gpu_layers = 0; // CPU only
llama_model* m = llama_model_load_from_file(p.c_str(), mp);
if (!m) { fprintf(stderr, "kazeia-engine: EMBEDDER load FAIL (%s)\n", p.c_str()); return 0; }
// Encoder-decoder (T5) non supporté pour l'embedding (cf embedding.cpp).
// BERT pur (encoder-only) est routé via llama_decode dans ce fork -> OK.
if (llama_model_has_encoder(m) && llama_model_has_decoder(m)) {
fprintf(stderr, "kazeia-engine: EMBEDDER refuse un modèle encoder-decoder (%s)\n", p.c_str());
llama_model_free(m); return 0;
}
auto cp = llama_context_default_params();
cp.n_threads = (nThreads > 0) ? nThreads : 4;
cp.n_threads_batch = cp.n_threads;
cp.embeddings = true; // <-- clé
cp.pooling_type = (pooling == 1) ? LLAMA_POOLING_TYPE_MEAN
: (pooling == 2) ? LLAMA_POOLING_TYPE_CLS
: LLAMA_POOLING_TYPE_UNSPECIFIED; // défaut = métadonnées du modèle
// Pooling MEAN/CLS exige que TOUTE la séquence tienne dans un seul ubatch :
cp.n_ctx = 512; cp.n_batch = 512; cp.n_ubatch = 512; // chunks <= 512 tokens
// (ne PAS réutiliser make_ctx : pas de flash_attn ni de KV f16 ici)
llama_context* c = llama_init_from_model(m, cp);
if (!c) { fprintf(stderr, "kazeia-engine: EMBEDDER ctx FAIL (%s)\n", p.c_str()); llama_model_free(m); return 0; }
auto* k = new KEmbedder{ m, c, llama_model_get_vocab(m), llama_model_n_embd(m) };
fprintf(stderr, "kazeia-engine: EMBEDDER chargé (%s, n_embd=%d, pooling=%d, threads=%d)\n",
p.c_str(), k->n_embd, pooling, cp.n_threads);
return (jlong) k;
}
// texte -> float[n_embd] L2-normalisé. nullptr en cas d'échec.
extern "C" JNIEXPORT jfloatArray JNICALL
Java_com_kazeia_llm_EngineJni_embedText(JNIEnv* e, jobject, jlong h, jstring text) {
auto* k = (KEmbedder*) h;
if (!k) return nullptr;
const char* tc = e->GetStringUTFChars(text, 0);
std::string s(tc ? tc : ""); e->ReleaseStringUTFChars(text, tc);
// 1) tokenize (add_special=true -> CLS/BOS + SEP/EOS selon le modèle)
int n = -llama_tokenize(k->v, s.c_str(), (int)s.size(), nullptr, 0, true, true);
if (n <= 0) return nullptr;
if (n > 512) n = 512; // garde-fou ubatch (chunks plus courts en amont)
std::vector<llama_token> toks(n);
if (llama_tokenize(k->v, s.c_str(), (int)s.size(), toks.data(), n, true, true) < 0) return nullptr;
// 2) batch : tous les tokens, seq 0, sortie activée sur chaque token (requis pour le pooling)
llama_memory_clear(llama_get_memory(k->c), true);
llama_batch b = llama_batch_init(n, 0, 1);
for (int i = 0; i < n; ++i) {
b.token[i] = toks[i]; b.pos[i] = i;
b.n_seq_id[i] = 1; b.seq_id[i][0] = 0; b.logits[i] = 1;
}
b.n_tokens = n;
if (llama_decode(k->c, b) != 0) { llama_batch_free(b); return nullptr; }
// 3) embedding poolé de la séquence 0
const float* emb = llama_get_embeddings_seq(k->c, 0);
if (!emb) { llama_batch_free(b); return nullptr; }
// 4) normalisation L2 (cosinus = produit scalaire côté Android)
std::vector<float> out(k->n_embd);
double norm = 0.0; for (int i = 0; i < k->n_embd; ++i) norm += (double)emb[i] * emb[i];
norm = norm > 0 ? 1.0 / std::sqrt(norm) : 0.0;
for (int i = 0; i < k->n_embd; ++i) out[i] = (float)(emb[i] * norm);
llama_batch_free(b);
jfloatArray arr = e->NewFloatArray(k->n_embd);
e->SetFloatArrayRegion(arr, 0, k->n_embd, out.data());
return arr;
}
extern "C" JNIEXPORT void JNICALL
Java_com_kazeia_llm_EngineJni_freeEmbedder(JNIEnv*, jobject, jlong h) {
auto* k = (KEmbedder*) h;
if (!k) return;
if (k->c) llama_free(k->c);
if (k->m) llama_model_free(k->m);
delete k;
}