141 lines
7.7 KiB
Kotlin
141 lines
7.7 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)
|
||
|
||
// -- 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) }
|
||
}
|