Kazeia-engine/dist/jni/EngineLlmEngine.kt

141 lines
7.7 KiB
Kotlin
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.

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)
// -- Streaming + session (latence in-app). cb.onToken(piece) -> false pour arrêter.
// generateStream = mono-tour streamé. session* = system prefillé 1×, réutilisé (cache de préfixe KV).
external fun generateStream(h: Long, sys: String, usr: String, maxTok: Int, cb: TokenCallback)
external fun sessionStart(h: Long, sys: String) // prefille le system + checkpoint KV
external fun sessionAsk(h: Long, usr: String, maxTok: Int, cb: TokenCallback) // ne prefille que le user
external fun sessionReset(h: Long) // repart du checkpoint système
external fun getLastStats(h: Long): LongArray // [prefill_ms, decode_ms, n_tokens]
// -- 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") } }
}
/** Callback de streaming token. Renvoyer false pour interrompre la génération. */
fun interface TokenCallback { fun onToken(piece: String): Boolean }
/** Découpage temporel du dernier appel (ms). */
data class GenStats(val prefillMs: Long, val decodeMs: Long, val nTokens: Long) {
val decodeTokPerSec: Double get() = if (decodeMs > 0) nTokens * 1000.0 / decodeMs else 0.0
}
// 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)
// Mono-tour STREAMÉ : émet chaque morceau via onToken (false = stop). Permet de démarrer le TTS
// dès la 1ʳᵉ phrase. Récupère le découpage prefill/decode via lastStats() après l'appel.
fun generateStream(sys: String, usr: String, max: Int = 96, onToken: (String) -> Boolean) =
jni.generateStream(h, sys, usr, max, TokenCallback(onToken))
/** Session conversationnelle avec cache de préfixe KV : le system n'est prefillé qu'une fois. */
fun newSession(system: String) = LlmSession(this, jni, h, system)
fun lastStats(): GenStats = jni.getLastStats(h).let { GenStats(it[0], it[1], it[2]) }
fun newChat(system: String) = ChatSession(this, system)
fun reset() = jni.reset(h)
fun release() = jni.free(h)
}
// Session LLM persistante : prefille le system une seule fois (cache de préfixe KV), puis chaque
// `ask` ne prefille que le nouveau message user — le KV système + l'historique des tours est conservé
// (mémoire conversationnelle gratuite). `reset()` repart du system (nouvelle conversation).
// Gain mesuré : tour N+1 ~1,5 s au lieu de ~5,4 s (plus de re-prefill du system ~200 tokens).
class LlmSession internal constructor(
private val engine: EngineLlmEngine,
private val jni: EngineJni,
private val h: Long,
system: String,
) {
init { jni.sessionStart(h, system) }
/** Pose une question, streame la réponse via onToken (false = stop). Renvoie la réponse complète. */
fun ask(user: String, max: Int = 96, onToken: ((String) -> Boolean)? = null): String {
val sb = StringBuilder()
jni.sessionAsk(h, user, max, TokenCallback { piece ->
sb.append(piece)
onToken?.invoke(piece) ?: true
})
return sb.toString()
}
fun lastStats(): GenStats = engine.lastStats()
/** Vide l'historique des tours, garde le system prefillé. */
fun reset() = jni.sessionReset(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) }
}