Kazeia-engine/dist/jni/EngineLlmEngine.kt

88 lines
4.8 KiB
Kotlin

package com.kazeia.llm
// Remplace ExecuTorchLlmEngine. Option C: prefill HTP / decode CPU (géré côté natif).
// Mono-tour: generate(sys, usr). Multi-tour: ChatSession (gère l'historique + le template ChatML).
class EngineJni {
/**
* @param nThreads decode CPU threads. Défaut 0 -> 6 (sweet spot Snapdragon 8 Elite,
* mémoire r3). Recommandations : 6 pour Speaker 4B/9B in-app, 4 si cascade
* simultanée Speaker+Thinker pour laisser cœurs au Thinker.
*/
external fun load(ggufPath: String, nCtx: Int, nThreads: Int): Long
external fun generate(h: Long, sys: String, usr: String, maxTok: Int): String
external fun generateRaw(h: Long, prompt: String, maxTok: Int): String // prompt complet déjà formaté
external fun reset(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
// -- Embeddings (RAG) : modèle dédié BERT-like (e5/bge), handle SÉPARÉ du LLM/TTS, CPU pur.
// pooling: -1 = défaut du modèle (recommandé) ; 1 = MEAN (e5) ; 2 = CLS (bge).
// embedText renvoie un float[n_embd] L2-normalisé (cosinus = dot product), ou null si échec.
// Préfixes e5 ("query:" / "passage:") = côté appelant, PAS ici.
external fun loadEmbedder(ggufPath: String, nThreads: Int, pooling: Int): Long
external fun embedText(handle: Long, text: String): FloatArray?
external fun freeEmbedder(handle: Long)
companion object { init { System.loadLibrary("kazeia_engine") } }
}
// Wrapper RAG-friendly. Un seul embedder par process (le modèle est petit ~120 MB).
// Le MÊME GGUF doit servir à l'ingestion ET à la requête (vecteurs incompatibles sinon) :
// versionne le modèle dans le contrat de distribution.
class EmbedderEngine(model: String, nThreads: Int = 4, pooling: Int = -1) {
private val jni = EngineJni()
private val h = jni.loadEmbedder(model, nThreads, pooling)
init { require(h != 0L) { "Kazeia-Engine: échec du chargement de l'embedder ($model)" } }
/** Texte -> vecteur L2-normalisé (float[n_embd]). Déterministe. */
fun embed(text: String): FloatArray =
jni.embedText(h, text) ?: error("embedText a échoué pour: \"${text.take(40)}\"")
fun release() = jni.freeEmbedder(h)
}
class EngineLlmEngine(model: String, ctx: Int = 4096, nThreads: Int = 6) {
private val jni = EngineJni()
private val h = jni.load(model, ctx, nThreads)
init { require(h != 0L) { "Kazeia-Engine: échec du chargement du modèle ($model)" } }
// Mono-tour : system + un message user.
fun generate(sys: String, usr: String, max: Int = 96) = jni.generate(h, sys, usr, max)
// Multi-tour : prompt complet pré-formaté (voir ChatSession).
fun generateRaw(prompt: String, max: Int = 96) = jni.generateRaw(h, prompt, max)
fun newChat(system: String) = ChatSession(this, system)
fun reset() = jni.reset(h)
fun release() = jni.free(h)
}
// Conversation multi-tour : accumule l'historique et construit le ChatML Qwen3.5 + thinking-off.
// Structure identique à celle validée (jni/dual_ctx_mt.cpp) : mémoire conversationnelle correcte.
class ChatSession(private val engine: EngineLlmEngine, private val system: String) {
private val turns = StringBuilder() // "<|im_start|>role\ntext<|im_end|>\n" accumulés
fun ask(user: String, max: Int = 96): String {
val prompt = buildString {
append("<|im_start|>system\n").append(system).append("<|im_end|>\n")
append(turns)
append("<|im_start|>user\n").append(user).append("<|im_end|>\n")
append("<|im_start|>assistant\n<think>\n\n</think>\n\n") // thinking OFF déterministe
}
val resp = engine.generateRaw(prompt, max)
turns.append("<|im_start|>user\n").append(user).append("<|im_end|>\n")
.append("<|im_start|>assistant\n").append(resp).append("<|im_end|>\n")
return resp
}
fun clear() { turns.setLength(0) }
}