Kazeia-engine/dist/jni/stt_engine.cpp

168 lines
6.2 KiB
C++

// Implementation stt_engine — voir stt_engine.h pour l'API publique.
//
// S1.3 (cette session) : squelette compilable, VAD complet, mel via kazeia_mel.
// stt_engine_load / transcribe / free : stubs renvoyant err=-99 jusqu'au port ORT
// QNN en S2.
#include "stt_engine.h"
#include "kazeia_mel.h"
#include <algorithm>
#include <cmath>
#include <cstdio>
#include <cstring>
#include <vector>
#include <string>
// ============================================================================
// VAD RMS énergie (port C++ de VadStage.kt)
// ============================================================================
//
// État interne : compteur de frames parole / silence consécutifs, sliding sur
// chunks PCM. L'appelant push des chunks de taille frame_size; chaque push fait
// avancer l'état. is_speech / is_end_of_speech consultent l'état.
struct SttVadState {
int sample_rate;
int frame_size; // 1600 = 100 ms @ 16 kHz
int rms_threshold; // 150 sur PCM16 brut (prod KazeiaService)
int min_speech_frames; // 3 -> 300 ms de parole pour déclencher
int silence_end_frames; // 8 -> 800 ms de silence pour finir
int speech_count = 0;
int silence_count = 0;
bool in_speech = false;
bool end_of_speech = false;
};
SttVadState * stt_vad_new(int sample_rate, int frame_size, int rms_threshold,
int min_speech_frames, int silence_end_frames) {
auto * s = new SttVadState();
s->sample_rate = sample_rate;
s->frame_size = frame_size;
s->rms_threshold = rms_threshold;
s->min_speech_frames = min_speech_frames;
s->silence_end_frames = silence_end_frames;
return s;
}
void stt_vad_reset(SttVadState * s) {
if (!s) return;
s->speech_count = 0;
s->silence_count = 0;
s->in_speech = false;
s->end_of_speech = false;
}
void stt_vad_free(SttVadState * s) { delete s; }
bool stt_vad_is_speech(const SttVadState * s) { return s && s->in_speech; }
bool stt_vad_is_end_of_speech(const SttVadState * s) { return s && s->end_of_speech; }
// Calcule RMS d'un buffer PCM16, renvoie l'énergie en unités PCM brutes
// (compatible seuil 150 prod).
static float rms_pcm16(const int16_t * pcm, int n) {
if (n <= 0) return 0.0f;
double acc = 0.0;
for (int i = 0; i < n; ++i) { double v = (double)pcm[i]; acc += v * v; }
return (float)std::sqrt(acc / (double)n);
}
bool stt_vad_push(SttVadState * s, const int16_t * pcm, int n_samples) {
if (!s || !pcm || n_samples <= 0) return false;
s->end_of_speech = false;
// Traite par chunks de frame_size. Si n_samples < frame_size, on accumulera
// mentalement mais ici on calcule simplement la RMS sur ce qu'on a (l'app
// appelante pousse des chunks pleins de 100 ms en pratique).
int offset = 0;
while (offset + s->frame_size <= n_samples) {
float r = rms_pcm16(pcm + offset, s->frame_size);
if (r >= (float)s->rms_threshold) {
s->speech_count += 1;
s->silence_count = 0;
if (!s->in_speech && s->speech_count >= s->min_speech_frames) {
s->in_speech = true;
}
} else {
s->silence_count += 1;
// Pas de reset speech_count tant qu'on est in_speech, c'est l'enchainement
// silence consécutifs qui clôt.
if (s->in_speech && s->silence_count >= s->silence_end_frames) {
s->in_speech = false;
s->end_of_speech = true;
s->speech_count = 0;
} else if (!s->in_speech) {
s->speech_count = 0; // reset compteur si on n'a pas atteint le seuil
}
}
offset += s->frame_size;
}
return true;
}
// ============================================================================
// SttEngine (stubs — port ORT QNN en S2)
// ============================================================================
struct SttEngine {
SttEngineLoadCfg cfg;
std::vector<float> mel_basis; // [80, 201] librosa whisper
// S2 :
// void * encoder_session; // Ort::Session*
// void * decoder_session;
// std::vector<int> vocab_ids;
// std::unordered_map<std::string,int> token_to_id;
bool loaded = false;
};
SttEngine * stt_engine_load(const SttEngineLoadCfg & cfg) {
if (!cfg.model_dir) {
fprintf(stderr, "stt_engine_load: model_dir requis\n"); return nullptr;
}
auto * eng = new SttEngine();
eng->cfg = cfg;
// Tente charger mel_filters.bin (préféré) ou mel_filters.json (sera porté en S2).
std::string mel_path = std::string(cfg.model_dir) + "/mel_filters.bin";
auto mel_cfg = kazeia_mel::config_whisper();
if (!kazeia_mel::load_mel_basis(mel_path.c_str(), mel_cfg, eng->mel_basis)) {
fprintf(stderr, "stt_engine_load: mel_filters.bin absent ou mauvaise taille à %s\n",
mel_path.c_str());
// Pas fatal : à S2 on chargera depuis mel_filters.json. Pour S1.3 on continue
// pour permettre les tests VAD.
}
eng->loaded = true;
return eng;
}
SttTranscribeResult stt_engine_transcribe(SttEngine * eng, const SttTranscribeCfg & cfg) {
SttTranscribeResult R{};
if (!eng || !eng->loaded || !cfg.pcm16 || cfg.n_samples <= 0) {
R.err = -1; return R;
}
// Étape 1 : mel via kazeia_mel (Whisper config). Implémentée en S1.3 pour exercer
// la lib partagée ; le encoder/decoder ORT viennent en S2.
std::vector<float> wav((size_t)cfg.n_samples);
for (int i = 0; i < cfg.n_samples; ++i) wav[i] = (float)cfg.pcm16[i] / 32768.0f;
if (eng->mel_basis.empty()) {
fprintf(stderr, "stt_engine_transcribe: mel_basis non chargé (mel_filters.bin manquant)\n");
R.err = -2; return R;
}
auto mel_cfg = kazeia_mel::config_whisper();
int T = 0;
auto mel = kazeia_mel::compute(wav, eng->mel_basis, mel_cfg, T);
if (mel.empty()) { R.err = -3; return R; }
R.mel_ms = 0; // TODO: timer
// Étape 2..N : encoder NPU + decoder loop + tokenizer BPE -> S2.
fprintf(stderr, "stt_engine_transcribe: stubs (port S2 requis). mel %d frames OK.\n", T);
R.err = -99; // not implemented yet
R.detected_language = cfg.language;
return R;
}
void stt_engine_free(SttEngine * eng) {
if (!eng) return;
delete eng;
}