kazeia/docs/RAG_EMBEDDINGS_ENGINE_SPEC.md

8.4 KiB
Raw Permalink Blame History

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) :

// ============================================================================
// 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é)

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) :

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.