438 lines
21 KiB
C++
438 lines
21 KiB
C++
// 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 "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)
|
|
};
|
|
|
|
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 = ≶
|
|
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());
|
|
}
|
|
|
|
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);
|
|
}
|
|
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;
|
|
}
|