kazeia/docs/RAG_EMBEDDINGS_ENGINE_SPEC.md

179 lines
8.4 KiB
Markdown
Raw 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.

# Spec — Ajouter le support embeddings (RAG) à kazeia-engine
> **Destinataire :** développeur de `kazeia-engine` (`/opt/Kazeia-engine`).
> **But :** exposer un entrypoint « texte → vecteur de phrase poolé » pour le RAG de Kazeia mobile.
> **Fichier cible :** `dist/jni/kazeia_engine_jni.cpp` (+ bindings Kotlin côté app).
> **Coût estimé :** ~4060 lignes C++, aucune nouvelle dépendance.
## Contexte
Le RAG de Kazeia mobile doit transformer un texte en **vecteur de phrase poolé**. La machinerie
existe déjà dans le fork (`llama_pooling_type`, champ `pooling_type` dans les context params,
`llama_get_embeddings_seq`), mais **aucun entrypoint JNI texte→vecteur** n'est exposé. Il faut
l'ajouter **sans toucher au chemin Speaker/TTS existant**.
## Principe (important)
- L'embedder est un **modèle dédié** (GGUF type **e5 / bge**, arch BERT), chargé comme **handle
séparé** — surtout **pas** le LLM Speaker (un décodeur causal est un mauvais embedder).
- **CPU pur, pas de HTP/HMX** (un embed coûte ~1050 ms ; le NPU est inutile et déstabilisant ici).
- **Ne pas réutiliser** `prefillEmbeds`/`decodeEmbed`/`llama_get_embeddings_ith` : c'est l'I/O float
du Talker TTS (hidden du *dernier token*), pas un embedding de phrase poolé.
---
## Tâche 1 — Nouvelle struct + `loadEmbedder`
À ajouter dans `kazeia_engine_jni.cpp` (section nouvelle, après le bloc LLM) :
```cpp
// ============================================================================
// EMBEDDINGS (RAG) : modèle dédié BERT-like (e5/bge), CPU, pooling de séquence.
// Indépendant du KEngine Speaker. Aucun HTP.
// ============================================================================
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) return 0;
auto cp = llama_context_default_params();
cp.n_threads = (nThreads > 0) ? nThreads : 4;
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) { 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)\n", p.c_str(), k->n_embd, pooling);
return (jlong) k;
}
```
## Tâche 2 — `embedText` (texte → vecteur L2-normalisé)
```cpp
extern "C" JNIEXPORT jfloatArray JNICALL
Java_com_kazeia_llm_EngineJni_embedText(JNIEnv* e, jobject, jlong h, jstring text) {
auto* k = (KEmbedder*) h;
const char* tc = e->GetStringUTFChars(text, 0);
std::string s(tc); e->ReleaseStringUTFChars(text, tc);
// 1) tokenize (add_special=true -> CLS/BOS selon le modèle)
int n = -llama_tokenize(k->v, s.c_str(), s.size(), nullptr, 0, true, true);
if (n <= 0) return nullptr;
if (n > 512) n = 512; // garde-fou ubatch
std::vector<llama_token> toks(n);
llama_tokenize(k->v, s.c_str(), s.size(), toks.data(), n, true, true);
// 2) batch : tous les tokens, seq 0, sortie activée (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);
llama_batch_free(b);
if (!emb) 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);
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->c) llama_free(k->c);
if (k->m) llama_model_free(k->m);
delete k;
}
```
> ⚠️ **À cross-checker contre `tools/llama-embedding` (ou `examples/embedding`) de TON fork.**
> Selon la version, un modèle encoder-only (BERT) peut nécessiter `llama_encode()` au lieu de
> `llama_decode()`, et le drapeau `logits` du batch peut différer. **Reproduis exactement ce que
> fait ton outil d'embedding de référence** — c'est la source de vérité.
## Tâche 3 — Bindings Kotlin
Dans `EngineJni` (`kazeia-android/.../llm/EngineLlmEngine.kt`) :
```kotlin
external fun loadEmbedder(ggufPath: String, nThreads: Int, pooling: Int): Long
external fun embedText(handle: Long, text: String): FloatArray?
external fun freeEmbedder(handle: Long)
```
---
## Points de vigilance (impératifs)
1. **Modèle = e5/bge dédié, FR-capable** (ex. `multilingual-e5-small`, 384-dim). Le **même** modèle
à l'ingestion et à la requête, sinon vecteurs incompatibles. **Versionne-le** (le n° de version
du modèle d'embedding fait partie du contrat de distribution).
2. **Vérifier d'abord que le fork charge l'arch BERT** (`bert` / `nomic-bert`) : le fork a été modifié
pour Qwen/TTS — **teste le GGUF d'embedding avec `llama-embedding` AVANT d'intégrer**. Si l'arch a
été retirée du fork, c'est LE point de blocage à traiter en priorité (≈ 5 min à vérifier).
3. **`n_ubatch ≥ nombre de tokens`** : le pooling MEAN/CLS impose la séquence entière dans un seul
ubatch. D'où `n_ctx = n_batch = n_ubatch = 512` et le cap à 512 tokens (chunks plus courts).
4. **`pooling = -1` (défaut modèle) recommandé** ; sinon e5→MEAN(1), bge→CLS(2).
5. **Préfixes e5 (`query:` / `passage:`) = côté appelant**, pas dans le moteur.
6. **CPU only** (`n_gpu_layers = 0`), pas de HTP/HMX.
7. **Déterminisme** : pas d'échantillonnage, l'embed est déterministe — indispensable pour
l'invariant « même modèle ingestion = requête ».
## Test d'acceptation (sur PC d'abord)
- `loadEmbedder` retourne un handle ≠ 0 sur le GGUF choisi ; `embedText` renvoie un `float[]` de
taille `n_embd` (384 pour e5-small).
- **Déterminisme** : même texte → vecteur **bit-identique** sur 2 appels.
- **Norme ≈ 1.0** (L2).
- **Cohérence sémantique** : `cos("j'ai du mal à dormir", "insomnie")` >
`cos("j'ai du mal à dormir", "recette de tarte")`.
- **Parité référence** : le vecteur ≈ celui de `llama-embedding` sur le même texte (à epsilon près).
## Build
Rebuild `libkazeia_engine.so` via le script habituel (`dist/build_kazeia_*.sh` / CMake).
**Aucune nouvelle dépendance** (tout est dans libllama). Livrer la `.so` + un GGUF d'embedding e5/bge.
---
## Ce qui reste côté app mobile (hors de cette spec, pour info)
Une fois `embedText` livré, l'intégration RAG dans `kazeia-android` est :
schéma SQLite (`chunk_text`, `source`, `embedding BLOB`) dans la base SQLCipher existante →
ingestion côté app admin (chunk sémantique + `embedText`) → au runtime, charger la matrice en RAM
→ brute-force cosinus (produit scalaire NEON) top-k=35 + **seuil de similarité** → injection d'un
bloc « CONTEXTE » budgété en tokens dans le system prompt. **Pas de vector-DB, pas d'ANN, pas
d'ONNX** (cf. analyse d'architecture). Le RAG ne porte jamais la gestion de crise.