Kazeia-engine/dist/jni/stt_cli.cpp

100 lines
3.6 KiB
C++

// kazeia_stt_cli : binaire S2 standalone transcrivant un WAV via stt_engine
// (Whisper NPU QNN + decoder KV-cache).
//
// Usage : kazeia_stt_cli <model_dir> <audio16k_mono.wav> [language=fr] [cpu|htp=htp]
// model_dir doit contenir : HfWhisperEncoder.onnx + _qairt_context.bin,
// HfWhisperDecoder.onnx + _qairt_context.bin,
// mel_filters.json (ou .bin), vocab.json
#include "stt_engine.h"
#include <cstdio>
#include <cstring>
#include <cstdint>
#include <vector>
#include <string>
static std::vector<int16_t> read_wav_pcm16(const char* path, int & sr_out) {
FILE* f = fopen(path, "rb");
if (!f) { fprintf(stderr, "open %s FAIL\n", path); return {}; }
char riff[12];
if (fread(riff, 1, 12, f) != 12 || memcmp(riff, "RIFF", 4) != 0 || memcmp(riff+8, "WAVE", 4) != 0) {
fprintf(stderr, "not RIFF/WAVE\n"); fclose(f); return {};
}
uint16_t channels = 0, bps = 0;
uint32_t sr = 0, data_sz = 0;
long data_off = -1;
while (!feof(f)) {
char cid[4]; uint32_t csz;
if (fread(cid, 1, 4, f) != 4) break;
if (fread(&csz, 4, 1, f) != 1) break;
if (memcmp(cid, "fmt ", 4) == 0) {
std::vector<uint8_t> buf(csz); fread(buf.data(), 1, csz, f);
channels = *(uint16_t*)&buf[2]; sr = *(uint32_t*)&buf[4]; bps = *(uint16_t*)&buf[14];
} else if (memcmp(cid, "data", 4) == 0) {
data_sz = csz; data_off = ftell(f); break;
} else { fseek(f, csz, SEEK_CUR); }
}
if (data_off < 0 || channels != 1 || bps != 16) {
fprintf(stderr, "WAV mono 16-bit attendu (ch=%d bps=%d)\n", channels, bps);
fclose(f); return {};
}
sr_out = (int)sr;
fseek(f, data_off, SEEK_SET);
std::vector<int16_t> pcm(data_sz / 2);
fread(pcm.data(), 2, pcm.size(), f);
fclose(f);
return pcm;
}
int main(int argc, char** argv) {
setvbuf(stderr, nullptr, _IONBF, 0);
setvbuf(stdout, nullptr, _IONBF, 0);
if (argc < 3) {
fprintf(stderr, "usage: %s <model_dir> <audio16k_mono.wav> [language=fr] [cpu|htp=htp]\n",
argv[0]);
return 1;
}
const char * model_dir = argv[1];
const char * wav_path = argv[2];
const char * lang = (argc >= 4) ? argv[3] : "fr";
const bool use_htp = !(argc >= 5 && !strcmp(argv[4], "cpu"));
int sr = 0;
auto pcm = read_wav_pcm16(wav_path, sr);
if (pcm.empty()) return 2;
fprintf(stderr, "wav : %zu samples @ %d Hz = %.2f s\n",
pcm.size(), sr, (double)pcm.size() / sr);
if (sr != 16000) {
fprintf(stderr, "ERROR: sr=%d != 16000 (Whisper requiert 16kHz)\n", sr); return 3;
}
SttEngineLoadCfg lc;
lc.model_dir = model_dir;
lc.use_htp = use_htp;
lc.n_threads = 6;
auto * eng = stt_engine_load(lc);
if (!eng) { fprintf(stderr, "stt_engine_load FAIL\n"); return 4; }
SttTranscribeCfg tc;
tc.pcm16 = pcm.data();
tc.n_samples = (int)pcm.size();
tc.sample_rate = 16000;
tc.language = lang;
tc.force_transcribe = true;
auto R = stt_engine_transcribe(eng, tc);
if (R.err) { fprintf(stderr, "transcribe FAIL err=%d\n", R.err); stt_engine_free(eng); return 5; }
printf("\n=== TRANSCRIPTION (%s, %s) ===\n", use_htp ? "NPU" : "CPU", lang);
printf("%s\n", R.text.c_str());
printf("\n=== TIMING ===\n");
printf("mel %4d ms\n", R.mel_ms);
printf("encoder %4d ms\n", R.encoder_ms);
printf("decoder %4d ms (%d tokens)\n", R.decoder_ms, R.n_tokens);
printf("total %4d ms for %.2f s audio => RTF %.3f\n",
R.total_ms, (double)pcm.size() / sr, R.rtf);
stt_engine_free(eng);
return 0;
}