168 lines
6.2 KiB
C++
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;
|
|
}
|