110 lines
5.1 KiB
C++
110 lines
5.1 KiB
C++
// kazeia_engine_jni.cpp — bridge LLM Kazeia-Engine (llama.cpp fork ql + Hexagon).
|
|
// OPTION C: prefill NPU/HTP (ngl99, SSM_CONV forcé CPU via OPFILTER) -> transfert KV -> decode CPU (ngl0, t4+fa).
|
|
// prefill ~100-180 t/s (vs ~14 CPU), decode CPU ~7-10 (KV f16). 2 instances du modèle (~4.7GB).
|
|
// SSM_CONV HTP cassé en prefill -> OPFILTER=SSM_CONV (le route sur ggml-cpu). cf HANDOFF.md / PERF.md.
|
|
// generate() est sans état entre appels (l'app passe l'historique complet dans le prompt). STT reste ORT-QAIRT.
|
|
#include <jni.h>
|
|
#include <cstdlib>
|
|
#include <string>
|
|
#include <vector>
|
|
#include <cstring>
|
|
#include "llama.h"
|
|
#include "ggml-backend.h"
|
|
|
|
struct KEngine {
|
|
llama_model* m_h; llama_context* c_h; // prefill HTP
|
|
llama_model* m_c; llama_context* c_c; // decode CPU
|
|
const llama_vocab* v; llama_sampler* s;
|
|
};
|
|
|
|
static llama_context* make_ctx(llama_model* m, int nctx, int nthreads) {
|
|
auto cp = llama_context_default_params();
|
|
cp.n_ctx = nctx; cp.n_batch = nctx; cp.n_threads = nthreads;
|
|
cp.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED; // t4+fa optimum decode
|
|
cp.type_k = GGML_TYPE_F16; cp.type_v = GGML_TYPE_F16; // KV f16 (q8_0 = -40% decode, mesuré)
|
|
return llama_init_from_model(m, cp);
|
|
}
|
|
|
|
extern "C" JNIEXPORT jlong JNICALL
|
|
Java_com_kazeia_llm_EngineJni_load(JNIEnv* e, jobject, jstring path, jint nctx) {
|
|
const char* p = e->GetStringUTFChars(path, 0);
|
|
// Doit être posé AVANT l'init du backend hexagon (lu au registre).
|
|
setenv("GGML_HEXAGON_GDN_PREFILL", "1", 1); // GDN prefill sur HTP
|
|
setenv("GGML_HEXAGON_OPFILTER", "SSM_CONV", 1); // conv1d HTP cassé en prefill -> sur CPU
|
|
llama_backend_init();
|
|
|
|
// device HTP0 explicite pour l'instance prefill
|
|
static ggml_backend_dev_t devs[2] = { nullptr, nullptr };
|
|
for (size_t i = 0; i < ggml_backend_dev_count(); ++i) {
|
|
auto d = ggml_backend_dev_get(i);
|
|
if (!strcmp(ggml_backend_dev_name(d), "HTP0")) devs[0] = d;
|
|
}
|
|
auto mp_h = llama_model_default_params(); mp_h.n_gpu_layers = 99; if (devs[0]) mp_h.devices = devs;
|
|
auto m_h = llama_model_load_from_file(p, mp_h);
|
|
auto mp_c = llama_model_default_params(); mp_c.n_gpu_layers = 0;
|
|
auto m_c = llama_model_load_from_file(p, mp_c);
|
|
e->ReleaseStringUTFChars(path, p);
|
|
if (!m_h || !m_c) return 0;
|
|
|
|
auto* k = new KEngine{ m_h, make_ctx(m_h, nctx, 8), // prefill t8
|
|
m_c, make_ctx(m_c, nctx, 4), // decode t4
|
|
llama_model_get_vocab(m_h), llama_sampler_init_greedy() };
|
|
return (jlong) k;
|
|
}
|
|
|
|
extern "C" JNIEXPORT jstring JNICALL
|
|
Java_com_kazeia_llm_EngineJni_generate(JNIEnv* e, jobject, jlong h, jstring sys, jstring usr, jint maxTok) {
|
|
auto* k = (KEngine*) h;
|
|
const char* sp = e->GetStringUTFChars(sys, 0); const char* up = e->GetStringUTFChars(usr, 0);
|
|
// ChatML Qwen3.5 + <think></think> vide = thinking OFF déterministe
|
|
std::string p = "<|im_start|>system\n"; p += sp; p += "<|im_end|>\n<|im_start|>user\n"; p += up;
|
|
p += "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n";
|
|
e->ReleaseStringUTFChars(sys, sp); e->ReleaseStringUTFChars(usr, up);
|
|
|
|
int n = -llama_tokenize(k->v, p.c_str(), p.size(), nullptr, 0, true, true);
|
|
std::vector<llama_token> t(n);
|
|
llama_tokenize(k->v, p.c_str(), p.size(), t.data(), n, true, true);
|
|
|
|
// PREFILL sur HTP (KV vidé -> historique complet re-prefillé à chaque appel)
|
|
llama_memory_clear(llama_get_memory(k->c_h), true);
|
|
llama_batch b = llama_batch_get_one(t.data(), n);
|
|
if (llama_decode(k->c_h, b) != 0) return e->NewStringUTF("");
|
|
|
|
// transfert état KV HTP -> CPU
|
|
size_t sz = llama_state_seq_get_size(k->c_h, 0);
|
|
std::vector<uint8_t> buf(sz);
|
|
llama_state_seq_get_data(k->c_h, buf.data(), sz, 0);
|
|
llama_memory_clear(llama_get_memory(k->c_c), true);
|
|
llama_state_seq_set_data(k->c_c, buf.data(), sz, 0);
|
|
|
|
// DECODE sur CPU (1er token depuis les logits prefill HTP, suite sur CPU)
|
|
llama_token id = llama_sampler_sample(k->s, k->c_h, -1);
|
|
std::string out; char zbuf[256]; int pos = n;
|
|
for (int i = 0; i < maxTok; ++i) {
|
|
if (llama_vocab_is_eog(k->v, id)) break;
|
|
int l = llama_token_to_piece(k->v, id, zbuf, sizeof zbuf, 0, true);
|
|
if (l > 0) out.append(zbuf, l);
|
|
llama_token tok = id; llama_pos pp = pos; int32_t ns = 1; llama_seq_id sd = 0, *spd = &sd; int8_t lg = 1;
|
|
llama_batch sb; memset(&sb, 0, sizeof sb);
|
|
sb.n_tokens = 1; sb.token = &tok; sb.pos = &pp; sb.n_seq_id = &ns; sb.seq_id = &spd; sb.logits = ≶
|
|
if (llama_decode(k->c_c, sb) != 0) break;
|
|
pos++; id = llama_sampler_sample(k->s, k->c_c, -1);
|
|
}
|
|
return e->NewStringUTF(out.c_str());
|
|
}
|
|
|
|
extern "C" JNIEXPORT void JNICALL
|
|
Java_com_kazeia_llm_EngineJni_reset(JNIEnv*, jobject, jlong h){
|
|
auto* k = (KEngine*) h;
|
|
llama_memory_clear(llama_get_memory(k->c_h), true);
|
|
llama_memory_clear(llama_get_memory(k->c_c), true);
|
|
}
|
|
extern "C" JNIEXPORT void JNICALL
|
|
Java_com_kazeia_llm_EngineJni_free(JNIEnv*, jobject, jlong h){
|
|
auto* k = (KEngine*) h;
|
|
llama_sampler_free(k->s);
|
|
llama_free(k->c_h); llama_free(k->c_c);
|
|
llama_model_free(k->m_h); llama_model_free(k->m_c);
|
|
delete k;
|
|
}
|