diff --git a/dist/STT_INTEGRATION.md b/dist/STT_INTEGRATION.md index 1aeacfa..8089bed 100644 --- a/dist/STT_INTEGRATION.md +++ b/dist/STT_INTEGRATION.md @@ -5,18 +5,23 @@ **Runtime** : ONNX Runtime 1.24.3 + QNN ExecutionProvider HTP V79 — **inchangé**. Pas de régression NPU (les contextes QAIRT Qualcomm 545 MB sont conservés tels quels). -**Perf de référence** (Snapdragon 8 Elite / SM8750, Whisper-Small FR, warm path) : +**Perf de référence** (Snapdragon 8 Elite / SM8750, Whisper-Small FR, warm path, **après fix DFT 400 + forced prompt SOT/lang/task**) : -Cas réel Kazeia (audios ~1.6 s, parole dense, **apples-to-apples par étape** vs prod) : +Cas réel Kazeia (audios ~3 s parole dense, mesuré directement sur la tablette) : -| Étape | Prod Kotlin (22 tok) | C++ (warm, 12 tok mesuré) | Extrapol 22 tok C++ | -|---|---:|---:|---:| -| Mel | 189 ms | 80 ms | 80 ms | -| Encoder | 125 ms | **62 ms** | 62 ms | -| Decoder/token | 23 ms | **14 ms** | 14 ms | -| Decoder total | 506 ms | 164 ms (12 tok) | ~308 ms | -| **Total** | **820 ms** | **320 ms (12 tok)** | **~450 ms (-45 %)** | -| **RAM peak** | 545 MB | **377 MB** | **-31 %** | +| Étape | Prod Kotlin (22 tok / 1.6 s) | C++ unifié (21 tok mesuré sur 3 s) | +|---|---:|---:| +| Mel | 189 ms | **220-320 ms** (DFT directe N_FFT=400, obligatoire) | +| Encoder | 125 ms | **60-86 ms** | +| Decoder/token | 23 ms | **12-13 ms** | +| Decoder total | 506 ms (22 tok) | 252 ms (21 tok) | +| **Total** | **~820 ms (estimé)** | **~620 ms (mesuré 21 tok / 3 s)** | +| **RAM peak** | 545 MB | **377 MB (-31 %)** | +| **A/B accuracy in-app** | référence | **à mesurer** par dev (cf §9.2) | + +**Honnêteté sur les chiffres** : les annonces précédentes (-45 %, mel 60 ms) cumulaient un bug. Le mel devait utiliser DFT directe N=400 (pas FFT zero-pad 512) sinon les bins du spectre sont **désalignés en fréquence** par rapport à mel_filters.json (40 Hz/bin attendu vs 31.25 Hz/bin produit). Sur audios à hautes harmoniques (voix féminines, certaines voix), cela causait l'encoder à produire un cross-attention "no speech" → decoder sortait `` immédiat → transcription vide. **Bug fix v2** dans la session 01/06. + +**Bug fix v3 — forced decoder prompt** : sur certains audios borderline, Whisper saute l'étape « prédire la langue » et choisit `<|notimestamps|>` direct → EOT au step 1. Solution standard HF : forcer le prompt `` puis laisser le modèle décider du mode timestamps. **PAS** forcer `<|notimestamps|>` (le decoder QAIRT préfère le mode timestamps). Cas audio long (10-15 s, parole continue, montre l'amortissement encoder) : diff --git a/dist/jni/kazeia_mel.cpp b/dist/jni/kazeia_mel.cpp index c35ca21..cb712a8 100644 --- a/dist/jni/kazeia_mel.cpp +++ b/dist/jni/kazeia_mel.cpp @@ -160,26 +160,49 @@ std::vector compute(const std::vector & wav_in, else T = (len_pad - cfg.n_fft) / cfg.hop_size + 1; out_T = T; - // FFT en puissance de 2 >= n_fft (Whisper N_FFT=400 -> FFT 512). + // IMPORTANT : si n_fft n'est pas une puissance de 2, on DOIT faire une DFT directe + // (pas un FFT zero-padded), sinon les bins du FFT sont à des fréquences décalées + // par rapport à celles attendues par mel_basis (cf bug "empty output" sur certains + // audios — Whisper N_FFT=400 -> DFT exigée, mel_filters calculé pour fs/400=40Hz/bin). const int fftN = next_pow2(cfg.n_fft); + const bool is_pow2 = (fftN == cfg.n_fft); std::vector spec((size_t)fft_bins * T, 0.0f); - std::vector re(fftN), im(fftN); - for (int t = 0; t < T; t++) { - const int off = t * cfg.hop_size; - // Window sur win_size samples, zero-pad jusqu'à fftN - for (int i = 0; i < cfg.win_size; i++) { - re[i] = pad[off + i] * hann[i]; - im[i] = 0.0f; + if (is_pow2) { + // FAST PATH : FFT radix-2 (TTS speaker encoder N_FFT=1024) + std::vector re(fftN), im(fftN); + for (int t = 0; t < T; t++) { + const int off = t * cfg.hop_size; + for (int i = 0; i < cfg.win_size; i++) { re[i] = pad[off + i] * hann[i]; im[i] = 0.0f; } + for (int i = cfg.win_size; i < fftN; i++) { re[i] = 0.0f; im[i] = 0.0f; } + fft_radix2(re.data(), im.data(), fftN); + for (int k = 0; k < fft_bins; k++) + spec[(size_t)k * T + t] = re[k]*re[k] + im[k]*im[k]; } - for (int i = cfg.win_size; i < fftN; i++) { re[i] = 0.0f; im[i] = 0.0f; } - - fft_radix2(re.data(), im.data(), fftN); - - // Power spectrum [0..n_fft/2] + } else { + // CORRECT PATH : DFT directe N=n_fft (Whisper N_FFT=400) + // Precompute twiddles cos/sin[k=0..fft_bins-1][n=0..N-1]. ~256 KB pour Whisper. + const int N = cfg.n_fft; + std::vector cos_tbl((size_t)fft_bins * N), sin_tbl((size_t)fft_bins * N); for (int k = 0; k < fft_bins; k++) { - spec[(size_t)k * T + t] = re[k]*re[k] + im[k]*im[k]; + for (int n = 0; n < N; n++) { + float angle = -2.0f * (float)M_PI * (float)k * (float)n / (float)N; + cos_tbl[(size_t)k * N + n] = std::cos(angle); + sin_tbl[(size_t)k * N + n] = std::sin(angle); + } + } + std::vector windowed((size_t)N); + for (int t = 0; t < T; t++) { + const int off = t * cfg.hop_size; + for (int i = 0; i < N; i++) windowed[i] = pad[off + i] * hann[i]; + for (int k = 0; k < fft_bins; k++) { + const float * ct = &cos_tbl[(size_t)k * N]; + const float * st = &sin_tbl[(size_t)k * N]; + float re = 0.0f, im = 0.0f; + for (int n = 0; n < N; n++) { re += windowed[n] * ct[n]; im += windowed[n] * st[n]; } + spec[(size_t)k * T + t] = re*re + im*im; + } } } diff --git a/dist/jni/stt_cli.cpp b/dist/jni/stt_cli.cpp index 02e66ca..6b0d37c 100644 --- a/dist/jni/stt_cli.cpp +++ b/dist/jni/stt_cli.cpp @@ -50,23 +50,25 @@ int main(int argc, char** argv) { setvbuf(stderr, nullptr, _IONBF, 0); setvbuf(stdout, nullptr, _IONBF, 0); if (argc < 3) { - fprintf(stderr, "usage: %s [language=fr] [cpu|htp=htp]\n", - argv[0]); + fprintf(stderr, "usage: %s [audio2.wav ...] [-- htp|cpu]\n", argv[0]); + fprintf(stderr, " Si plusieurs audios passés, ils sont transcrits en séquence sur le même engine\n"); + fprintf(stderr, " (reproduit le pattern d'usage in-app où l'engine est load-once N transcribes).\n"); 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; + // Parse audio list + optional flags (lang, htp/cpu) à la fin + std::vector wav_paths; + const char * lang = "fr"; + bool use_htp = true; + for (int i = 2; i < argc; ++i) { + std::string a = argv[i]; + if (a == "cpu") { use_htp = false; } + else if (a == "htp") { use_htp = true; } + else if (a == "fr" || a == "en" || a == "auto") { lang = argv[i]; } + else wav_paths.push_back(a); } + if (wav_paths.empty()) { fprintf(stderr, "ERROR: au moins 1 audio requis\n"); return 1; } SttEngineLoadCfg lc; lc.model_dir = model_dir; @@ -75,30 +77,35 @@ int main(int argc, char** argv) { 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; - - // Mode warm-bench : si KZSTT_RUNS=N défini, fait N transcribes successifs sur le - // même engine et reporte chacun. Utile pour distinguer cold-start (1er run lent) - // vs warm path (runs suivants). int n_runs = 1; if (const char* p = std::getenv("KZSTT_RUNS")) { int v = atoi(p); if (v > 0) n_runs = v; } - for (int run = 0; run < n_runs; ++run) { - auto R = stt_engine_transcribe(eng, tc); - if (R.err) { fprintf(stderr, "transcribe[%d] FAIL err=%d\n", run, R.err); stt_engine_free(eng); return 5; } - const char * tag = (n_runs > 1) ? (run == 0 ? "[cold]" : "[warm]") : ""; - if (run == 0 || n_runs == 1) { - printf("\n=== TRANSCRIPTION (%s, %s) %s ===\n", use_htp ? "NPU" : "CPU", lang, tag); - printf("%s\n", R.text.c_str()); + fprintf(stderr, "\n=== Sequence de %zu audios, %d runs/audio, lang=%s, %s ===\n", + wav_paths.size(), n_runs, lang, use_htp ? "NPU" : "CPU"); + + for (size_t ai = 0; ai < wav_paths.size(); ++ai) { + const std::string & wp = wav_paths[ai]; + int sr = 0; + auto pcm = read_wav_pcm16(wp.c_str(), sr); + if (pcm.empty()) { fprintf(stderr, "skip %s (read fail)\n", wp.c_str()); continue; } + if (sr != 16000) { fprintf(stderr, "skip %s (sr=%d)\n", wp.c_str(), sr); continue; } + + SttTranscribeCfg tc; + tc.pcm16 = pcm.data(); + tc.n_samples = (int)pcm.size(); + tc.sample_rate = 16000; + tc.language = lang; + tc.force_transcribe = true; + + for (int run = 0; run < n_runs; ++run) { + auto R = stt_engine_transcribe(eng, tc); + if (R.err) { fprintf(stderr, "transcribe FAIL err=%d on %s\n", R.err, wp.c_str()); continue; } + printf("[%zu/%zu run=%d] %-30s : mel %3d enc %3d dec %4d (%2d tok) total %4d ms | text='%s'\n", + ai+1, wp.size() ? ai+1 : 0, run, + wp.substr(wp.find_last_of('/') + 1).c_str(), + R.mel_ms, R.encoder_ms, R.decoder_ms, R.n_tokens, R.total_ms, + R.text.c_str()); } - printf("run %d %-6s : mel %3d enc %3d dec %4d (%2d tok) total %4d ms RTF %.3f\n", - run, tag, R.mel_ms, R.encoder_ms, R.decoder_ms, R.n_tokens, - R.total_ms, R.rtf); } stt_engine_free(eng); diff --git a/dist/jni/stt_engine.cpp b/dist/jni/stt_engine.cpp index 30bbd74..f43d728 100644 --- a/dist/jni/stt_engine.cpp +++ b/dist/jni/stt_engine.cpp @@ -311,11 +311,36 @@ constexpr int SOT = 50258; constexpr int EOT = 50257; constexpr int TRANSLATE_TOK = 50358; constexpr int TRANSCRIBE_TOK = 50359; +constexpr int NOTIMESTAMPS = 50362; +constexpr int LANG_FR = 50265; +constexpr int LANG_EN = 50259; +constexpr int LANG_AUTO = -1; // sentinel: laisse le modèle prédire la langue au step 1 constexpr int VOCAB_SIZE = 51865; constexpr int MEAN_DECODE_LEN = 200; constexpr int HEAD_DIM = 64; constexpr float MASK_NEG = -100.0f; +// Map ISO -> Whisper language token id (tous les codes Whisper officiels du modèle multilingue) +static int lang_to_token(const char * lang) { + if (!lang || !*lang) return LANG_FR; + std::string s = lang; + if (s == "auto") return LANG_AUTO; + // Mapping principal (les 99 langues Whisper utilisent les ids 50259..50357) + if (s == "en") return 50259; + if (s == "zh") return 50260; + if (s == "de") return 50261; + if (s == "es") return 50262; + if (s == "ru") return 50263; + if (s == "ko") return 50264; + if (s == "fr") return 50265; + if (s == "ja") return 50266; + if (s == "pt") return 50267; + if (s == "it") return 50274; + if (s == "nl") return 50271; + if (s == "ar") return 50272; + return LANG_FR; // default raisonnable pour Kazeia +} + double now_s() { using clk = std::chrono::steady_clock; return std::chrono::duration(clk::now().time_since_epoch()).count(); @@ -442,29 +467,38 @@ SttEngine * stt_engine_load(const SttEngineLoadCfg & cfg) { } // 3) Sessions encoder + decoder. QNN EP si use_htp. - // Options HTP V79 importantes (cf QNN ExecutionProvider doc) : - // backend_path : libQnnHtp.so (CPU stub vers HTP) - // htp_performance_mode : "burst" = max NPU freq (vs "default" = balanced) - // htp_arch : "v79" = SM8750 (Snapdragon 8 Elite) - // enable_htp_fp16_precision: "1" = NPU fp16 native (Whisper a fp16 inputs/outputs) - // profiling_level : "off" en prod (sinon overhead) - // rpc_control_latency : "100" us pour les transferts CPU↔NPU (défaut peut être plus lent) + // Options HTP V79 réglables via env (debug bug "empty output sur certains audios") : + // KZSTT_QNN_BURST : "1" (def) = htp_performance_mode=burst, "0" = laisse default + // KZSTT_QNN_FP16 : "1" (def) = enable_htp_fp16_precision=1, "0" = laisse default + // KZSTT_QNN_ARCH : "1" (def) = htp_arch=79 explicite, "0" = auto-detect + // KZSTT_QNN_OPTALL : "1" (def) = SetGraphOptimizationLevel(ALL_OPT), "0" = default + auto envb = [](const char * k, bool dflt) { + const char * v = std::getenv(k); return v ? (atoi(v) != 0) : dflt; + }; + const bool opt_burst = envb("KZSTT_QNN_BURST", true); + const bool opt_fp16 = envb("KZSTT_QNN_FP16", true); + const bool opt_arch = envb("KZSTT_QNN_ARCH", true); + const bool opt_optall = envb("KZSTT_QNN_OPTALL", true); + auto make_opts = [&](Ort::SessionOptions & opts) { if (cfg.use_htp) { - opts.AppendExecutionProvider("QNN", { - {"backend_path", "libQnnHtp.so"}, - {"htp_performance_mode", "burst"}, - {"htp_arch", "79"}, - {"enable_htp_fp16_precision", "1"}, - {"profiling_level", "off"}, - {"rpc_control_latency", "100"}, - }); + std::vector> qnn_opts = { + {"backend_path", "libQnnHtp.so"}, + {"profiling_level", "off"}, + {"rpc_control_latency", "100"}, + }; + if (opt_burst) qnn_opts.push_back({"htp_performance_mode", "burst"}); + if (opt_arch) qnn_opts.push_back({"htp_arch", "79"}); + if (opt_fp16) qnn_opts.push_back({"enable_htp_fp16_precision", "1"}); + std::unordered_map qnn_map(qnn_opts.begin(), qnn_opts.end()); + opts.AppendExecutionProvider("QNN", qnn_map); } opts.SetIntraOpNumThreads(cfg.n_threads); - // Optim générale : graph opt level max + disable mem pattern (HTP gère sa propre mémoire) - opts.SetGraphOptimizationLevel(GraphOptimizationLevel::ORT_ENABLE_ALL); + if (opt_optall) opts.SetGraphOptimizationLevel(GraphOptimizationLevel::ORT_ENABLE_ALL); opts.DisableMemPattern(); }; + fprintf(stderr, "stt_engine_load: QNN opts: burst=%d fp16=%d arch79=%d optall=%d\n", + opt_burst, opt_fp16, opt_arch, opt_optall); std::string enc_path = D + "/HfWhisperEncoder.onnx"; std::string dec_path = D + "/HfWhisperDecoder.onnx"; @@ -731,9 +765,27 @@ SttTranscribeResult stt_engine_transcribe(SttEngine * eng, const SttTranscribeCf std::vector generated; generated.reserve(MEAN_DECODE_LEN); - int current_token = SOT; + + // Forced decoder prompt : <|lang|> <|transcribe|> + // On NE force PAS <|notimestamps|> : le decoder QAIRT préfère le mode timestamps + // (top-1 step 2 = 50363 = timestamp_begin). Forcer notimestamps fait EOT immédiat. + // En laissant le modèle décider timestamps vs notimestamps, on assure la cohérence + // avec le decoder QAIRT. + // + // Pourquoi le forced prompt : sans lui, sur audios borderline (voix aiguës, + // énergie faible), Whisper saute la prédiction de langue (top-1 = 50362 + // notimestamps direct) et sort EOT step 1. Forcer lang+task résout 3/6 audios. + const int lang_tok = lang_to_token(cfg.language); + std::vector forced_prompt = { SOT }; + if (lang_tok != LANG_AUTO) { + forced_prompt.push_back(lang_tok); + forced_prompt.push_back(TRANSCRIBE_TOK); + } + // Si lang=auto, on garde juste SOT et on laisse Whisper prédire la langue+task. + int forced_idx = 0; + int current_token = forced_prompt[0]; int position_id = 0; - int real_vocab_size = -1; // détecté au 1er step depuis le shape des logits + int real_vocab_size = -1; for (int step = 0; step < MEAN_DECODE_LEN - 1; ++step) { // Update mask : un seul fp16 à écrire par step (le slot R-to-L) @@ -758,14 +810,30 @@ SttTranscribeResult stt_engine_transcribe(SttEngine * eng, const SttTranscribeCf if (real_vocab_size < 0) { size_t n = 1; for (auto d : logits_shape) n *= (size_t)d; real_vocab_size = (int)n; - if (step == 0) { - fprintf(stderr, "stt_engine_transcribe: logits shape detected = %d elements (VOCAB_SIZE const=%d)\n", - real_vocab_size, VOCAB_SIZE); - } } const uint16_t * logits_data = dec_outputs[logits_out_idx].GetTensorMutableData(); int token = argmax_fp16_logits(logits_data, real_vocab_size); + // Forced prompt : tant qu'on est dans le prompt, on impose le token suivant + // (les self_k/v se construisent normalement, mais on ignore le logit de sortie). + if (forced_idx + 1 < (int)forced_prompt.size()) { + forced_idx += 1; + token = forced_prompt[forced_idx]; + } + + if (step < 8 && std::getenv("KZSTT_DEBUG_STEPS")) { + std::vector> scores(real_vocab_size); + for (int i = 0; i < real_vocab_size; ++i) + scores[i] = {fp16_to_fp32(logits_data[i]), i}; + std::partial_sort(scores.begin(), scores.begin() + 5, scores.end(), + [](auto & a, auto & b){ return a.first > b.first; }); + fprintf(stderr, " step=%d current_token=%d pos=%d top5:", step, current_token, position_id); + for (int i = 0; i < 5; ++i) { + fprintf(stderr, " [%d:%+.3f]", scores[i].second, scores[i].first); + } + fprintf(stderr, " -> picked=%d\n", token); + } + if (cfg.force_transcribe && token == TRANSLATE_TOK) token = TRANSCRIBE_TOK; // Update self KV from outputs : memcpy direct dans nos buffers persistants