diff --git a/electron/native/whisper-stt/CMakeLists.txt b/electron/native/whisper-stt/CMakeLists.txt index 0f6f358da..a22d02eed 100644 --- a/electron/native/whisper-stt/CMakeLists.txt +++ b/electron/native/whisper-stt/CMakeLists.txt @@ -179,9 +179,10 @@ if(NOT "${CMAKE_CURRENT_BINARY_DIR}/whisper-patched/whisper.cpp" IN_LIST osc_whi endif() set_property(TARGET whisper PROPERTY SOURCES ${osc_whisper_srcs}) -add_executable(whisper-stt-server src/main.cpp) +add_executable(whisper-stt-server src/main.cpp src/ctc_aligner.cpp) target_link_libraries(whisper-stt-server PRIVATE whisper + ggml httplib::httplib nlohmann_json::nlohmann_json ) diff --git a/electron/native/whisper-stt/src/ctc_aligner.cpp b/electron/native/whisper-stt/src/ctc_aligner.cpp new file mode 100644 index 000000000..18d514596 --- /dev/null +++ b/electron/native/whisper-stt/src/ctc_aligner.cpp @@ -0,0 +1,433 @@ +// wav2vec2-for-CTC forward pass on ggml (HuggingFace Wav2Vec2ForCTC, both the +// "group norm, post-LN" base layout and the "layer norm, pre-LN" large one). +// Weights come from scripts/convert-wav2vec2-gguf.mjs. +// +// Layout: activations are [channels, time] (ne0 = channels) for the linear +// layers, and transposed to [time, channels] only where a convolution or a +// per-channel norm needs time contiguous. + +#include "ctc_aligner.h" + +#include +#include +#include +#include + +#include +#include +#include + +namespace { + +constexpr int kGraphSize = 8192; +// Longest input one graph sees; longer ones run in windows. Attention is +// quadratic in time and the first convolutions are wide, so this bounds memory +// (about 265 MB of compute buffers at 20 s, for either model: the first +// convolutions dominate). +constexpr int kWindowFrames = 1000; // 20 s +constexpr int kContextFrames = 50; // 1 s of context kept on each side of a window + +struct Layer { + ggml_tensor *q_w, *q_b, *k_w, *k_b, *v_w, *v_b, *o_w, *o_b; + ggml_tensor *ln1_w, *ln1_b, *ff1_w, *ff1_b, *ff2_w, *ff2_b, *ln2_w, *ln2_b; +}; + +} // namespace + +struct CtcModel { + CtcModelInfo info; + int hidden = 0, n_heads = 0, pos_k = 0, pos_groups = 0; + float eps = 1e-5f; + bool group_norm = false, stable = false; + std::vector conv_k, conv_s; + + std::vector conv_w, conv_b, conv_ln_w, conv_ln_b; + ggml_tensor *fp_ln_w = nullptr, *fp_ln_b = nullptr, *fp_w = nullptr, *fp_b = nullptr; + ggml_tensor *pos_w = nullptr, *pos_b = nullptr, *enc_ln_w = nullptr, *enc_ln_b = nullptr; + ggml_tensor *head_w = nullptr, *head_b = nullptr; + std::vector layers; + + ggml_context* wctx = nullptr; + gguf_context* gctx = nullptr; + ggml_backend_buffer_t wbuf = nullptr; + ggml_backend_t gpu = nullptr, cpu = nullptr; + ggml_backend_sched_t sched = nullptr; + std::string device = "CPU"; + + ~CtcModel() { + if (sched) ggml_backend_sched_free(sched); + if (wbuf) ggml_backend_buffer_free(wbuf); + if (gpu) ggml_backend_free(gpu); + if (cpu) ggml_backend_free(cpu); + if (wctx) ggml_free(wctx); + if (gctx) gguf_free(gctx); + } +}; + +void CtcModelDeleter::operator()(CtcModel* m) const { delete m; } + +const CtcModelInfo& ctc_info(const CtcModel& model) { return model.info; } +std::string ctc_device(const CtcModel& model) { return model.device; } + +namespace { + +int64_t key(gguf_context* g, const char* k, std::string& err) { + const int64_t id = gguf_find_key(g, k); + if (id < 0 && err.empty()) err = std::string("missing metadata ") + k; + return id; +} + +std::vector str_array(gguf_context* g, int64_t id) { + std::vector out; + if (id < 0) return out; + const size_t n = gguf_get_arr_n(g, id); + for (size_t i = 0; i < n; ++i) out.emplace_back(gguf_get_arr_str(g, id, i)); + return out; +} + +std::vector u32_array(gguf_context* g, int64_t id) { + std::vector out; + if (id < 0) return out; + const auto* p = static_cast(gguf_get_arr_data(g, id)); + out.assign(p, p + gguf_get_arr_n(g, id)); + return out; +} + +ggml_tensor* layer_norm(ggml_context* ctx, ggml_tensor* x, ggml_tensor* w, ggml_tensor* b, float eps) { + return ggml_add(ctx, ggml_mul(ctx, ggml_norm(ctx, x, eps), w), b); +} + +ggml_tensor* linear(ggml_context* ctx, ggml_tensor* x, ggml_tensor* w, ggml_tensor* b) { + return ggml_add(ctx, ggml_mul_mat(ctx, w, x), b); +} + +// 1-D convolution of `x` [time, in] with `w` [k, in, out] -> [out, time']. The +// columns are F16, as in whisper.cpp's encoder: half the memory traffic of F32, +// which made the CPU path 1.7x slower, and `w` must be F16 to match. +ggml_tensor* conv1d(ggml_context* ctx, ggml_tensor* w, ggml_tensor* x, int stride, int pad) { + ggml_tensor* cols = ggml_im2col(ctx, w, x, stride, 0, pad, 0, 1, 0, false, GGML_TYPE_F16); + // cols: [k * in, time'] + cols = ggml_reshape_2d(ctx, cols, cols->ne[0], cols->ne[1]); + ggml_tensor* w2 = ggml_reshape_2d(ctx, w, w->ne[0] * w->ne[1], w->ne[2]); + return ggml_mul_mat(ctx, w2, cols); +} + +ggml_tensor* build(CtcModel& m, ggml_context* ctx, ggml_tensor* wave) { + // ---- feature encoder ---- + ggml_tensor* x = wave; // [time, 1] + ggml_tensor* y = nullptr; + for (size_t i = 0; i < m.conv_w.size(); ++i) { + y = conv1d(ctx, m.conv_w[i], x, m.conv_s[i], 0); // [512, t] + if (m.conv_b[i]) y = ggml_add(ctx, y, m.conv_b[i]); + if (m.group_norm) { + // Time-major for the next convolution. Only the first layer is normed: + // GroupNorm with one group per channel, i.e. each channel over time. + ggml_tensor* t = ggml_cont(ctx, ggml_transpose(ctx, y)); // [t, 512] + if (i == 0) { + t = ggml_norm(ctx, t, m.eps); + t = ggml_add(ctx, ggml_mul(ctx, t, ggml_reshape_2d(ctx, m.conv_ln_w[0], 1, t->ne[1])), + ggml_reshape_2d(ctx, m.conv_ln_b[0], 1, t->ne[1])); + } + x = ggml_gelu_erf(ctx, t); + } else { + y = layer_norm(ctx, y, m.conv_ln_w[i], m.conv_ln_b[i], m.eps); + x = ggml_cont(ctx, ggml_transpose(ctx, ggml_gelu_erf(ctx, y))); // [t, 512] + } + } + // x: [t, 512] -> [512, t] + ggml_tensor* h = ggml_cont(ctx, ggml_transpose(ctx, x)); + const int64_t T = h->ne[1]; + + // ---- feature projection ---- + h = layer_norm(ctx, h, m.fp_ln_w, m.fp_ln_b, m.eps); + h = linear(ctx, h, m.fp_w, m.fp_b); // [H, T] + + // ---- positional convolution (grouped, same padding, drop the last frame) ---- + { + ggml_tensor* ht = ggml_cont(ctx, ggml_transpose(ctx, h)); // [T, H] + const int64_t cg = m.hidden / m.pos_groups; + ggml_tensor* pos = nullptr; + for (int g = 0; g < m.pos_groups; ++g) { + ggml_tensor* xin = ggml_view_2d(ctx, ht, T, cg, ht->nb[1], g * cg * ht->nb[1]); + ggml_tensor* wg = ggml_view_3d(ctx, m.pos_w, m.pos_k, cg, cg, m.pos_w->nb[1], + m.pos_w->nb[2], g * cg * m.pos_w->nb[2]); + ggml_tensor* part = conv1d(ctx, wg, ggml_cont(ctx, xin), 1, m.pos_k / 2); // [cg, T+1] + pos = pos ? ggml_concat(ctx, pos, part, 0) : part; + } + pos = ggml_view_2d(ctx, pos, m.hidden, T, pos->nb[1], 0); + pos = ggml_gelu_erf(ctx, ggml_add(ctx, pos, m.pos_b)); + h = ggml_add(ctx, h, pos); + } + if (!m.stable) h = layer_norm(ctx, h, m.enc_ln_w, m.enc_ln_b, m.eps); + + // ---- transformer ---- + const int64_t dh = m.hidden / m.n_heads; + const float scale = 1.0f / std::sqrt(static_cast(dh)); + for (const Layer& L : m.layers) { + ggml_tensor* res = h; + ggml_tensor* a = m.stable ? layer_norm(ctx, h, L.ln1_w, L.ln1_b, m.eps) : h; + ggml_tensor* q = ggml_reshape_3d(ctx, linear(ctx, a, L.q_w, L.q_b), dh, m.n_heads, T); + ggml_tensor* k = ggml_reshape_3d(ctx, linear(ctx, a, L.k_w, L.k_b), dh, m.n_heads, T); + ggml_tensor* v = ggml_reshape_3d(ctx, linear(ctx, a, L.v_w, L.v_b), dh, m.n_heads, T); + q = ggml_permute(ctx, q, 0, 2, 1, 3); // [dh, T, nh] + k = ggml_permute(ctx, k, 0, 2, 1, 3); // [dh, T, nh] + v = ggml_cont(ctx, ggml_permute(ctx, v, 1, 2, 0, 3)); // [T, dh, nh] + ggml_tensor* kq = ggml_mul_mat(ctx, k, q); // [Tk, Tq, nh] + kq = ggml_soft_max_ext(ctx, kq, nullptr, scale, 0.0f); + ggml_tensor* o = ggml_mul_mat(ctx, v, kq); // [dh, Tq, nh] + o = ggml_cont(ctx, ggml_permute(ctx, o, 0, 2, 1, 3)); // [dh, nh, T] + o = linear(ctx, ggml_reshape_2d(ctx, o, m.hidden, T), L.o_w, L.o_b); + h = ggml_add(ctx, res, o); + if (m.stable) { + ggml_tensor* f = layer_norm(ctx, h, L.ln2_w, L.ln2_b, m.eps); + f = linear(ctx, ggml_gelu_erf(ctx, linear(ctx, f, L.ff1_w, L.ff1_b)), L.ff2_w, L.ff2_b); + h = ggml_add(ctx, h, f); + } else { + h = layer_norm(ctx, h, L.ln1_w, L.ln1_b, m.eps); + ggml_tensor* f = linear(ctx, ggml_gelu_erf(ctx, linear(ctx, h, L.ff1_w, L.ff1_b)), L.ff2_w, L.ff2_b); + h = layer_norm(ctx, ggml_add(ctx, h, f), L.ln2_w, L.ln2_b, m.eps); + } + } + if (m.stable) h = layer_norm(ctx, h, m.enc_ln_w, m.enc_ln_b, m.eps); + return linear(ctx, h, m.head_w, m.head_b); // [V, T] +} + +int frames_for(const CtcModel& m, int64_t n) { + for (size_t i = 0; i < m.conv_k.size(); ++i) { + if (n < m.conv_k[i]) return 0; + n = (n - m.conv_k[i]) / m.conv_s[i] + 1; + } + return static_cast(n); +} + +// Logits for one window of already-normalized samples. +bool run_window(CtcModel& m, const float* pcm, int64_t n, std::vector& logits, int& frames, + std::string& err) { + frames = frames_for(m, n); + if (frames <= 0) return true; + ggml_init_params p = {ggml_tensor_overhead() * kGraphSize + ggml_graph_overhead_custom(kGraphSize, false), + nullptr, true}; + ggml_context* ctx = ggml_init(p); + ggml_tensor* wave = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n, 1); + ggml_set_input(wave); + ggml_tensor* out = build(m, ctx, wave); + ggml_set_output(out); + ggml_cgraph* gf = ggml_new_graph_custom(ctx, kGraphSize, false); + ggml_build_forward_expand(gf, out); + ggml_backend_sched_reset(m.sched); + bool ok = ggml_backend_sched_alloc_graph(m.sched, gf); + if (ok) { + ggml_backend_tensor_set(wave, pcm, 0, n * sizeof(float)); + ok = ggml_backend_sched_graph_compute(m.sched, gf) == GGML_STATUS_SUCCESS; + } + if (ok) { + if (out->ne[1] != frames) { + err = "unexpected frame count"; + ok = false; + } else { + logits.resize(static_cast(out->ne[0]) * frames); + ggml_backend_tensor_get(out, logits.data(), 0, logits.size() * sizeof(float)); + } + } else if (err.empty()) { + err = "ggml graph allocation or compute failed"; + } + ggml_free(ctx); + return ok; +} + +} // namespace + +CtcModelPtr ctc_load(const std::string& path, bool use_gpu, int threads, std::string& err) { + CtcModelPtr m(new CtcModel()); + gguf_init_params gp = {true, &m->wctx}; + m->gctx = gguf_init_from_file(path.c_str(), gp); + if (!m->gctx) { + err = "cannot read " + path; + return nullptr; + } + gguf_context* g = m->gctx; + CtcModelInfo& info = m->info; + info.vocab = str_array(g, key(g, "w2v.vocab", err)); + info.languages = str_array(g, key(g, "w2v.languages", err)); + int64_t id; + if ((id = key(g, "w2v.blank_id", err)) >= 0) info.blank = gguf_get_val_u32(g, id); + if ((id = key(g, "w2v.hidden_size", err)) >= 0) m->hidden = gguf_get_val_u32(g, id); + if ((id = key(g, "w2v.num_attention_heads", err)) >= 0) m->n_heads = gguf_get_val_u32(g, id); + if ((id = key(g, "w2v.layer_norm_eps", err)) >= 0) m->eps = gguf_get_val_f32(g, id); + if ((id = key(g, "w2v.num_conv_pos_embeddings", err)) >= 0) m->pos_k = gguf_get_val_u32(g, id); + if ((id = key(g, "w2v.num_conv_pos_embedding_groups", err)) >= 0) m->pos_groups = gguf_get_val_u32(g, id); + if ((id = key(g, "w2v.feat_extract_norm", err)) >= 0) m->group_norm = std::string(gguf_get_val_str(g, id)) == "group"; + if ((id = key(g, "w2v.stable_layer_norm", err)) >= 0) m->stable = gguf_get_val_bool(g, id); + m->conv_k = u32_array(g, key(g, "w2v.conv_kernel", err)); + m->conv_s = u32_array(g, key(g, "w2v.conv_stride", err)); + int n_layers = 0; + if ((id = key(g, "w2v.num_hidden_layers", err)) >= 0) n_layers = gguf_get_val_u32(g, id); + if (!err.empty()) return nullptr; + info.stride = 1; + for (int s : m->conv_s) info.stride *= s; + info.receptive_field = 1; + for (int i = static_cast(m->conv_k.size()) - 1; i >= 0; --i) + info.receptive_field = (info.receptive_field - 1) * m->conv_s[i] + m->conv_k[i]; + + auto T = [&](const std::string& name, bool required = true) -> ggml_tensor* { + ggml_tensor* t = ggml_get_tensor(m->wctx, name.c_str()); + if (!t && required && err.empty()) err = "missing tensor " + name; + return t; + }; + for (size_t i = 0; i < m->conv_k.size(); ++i) { + const std::string p = "feature_extractor.conv_layers." + std::to_string(i) + "."; + m->conv_w.push_back(T(p + "conv.weight")); + m->conv_b.push_back(T(p + "conv.bias", false)); + const bool has_ln = !m->group_norm || i == 0; + m->conv_ln_w.push_back(has_ln ? T(p + "layer_norm.weight") : nullptr); + m->conv_ln_b.push_back(has_ln ? T(p + "layer_norm.bias") : nullptr); + } + m->fp_ln_w = T("feature_projection.layer_norm.weight"); + m->fp_ln_b = T("feature_projection.layer_norm.bias"); + m->fp_w = T("feature_projection.projection.weight"); + m->fp_b = T("feature_projection.projection.bias"); + m->pos_w = T("encoder.pos_conv_embed.conv.weight"); + m->pos_b = T("encoder.pos_conv_embed.conv.bias"); + m->enc_ln_w = T("encoder.layer_norm.weight"); + m->enc_ln_b = T("encoder.layer_norm.bias"); + m->head_w = T("lm_head.weight"); + m->head_b = T("lm_head.bias"); + for (int l = 0; l < n_layers; ++l) { + const std::string p = "encoder.layers." + std::to_string(l) + "."; + m->layers.push_back({ + T(p + "attention.q_proj.weight"), T(p + "attention.q_proj.bias"), + T(p + "attention.k_proj.weight"), T(p + "attention.k_proj.bias"), + T(p + "attention.v_proj.weight"), T(p + "attention.v_proj.bias"), + T(p + "attention.out_proj.weight"), T(p + "attention.out_proj.bias"), + T(p + "layer_norm.weight"), T(p + "layer_norm.bias"), + T(p + "feed_forward.intermediate_dense.weight"), T(p + "feed_forward.intermediate_dense.bias"), + T(p + "feed_forward.output_dense.weight"), T(p + "feed_forward.output_dense.bias"), + T(p + "final_layer_norm.weight"), T(p + "final_layer_norm.bias"), + }); + } + if (!err.empty()) return nullptr; + // conv1d() multiplies them with F16 columns; anything else aborts inside ggml. + for (const ggml_tensor* w : m->conv_w) { + if (w->type != GGML_TYPE_F16) err = "convolution weights must be F16"; + } + if (m->pos_w->type != GGML_TYPE_F16) err = "convolution weights must be F16"; + if (!err.empty()) return nullptr; + if (static_cast(info.vocab.size()) != m->head_w->ne[1]) { + err = "vocabulary does not match lm_head"; + return nullptr; + } + + // ---- backends ---- + if (use_gpu) { + for (size_t i = 0; i < ggml_backend_dev_count(); ++i) { + ggml_backend_dev_t dev = ggml_backend_dev_get(i); + const auto type = ggml_backend_dev_type(dev); + if (type != GGML_BACKEND_DEVICE_TYPE_GPU && type != GGML_BACKEND_DEVICE_TYPE_IGPU) continue; + m->gpu = ggml_backend_dev_init(dev, nullptr); + if (m->gpu) { + m->device = ggml_backend_dev_name(dev); + break; + } + } + } + m->cpu = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + if (!m->cpu) { + err = "no CPU backend"; + return nullptr; + } + { + ggml_backend_reg_t reg = ggml_backend_dev_backend_reg(ggml_backend_get_device(m->cpu)); + using set_threads_t = void (*)(ggml_backend_t, int); + auto set_threads = reinterpret_cast( + ggml_backend_reg_get_proc_address(reg, "ggml_backend_set_n_threads")); + if (set_threads) set_threads(m->cpu, threads); + } + ggml_backend_t weights_on = m->gpu ? m->gpu : m->cpu; + m->wbuf = ggml_backend_alloc_ctx_tensors(m->wctx, weights_on); + if (!m->wbuf) { + err = "cannot allocate the model's weights"; + return nullptr; + } + { + FILE* f = std::fopen(path.c_str(), "rb"); + if (!f) { + err = "cannot open " + path; + return nullptr; + } + std::vector buf; + const size_t data_off = gguf_get_data_offset(g); + for (int64_t i = 0; i < gguf_get_n_tensors(g); ++i) { + ggml_tensor* t = ggml_get_tensor(m->wctx, gguf_get_tensor_name(g, i)); + const size_t nbytes = ggml_nbytes(t); + buf.resize(nbytes); +#ifdef _WIN32 + _fseeki64(f, static_cast(data_off + gguf_get_tensor_offset(g, i)), SEEK_SET); +#else + fseeko(f, static_cast(data_off + gguf_get_tensor_offset(g, i)), SEEK_SET); +#endif + if (std::fread(buf.data(), 1, nbytes, f) != nbytes) { + std::fclose(f); + err = "truncated model file " + path; + return nullptr; + } + ggml_backend_tensor_set(t, buf.data(), 0, nbytes); + } + std::fclose(f); + } + std::vector backends; + if (m->gpu) backends.push_back(m->gpu); + backends.push_back(m->cpu); + m->sched = ggml_backend_sched_new(backends.data(), nullptr, static_cast(backends.size()), + kGraphSize, false, true); + if (!m->sched) { + err = "cannot create the ggml scheduler"; + return nullptr; + } + return m; +} + +bool ctc_emissions(CtcModel& m, const float* pcm, size_t n, std::vector& out, int& frames, + std::string& err) { + const int V = static_cast(m.info.vocab.size()); + const int64_t stride = m.info.stride; + frames = frames_for(m, static_cast(n)); + out.assign(static_cast(std::max(frames, 0)) * V, 0.0f); + if (frames <= 0) return true; + + // Wav2Vec2FeatureExtractor(do_normalize=True): zero mean, unit variance, per input. + double mean = 0.0, var = 0.0; + for (size_t i = 0; i < n; ++i) mean += pcm[i]; + mean /= static_cast(n); + for (size_t i = 0; i < n; ++i) var += (pcm[i] - mean) * (pcm[i] - mean); + var /= static_cast(n); + const double inv = 1.0 / std::sqrt(var + 1e-7); + std::vector x(n); + for (size_t i = 0; i < n; ++i) x[i] = static_cast((pcm[i] - mean) * inv); + + // Windows of kWindowFrames, stepping by kWindowFrames - 2 * kContextFrames; each + // keeps its frames away from the edges it shares with a neighbour. + const int64_t win_samples = kWindowFrames * stride + (m.info.receptive_field - stride); + const int hop = kWindowFrames - 2 * kContextFrames; + std::vector logits; + for (int f0 = 0;; f0 += hop) { + const int64_t s0 = static_cast(f0) * stride; + const int64_t len = std::min(win_samples, static_cast(n) - s0); + int wf = 0; + if (!run_window(m, x.data() + s0, len, logits, wf, err)) return false; + const bool last = s0 + len >= static_cast(n); + const int keep_from = f0 == 0 ? 0 : kContextFrames; + const int keep_to = last ? wf : std::min(wf, kWindowFrames - kContextFrames); + for (int i = keep_from; i < keep_to && f0 + i < frames; ++i) { + const float* row = logits.data() + static_cast(i) * V; + float mx = row[0]; + for (int v = 1; v < V; ++v) mx = std::max(mx, row[v]); + double sum = 0.0; + for (int v = 0; v < V; ++v) sum += std::exp(row[v] - mx); + const float lse = mx + static_cast(std::log(sum)); + float* dst = out.data() + static_cast(f0 + i) * V; + for (int v = 0; v < V; ++v) dst[v] = row[v] - lse; + } + if (last) break; + } + return true; +} diff --git a/electron/native/whisper-stt/src/ctc_aligner.h b/electron/native/whisper-stt/src/ctc_aligner.h new file mode 100644 index 000000000..80997505f --- /dev/null +++ b/electron/native/whisper-stt/src/ctc_aligner.h @@ -0,0 +1,37 @@ +// CTC acoustic model (wav2vec2 fine-tuned for CTC) run on ggml, for the word +// aligner's second pass. The helper only computes the per-frame log-probs; the +// forced alignment of whisper's words over them lives on the Node side +// (electron/stt/ctcAlign.ts). See transcription-and-captions.md § Word-level alignment. +#pragma once + +#include +#include +#include + +struct CtcModel; + +struct CtcModelInfo { + std::vector vocab; // token id -> text; "|" is the word delimiter + std::vector languages; // what the model was fine-tuned on + int blank = 0; + int stride = 320; // samples per output frame (20 ms at 16 kHz) + int receptive_field = 400; // samples the first frame sees +}; + +struct CtcModelDeleter { void operator()(CtcModel* m) const; }; +using CtcModelPtr = std::unique_ptr; + +// Loads a GGUF written by scripts/convert-wav2vec2-gguf.mjs. On the GPU when +// `use_gpu` and a GPU device is registered, else on the CPU. Null + `err` on failure. +CtcModelPtr ctc_load(const std::string& path, bool use_gpu, int threads, std::string& err); + +const CtcModelInfo& ctc_info(const CtcModel& model); + +// Name of the device the weights live on ("Vulkan0", "MTL0", "CPU"). +std::string ctc_device(const CtcModel& model); + +// Log-softmax emissions for 16 kHz mono `pcm`: `frames` rows of vocab-size +// floats. Frame i sees samples [i * stride, i * stride + receptive_field). +// Long inputs run in overlapping windows, so memory stays bounded. +bool ctc_emissions(CtcModel& model, const float* pcm, size_t n, std::vector& out, + int& frames, std::string& err); diff --git a/electron/native/whisper-stt/src/main.cpp b/electron/native/whisper-stt/src/main.cpp index 037230a57..77440e54d 100644 --- a/electron/native/whisper-stt/src/main.cpp +++ b/electron/native/whisper-stt/src/main.cpp @@ -36,6 +36,7 @@ // this is a belt-and-braces guarantee against a future bug or parallel invoker). #include "whisper.h" +#include "ctc_aligner.h" #include #include @@ -259,6 +260,21 @@ double to_original_sec(int64_t cs, const std::vector& kept) { return (kept.back().from + kept.back().len) / 16000.0; } +std::string base64(const void* data, size_t n) { + static const char* abc = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"; + const auto* p = static_cast(data); + std::string out; + out.reserve((n + 2) / 3 * 4); + for (size_t i = 0; i < n; i += 3) { + const uint32_t v = (p[i] << 16) | (i + 1 < n ? p[i + 1] << 8 : 0) | (i + 2 < n ? p[i + 2] : 0); + out += abc[(v >> 18) & 63]; + out += abc[(v >> 12) & 63]; + out += i + 1 < n ? abc[(v >> 6) & 63] : '='; + out += i + 2 < n ? abc[v & 63] : '='; + } + return out; +} + } // namespace int main(int argc, char** argv) { @@ -376,13 +392,15 @@ int main(int argc, char** argv) { res.set_content(probe.dump(), "application/json"); }); - // POST /inference — multipart form with `file` (WAV) + `language` + `response_format`. - svr.Post("/inference", [&](const httplib::Request& req, httplib::Response& res) { + // The upload of /inference and /emissions: a 16 kHz mono PCM16 WAV in the + // multipart field `file`. False after answering 400. + const auto read_upload = [](const httplib::Request& req, httplib::Response& res, + std::vector& pcm) -> bool { auto it = req.files.find("file"); if (it == req.files.end()) { res.status = 400; res.set_content(R"({"error":"missing 'file' form field"})", "application/json"); - return; + return false; } const auto& file_entry = it->second; @@ -404,7 +422,6 @@ int main(int argc, char** argv) { out.write(file_entry.content.data(), static_cast(file_entry.content.size())); } - std::vector pcm; int sample_rate = 0, channels = 0; const bool ok = read_wav_pcm16(tmp_wav, pcm, sample_rate, channels); std::error_code ec; @@ -412,15 +429,22 @@ int main(int argc, char** argv) { if (!ok) { res.status = 400; res.set_content(R"({"error":"failed to parse WAV"})", "application/json"); - return; + return false; } if (sample_rate != 16000 || channels != 1) { res.status = 400; res.set_content( R"({"error":"expected 16 kHz mono PCM16 WAV"})", "application/json"); - return; + return false; } + return true; + }; + + // POST /inference — multipart form with `file` (WAV) + `language` + `response_format`. + svr.Post("/inference", [&](const httplib::Request& req, httplib::Response& res) { + std::vector pcm; + if (!read_upload(req, res, pcm)) return; // language param std::string language = "auto"; @@ -669,6 +693,77 @@ int main(int argc, char** argv) { res.set_content(reply.dump(), "application/json"); }); + // POST /emissions — the CTC aligner's acoustic pass (issue #948, phase 3). + // Multipart form: `file` (the same WAV as /inference), `model` (path of a + // wav2vec2 GGUF, see ctc_aligner.h) and `regions` (JSON [[start_s, end_s], ...]). + // Answers the model's vocabulary and, per region, base64 float32 log-probs + // [frames x vocab]; frame i of a region sees the audio from + // `start + i * stride_s` for `receptive_s`. The forced alignment itself runs + // on the Node side (electron/stt/ctcAlign.ts). The model stays loaded until a + // request names another one. + CtcModelPtr aligner; + std::string aligner_path; + svr.Post("/emissions", [&](const httplib::Request& req, httplib::Response& res) { + std::vector pcm; + if (!read_upload(req, res, pcm)) return; + const std::string model = req.get_file_value("model").content; + nlohmann::json regions = nlohmann::json::parse(req.get_file_value("regions").content, nullptr, false); + if (model.empty() || !regions.is_array()) { + res.status = 400; + res.set_content(R"({"error":"need 'model' and a JSON 'regions' array"})", "application/json"); + return; + } + const std::lock_guard lk(infer_mu); + const auto t0 = std::chrono::steady_clock::now(); + if (!aligner || aligner_path != model) { + aligner.reset(); + std::string err; + aligner = ctc_load(model, cparams.use_gpu, threads, err); + if (!aligner) { + log("aligner: " + err); + res.status = 500; + res.set_content(nlohmann::json{{"error", "aligner: " + err}}.dump(), "application/json"); + return; + } + aligner_path = model; + log("aligner loaded on " + ctc_device(*aligner) + ": " + model); + } + const CtcModelInfo& info = ctc_info(*aligner); + nlohmann::json out_regions = nlohmann::json::array(); + const int64_t n_pcm = static_cast(pcm.size()); + for (const auto& r : regions) { + if (!r.is_array() || r.size() != 2 || !r[0].is_number() || !r[1].is_number()) continue; + const int64_t from = std::clamp(std::llround(r[0].get() * 16000.0), 0, n_pcm); + const int64_t to = std::clamp(std::llround(r[1].get() * 16000.0), from, n_pcm); + std::vector lp; + int frames = 0; + std::string err; + if (!ctc_emissions(*aligner, pcm.data() + from, static_cast(to - from), lp, frames, err)) { + log("aligner: " + err); + res.status = 500; + res.set_content(nlohmann::json{{"error", "aligner: " + err}}.dump(), "application/json"); + return; + } + out_regions.push_back({ + {"start", from / 16000.0}, + {"frames", frames}, + {"logprobs", base64(lp.data(), lp.size() * sizeof(float))}, + }); + } + const double elapsed_s = std::chrono::duration(std::chrono::steady_clock::now() - t0).count(); + nlohmann::json reply = { + {"vocab", info.vocab}, + {"blank", info.blank}, + {"languages", info.languages}, + {"stride_s", info.stride / 16000.0}, + {"receptive_s", info.receptive_field / 16000.0}, + {"device", ctc_device(*aligner)}, + {"elapsed_s", elapsed_s}, + {"regions", std::move(out_regions)}, + }; + res.set_content(reply.dump(), "application/json"); + }); + // ---- bind + listen ---- int bound_port = port; if (bound_port == 0) { diff --git a/electron/stt/ctcAlign.test.ts b/electron/stt/ctcAlign.test.ts new file mode 100644 index 000000000..91cc15d4c --- /dev/null +++ b/electron/stt/ctcAlign.test.ts @@ -0,0 +1,223 @@ +import { describe, expect, it } from "vitest"; +import { + type AlignableWord, + alignWordsOnEmissions, + type CtcEmissions, + ctcViterbi, + END_OFFSET_SEC, + emissionRegions, + PAUSE_END_OFFSET_SEC, + PAUSE_START_OFFSET_SEC, + parseEmissions, + START_OFFSET_SEC, + wordTokens, +} from "./ctcAlign"; + +const VOCAB = ["", "|", "a", "b", "c", "é", "'"]; +const index = new Map(VOCAB.map((t, i) => [t, i])); +const id = (t: string) => index.get(t) as number; + +/** Log-probs where frame t is sure of `tokens[t]` (a vocab entry, "" for blank). */ +function logprobs(tokens: string[]): Float32Array { + const V = VOCAB.length; + const out = new Float32Array(tokens.length * V).fill(-12); + tokens.forEach((t, f) => { + out[f * V + (t ? id(t) : 0)] = -0.01; + }); + return out; +} + +const STRIDE = 0.02; +const emissions = (tokens: string[], startSec = 0): CtcEmissions => ({ + vocab: VOCAB, + blank: 0, + strideSec: STRIDE, + regions: [{ startSec, frames: tokens.length, logprobs: logprobs(tokens) }], +}); +/** `n` blank frames. */ +const gap = (n: number) => Array(n).fill(""); +const word = ( + text: string, + startSec: number, + endSec: number, + anchorSec = startSec, +): AlignableWord => ({ + word: text, + startSec, + endSec, + anchorSec, +}); +const r = (x: number) => Math.round(x * 1000) / 1000; + +describe("wordTokens", () => { + it("spells a word in the vocabulary's case, dropping silent punctuation", () => { + const upper = new Map(["", "|", "A", "B", "C", "'"].map((t, i) => [t, i])); + expect(wordTokens("Cab,", upper)).toEqual([4, 2, 3]); + expect(wordTokens("ab’c", upper)).toEqual([2, 3, 5, 4]); + expect(wordTokens("cab", index)).toEqual([id("c"), id("a"), id("b")]); + }); + + it("keeps an accented letter the vocabulary has and strips the accent otherwise", () => { + expect(wordTokens("É", index)).toEqual([id("é")]); + expect(wordTokens("àbc", index)).toEqual([id("a"), id("b"), id("c")]); + }); + + it("gives nothing to say for punctuation, and no spelling for digits and symbols", () => { + expect(wordTokens("?", index)).toEqual([]); + expect(wordTokens("...", index)).toEqual([]); + expect(wordTokens("1080p", index)).toBeNull(); + expect(wordTokens("€", index)).toBeNull(); + expect(wordTokens("日本", index)).toBeNull(); + }); +}); + +describe("ctcViterbi", () => { + it("finds the frames each label is spoken on", () => { + const tokens = ["", "a", "a", "", "b", ""]; + expect(ctcViterbi(logprobs(tokens), 6, VOCAB.length, 0, [id("a"), id("b")])).toEqual([ + [1, 2], + [4, 4], + ]); + }); + + it("needs a blank between two repeats of a label, and fails when the frames are too few", () => { + const aba = logprobs(["a", "", "a"]); + expect(ctcViterbi(aba, 3, VOCAB.length, 0, [id("a"), id("a")])).toEqual([ + [0, 0], + [2, 2], + ]); + expect(ctcViterbi(logprobs(["a", "a"]), 2, VOCAB.length, 0, [id("a"), id("a")])).toBeNull(); + }); + + it("scores a wildcard as whatever is spoken", () => { + const spans = ctcViterbi(logprobs(["a", "", "c", "c", "", "b"]), 6, VOCAB.length, 0, [ + id("a"), + -1, + id("b"), + ]); + expect(spans?.[1]).toEqual([2, 3]); + }); +}); + +describe("emissionRegions", () => { + it("pads each stretch of speech and merges the ones that touch", () => { + expect( + emissionRegions( + [ + { startSec: 0.1, endSec: 1 }, + { startSec: 1.4, endSec: 2 }, + { startSec: 5, endSec: 6 }, + ], + 6.1, + ).map(([a, b]) => [r(a), r(b)]), + ).toEqual([ + [0, 2.3], + [4.7, 6.1], + ]); + }); +}); + +describe("alignWordsOnEmissions", () => { + const speech = [{ startSec: 0.2, endSec: 1 }]; + + it("times words on their letters, calibrated outward, meeting mid-way inside speech", () => { + // "ab" on frames 20-21, the delimiter on 22, "c" on 24: continuous speech. + const tokens = [...gap(20), "a", "b", "|", "", "c", ...gap(40)]; + const out = alignWordsOnEmissions( + [word("ab", 0.3, 0.5, 0.45), word("c", 0.5, 0.7, 0.6)], + speech, + emissions(tokens), + ); + const abEnd = 22 * STRIDE + END_OFFSET_SEC; + const cStart = 24 * STRIDE + START_OFFSET_SEC; + const boundary = abEnd > cStart ? (abEnd + cStart) / 2 : null; + expect(r(out[0].startSec)).toBe(r(20 * STRIDE + PAUSE_START_OFFSET_SEC)); + expect(r(out[0].endSec)).toBe(r(boundary ?? abEnd)); + expect(r(out[1].startSec)).toBe(r(boundary ?? cStart)); + expect(r(out[1].endSec)).toBe(r(25 * STRIDE + PAUSE_END_OFFSET_SEC)); + expect(out[0].anchorSec).toBe(0.45); + }); + + it("keeps a pause between two words as a pause", () => { + // "a" on frame 12, "c" on frame 40: 0.54 s apart. + const tokens = [...gap(12), "a", ...gap(27), "c", ...gap(40)]; + const out = alignWordsOnEmissions( + [word("a", 0.2, 0.4, 0.3), word("c", 0.4, 0.9, 0.85)], + speech, + emissions(tokens), + ); + expect(r(out[0].endSec)).toBe(r(13 * STRIDE + PAUSE_END_OFFSET_SEC)); + expect(r(out[1].startSec)).toBe(r(40 * STRIDE + PAUSE_START_OFFSET_SEC)); + }); + + it("gives a word it cannot spell the speech between its neighbours", () => { + const tokens = [...gap(15), "a", "|", "c", "b", "c", "|", "b", ...gap(40)]; + const out = alignWordsOnEmissions( + [word("a", 0.3, 0.32, 0.31), word("42", 0.32, 0.4, 0.35), word("b", 0.4, 0.5, 0.42)], + speech, + emissions(tokens), + ); + expect(out[1].startSec).toBeGreaterThanOrEqual(out[0].endSec); + expect(out[1].endSec).toBeLessThanOrEqual(out[2].startSec); + expect(r(out[1].startSec)).toBeLessThan(0.34); + expect(r(out[2].startSec)).toBeLessThan(0.42); + }); + + it("puts punctuation on the end of the word before it", () => { + const tokens = [...gap(15), "a", "b", ...gap(40)]; + const out = alignWordsOnEmissions( + [word("ab", 0.3, 0.5, 0.35), word("!", 0.5, 0.6, 0.55)], + speech, + emissions(tokens), + ); + expect(out[1].startSec).toBe(out[0].endSec); + expect(out[1].endSec).toBe(out[0].endSec); + }); + + it("aligns a stretch that runs to the end of the upload, past the last frame", () => { + // 1 s of audio gives 49 frames (0.98 s): the stretch ends after them. + const tokens = [...gap(30), "a", "b", ...gap(17)]; + const out = alignWordsOnEmissions( + [word("ab", 0.65, 1, 0.7)], + [{ startSec: 0.5, endSec: 1 }], + emissions(tokens), + ); + expect(r(out[0].startSec)).toBe(r(30 * STRIDE + PAUSE_START_OFFSET_SEC)); + }); + + it("keeps the helper's times where it cannot align", () => { + const words = [word("ab", 0.3, 0.5, 0.35), word("c", 1.5, 1.7, 1.6)]; + // The second word is outside every stretch; the first one's letters do not fit. + const out = alignWordsOnEmissions(words, speech, emissions(gap(2))); + expect(out).toEqual(words); + // No region covers the stretch at all. + expect(alignWordsOnEmissions(words, speech, emissions(gap(60), 5))).toEqual(words); + }); +}); + +describe("parseEmissions", () => { + const floats = new Float32Array([-0.1, -2, -3, -4, -5, -6, -7]); + const logprobs = Buffer.from(floats.buffer).toString("base64"); + + it("decodes the helper's base64 log-probs", () => { + const parsed = parseEmissions({ + vocab: VOCAB, + blank: 0, + stride_s: 0.02, + regions: [{ start: 1.5, frames: 1, logprobs }], + }); + expect(parsed?.regions[0].startSec).toBe(1.5); + expect([...(parsed?.regions[0].logprobs ?? [])]).toEqual([...floats]); + }); + + it("refuses a reply whose sizes do not add up", () => { + expect( + parseEmissions({ + vocab: VOCAB, + blank: 0, + stride_s: 0.02, + regions: [{ start: 0, frames: 2, logprobs }], + }), + ).toBeNull(); + }); +}); diff --git a/electron/stt/ctcAlign.ts b/electron/stt/ctcAlign.ts new file mode 100644 index 000000000..335b6c196 --- /dev/null +++ b/electron/stt/ctcAlign.ts @@ -0,0 +1,294 @@ +// Second alignment pass: re-times whisper's words on a CTC acoustic model +// (issue #948, phase 3). +// +// whisper's DTW times words within ~30 ms median inside continuous speech, but +// its 20 ms frames and diffuse cross-attention cap it there (P90 ~110 ms). A +// wav2vec2 model fine-tuned for CTC scores every 20 ms frame against every +// letter, so a Viterbi pass that forces whisper's own words through those +// scores puts each letter where the audio has it. The helper computes the +// scores (`/emissions`, electron/native/whisper-stt/src/ctc_aligner.cpp); this +// file spells whisper's words in the model's letters, aligns them, and maps the +// letters back to word times. +// +// Words the model cannot spell (digits, symbols, another script) become one +// wildcard token that matches any speech, so they hold their place and their +// neighbours stay aligned. A stretch the model cannot fit keeps the times it +// came with, and so does every word when there is no aligner for the language. +// +// Self-contained on purpose: tools/stt-eval/word-timing imports this file +// straight into Node, which resolves no extensionless import. + +import type { SttVadSegment } from "./transcriptionContract"; + +/** What `/emissions` answers, decoded. */ +export interface CtcEmissions { + vocab: string[]; + blank: number; + /** Seconds between two frames. */ + strideSec: number; + /** Per requested region: its first frame's time and `frames * vocab` log-probs. */ + regions: Array<{ startSec: number; frames: number; logprobs: Float32Array }>; +} + +/** The fields of a helper word this pass reads and rewrites. */ +export interface AlignableWord { + word: string; + startSec: number; + endSec: number; + /** Always inside the word; decides which stretch of speech owns it. */ + anchorSec: number; +} + +// Calibration, fitted on tools/stt-eval/word-timing (TTS and LibriSpeech agree). +// The model is sure of a letter a little after the sound starts and before it +// ends, so CTC shrinks words: the edges move outward. Inside continuous speech +// the word boundary lies between the last letter of one word and the first of +// the next; where the two estimates cross, the boundary is their midpoint. +/** Added to the frame of a word's first letter. */ +export const START_OFFSET_SEC = -0.045; +/** Added to the end of the frame of a word's last letter. */ +export const END_OFFSET_SEC = 0.025; +/** A gap between two words' letters this long is a pause, not a boundary. */ +export const PAUSE_SEC = 0.1; +/** After a pause, a word's first letter trails its onset further (breath, closure, soft attack). */ +export const PAUSE_START_OFFSET_SEC = -0.06; +/** Before a pause, a word decays past its last letter; early there is audible, late is silence. */ +export const PAUSE_END_OFFSET_SEC = 0.15; + +/** Audio kept on each side of a stretch of speech, so its first and last letters have context. */ +export const REGION_MARGIN_SEC = 0.3; + +/** As in snapWordBoundaries.ts: a stretch owns the words anchored before its end plus this tail. */ +const TAIL_SEC = 0.1; + +const WILDCARD = -1; + +/** Seconds of the upload the helper should score: the speech plus a margin, merged. */ +export function emissionRegions( + speech: SttVadSegment[], + durationSec: number, +): Array<[number, number]> { + const out: Array<[number, number]> = []; + for (const s of speech) { + const a = Math.max(0, s.startSec - REGION_MARGIN_SEC); + const b = Math.min(durationSec, s.endSec + REGION_MARGIN_SEC); + const last = out[out.length - 1]; + if (last && a <= last[1]) last[1] = Math.max(last[1], b); + else if (b > a) out.push([a, b]); + } + return out; +} + +/** + * A word as the model's token ids: `[]` when there is nothing to say + * (punctuation), `null` when a letter is not in the vocabulary (digits, symbols, + * another script). Case follows the vocabulary, an accented letter the + * vocabulary lacks falls back to its base letter, and the apostrophe is + * normalised. + */ +export function wordTokens(word: string, index: Map): number[] | null { + const upper = index.has("A") && !index.has("a"); + let text = word.normalize("NFC").replace(/[’‘`´]/g, "'"); + text = upper ? text.toUpperCase() : text.toLowerCase(); + const out: number[] = []; + for (const ch of text) { + const id = index.get(ch); + if (id !== undefined) { + out.push(id); + continue; + } + if (/\p{S}/u.test(ch)) return null; // "%", "+", "€": said aloud, but not spelled so + if (/[\p{P}\s]/u.test(ch)) continue; // silent, apostrophes and hyphens included + const base = [...ch.normalize("NFD").replace(/\p{M}/gu, "")]; + const ids = base.map((c) => index.get(c)); + if (ids.length === 0 || ids.some((x) => x === undefined)) return null; + out.push(...(ids as number[])); + } + return out; +} + +/** + * CTC Viterbi forced alignment of `labels` over `T` frames of `V` log-probs. + * `WILDCARD` scores as the frame's best non-blank token. Returns, per label, the + * first and last frame the best path spends on it, or null when the labels + * cannot fit in the frames. + */ +export function ctcViterbi( + logprobs: Float32Array, + T: number, + V: number, + blank: number, + labels: number[], +): Array<[number, number]> | null { + const L = labels.length; + const S = 2 * L + 1; + if (L === 0 || T === 0) return null; + const wild = new Float32Array(T); + if (labels.includes(WILDCARD)) { + for (let t = 0; t < T; t++) { + let best = Number.NEGATIVE_INFINITY; + for (let v = 0; v < V; v++) if (v !== blank) best = Math.max(best, logprobs[t * V + v]); + wild[t] = best; + } + } + const lab = (s: number) => (s % 2 === 0 ? blank : labels[(s - 1) >> 1]); + const em = (t: number, s: number) => { + const l = lab(s); + return l === WILDCARD ? wild[t] : logprobs[t * V + l]; + }; + const NEG = Number.NEGATIVE_INFINITY; + let prev = new Float64Array(S).fill(NEG); + let cur = new Float64Array(S); + const back = new Uint8Array(T * S); // how far back the best predecessor is: 0, 1 or 2 + prev[0] = em(0, 0); + prev[1] = em(0, 1); + for (let t = 1; t < T; t++) { + for (let s = 0; s < S; s++) { + let best = prev[s]; + let from = 0; + if (s >= 1 && prev[s - 1] > best) { + best = prev[s - 1]; + from = 1; + } + // A label may follow the one before the blank directly, unless it repeats it. + if (s >= 2 && s % 2 === 1 && lab(s) !== lab(s - 2) && prev[s - 2] > best) { + best = prev[s - 2]; + from = 2; + } + cur[s] = best === NEG ? NEG : best + em(t, s); + back[t * S + s] = from; + } + [prev, cur] = [cur, prev]; + } + let s = prev[S - 2] > prev[S - 1] ? S - 2 : S - 1; + if (prev[s] === NEG) return null; + const spans: Array<[number, number]> = labels.map(() => [-1, -1]); + for (let t = T - 1; t >= 0; t--) { + if (s % 2 === 1) { + const k = (s - 1) >> 1; + spans[k][0] = t; + if (spans[k][1] < 0) spans[k][1] = t; + } + s -= back[t * S + s]; + } + return spans; +} + +/** + * Re-time `words` on `emissions`. Each speech stretch owns the words anchored + * in it (the rule of snapWordBoundaries.ts); they are aligned on the frames + * between the neighbouring stretches, with the vocabulary's word delimiter + * between words when it has one. Words outside every stretch, stretches the + * model cannot fit, and a vocabulary with no word in it keep the helper's + * times. Punctuation collapses onto the end of the word before it. + */ +export function alignWordsOnEmissions( + words: W[], + speech: SttVadSegment[], + emissions: CtcEmissions, +): W[] { + const out = words.map((w) => ({ ...w })); + const index = new Map(emissions.vocab.map((t, i) => [t, i])); + const delimiter = index.get("|"); + const V = emissions.vocab.length; + const { strideSec } = emissions; + let k = 0; + for (let i = 0; i < speech.length; i++) { + const { startSec: onset, endSec: offset } = speech[i]; + const until = Math.min(offset + TAIL_SEC, speech[i + 1]?.startSec ?? Number.POSITIVE_INFINITY); + const owned: number[] = []; + for (; k < out.length && words[k].anchorSec < until; k++) owned.push(k); + // A region's frames stop up to one receptive field (25 ms) short of the audio + // it was cut from, so a stretch that runs to the end of the upload ends + // past them; two frames of slack keep it. `hi` below stays on the frames. + const region = emissions.regions.find( + (r) => r.startSec <= onset && r.startSec + (r.frames + 2) * strideSec >= offset, + ); + if (!region || owned.length === 0) continue; + // The frames of this stretch: its margin, stopping at the neighbours' speech. + const lo = Math.max( + region.startSec, + (speech[i - 1]?.endSec ?? Number.NEGATIVE_INFINITY) + TAIL_SEC, + onset - REGION_MARGIN_SEC, + ); + const hi = Math.min( + region.startSec + region.frames * strideSec, + speech[i + 1]?.startSec ?? Number.POSITIVE_INFINITY, + offset + REGION_MARGIN_SEC, + ); + const f0 = Math.max(0, Math.ceil((lo - region.startSec) / strideSec)); + const f1 = Math.min(region.frames, Math.floor((hi - region.startSec) / strideSec)); + if (f1 <= f0) continue; + + const labels: number[] = []; + const timed: Array<{ j: number; first: number; last: number }> = []; + for (const j of owned) { + const toks = wordTokens(out[j].word, index); + if (toks?.length === 0) continue; + if (labels.length > 0 && delimiter !== undefined) labels.push(delimiter); + timed.push({ j, first: labels.length, last: labels.length + (toks?.length ?? 1) - 1 }); + labels.push(...(toks ?? [WILDCARD])); + } + const spans = ctcViterbi( + region.logprobs.subarray(f0 * V, f1 * V), + f1 - f0, + V, + emissions.blank, + labels, + ); + if (!spans) continue; + const at = (f: number) => region.startSec + (f0 + f) * strideSec; + const s = timed.map((w) => at(spans[w.first][0])); + const e = timed.map((w) => at(spans[w.last][1] + 1)); + timed.forEach(({ j }, q) => { + const prevEnd = e[q - 1] ?? Number.NEGATIVE_INFINITY; + const nextStart = s[q + 1] ?? Number.POSITIVE_INFINITY; + out[j].startSec = + s[q] - prevEnd >= PAUSE_SEC + ? Math.max(s[q] + PAUSE_START_OFFSET_SEC, (prevEnd + s[q]) / 2) + : s[q] + START_OFFSET_SEC; + out[j].endSec = + nextStart - e[q] >= PAUSE_SEC + ? Math.min(e[q] + PAUSE_END_OFFSET_SEC, (e[q] + nextStart) / 2) + : e[q] + END_OFFSET_SEC; + const before = q > 0 ? out[timed[q - 1].j] : null; + if (before && before.endSec > out[j].startSec) { + const mid = (before.endSec + out[j].startSec) / 2; + before.endSec = mid; + out[j].startSec = mid; + } + }); + for (let n = 1; n < owned.length; n++) { + const j = owned[n]; + if (timed.some((w) => w.j === j)) continue; + out[j].startSec = out[owned[n - 1]].endSec; + out[j].endSec = out[j].startSec; + } + } + return out; +} + +/** The helper's `/emissions` JSON. */ +export interface EmissionsJson { + vocab: string[]; + blank: number; + stride_s: number; + regions: Array<{ start: number; frames: number; logprobs: string }>; + /** Seconds the helper spent, model load included. */ + elapsed_s?: number; +} + +/** Decodes the helper's `/emissions` JSON; null when it is not one. */ +export function parseEmissions(json: EmissionsJson): CtcEmissions | null { + if (!Array.isArray(json?.vocab) || !Array.isArray(json.regions) || !(json.stride_s > 0)) + return null; + const V = json.vocab.length; + const regions: CtcEmissions["regions"] = []; + for (const r of json.regions) { + // Copied: a pooled Buffer can start on an offset Float32Array refuses. + const bytes = Uint8Array.from(Buffer.from(r.logprobs, "base64")); + if (bytes.byteLength !== r.frames * V * 4) return null; + regions.push({ startSec: r.start, frames: r.frames, logprobs: new Float32Array(bytes.buffer) }); + } + return { vocab: json.vocab, blank: json.blank, strideSec: json.stride_s, regions }; +} diff --git a/electron/stt/index.test.ts b/electron/stt/index.test.ts index 69ef1723b..c5af87fe9 100644 --- a/electron/stt/index.test.ts +++ b/electron/stt/index.test.ts @@ -34,6 +34,9 @@ vi.mock("./whisperServer", () => { vi.mock("./modelManager", () => ({ ensureModels: vi.fn(async () => undefined), + ensureAligner: vi.fn(async ({ language }: { language: string }) => `/fake/${language}.gguf`), + cachedAligner: vi.fn(async () => null), + alignerPath: vi.fn((base: string, language: string): string | null => `${base}/${language}.gguf`), modelPaths: (base: string) => ({ whisper: `${base}/whisper-ggml/ggml-small-q8_0.bin`, }), @@ -470,6 +473,110 @@ describe("SttManager", () => { expect(fakeWhisperServer.start).toHaveBeenCalledOnce(); }); + describe("word aligner", () => { + /** Chunks report what `alignerFor("fr")` gave them, as the helper would use it. */ + function recordAligner() { + const asked: Array = []; + fakeWhisperServer.transcribe.mockImplementation( + async (opts: { alignerFor?: (language: string) => Promise }) => { + asked.push((await opts.alignerFor?.("fr")) ?? null); + return { + segments: [], + wordSegments: [], + detectedLanguage: "fr", + backend: "whispercpp-cpu" as const, + }; + }, + ); + return asked; + } + const long = new Float32Array(16_000 * 400); + + async function mocks() { + const mm = await import("./modelManager"); + const ensure = vi.mocked(mm.ensureAligner); + const cached = vi.mocked(mm.cachedAligner); + ensure.mockReset(); + cached.mockReset(); + cached.mockResolvedValue(null); + return { ensure, cached }; + } + + it("never holds a chunk on the download, and uses the aligner once it lands", async () => { + const { ensure } = await mocks(); + let land: (file: string) => void = () => undefined; + ensure.mockImplementation(() => new Promise((resolve) => (land = resolve))); + const asked = recordAligner(); + const mgr = new SttManager(); + await mgr.init({ modelsBaseDir: "/tmp/fake-stt-models" }); + + // The download never finishes during this run: every chunk goes on without it. + await mgr.transcribe({ samples: long, language: "fr" }); + expect(asked.length).toBeGreaterThan(1); + expect(asked.every((x) => x === null)).toBe(true); + expect(ensure).toHaveBeenCalledOnce(); + + land("/fake/fr.gguf"); + await new Promise((resolve) => setTimeout(resolve, 0)); + asked.length = 0; + await mgr.transcribe({ samples: new Float32Array(16_000), language: "fr" }); + expect(asked).toEqual(["/fake/fr.gguf"]); + expect(ensure).toHaveBeenCalledOnce(); + }); + + it("uses a verified copy already on disk without downloading", async () => { + const { ensure, cached } = await mocks(); + cached.mockResolvedValue("/fake/fr.gguf"); + const asked = recordAligner(); + const mgr = new SttManager(); + await mgr.init({ modelsBaseDir: "/tmp/fake-stt-models" }); + await mgr.transcribe({ samples: long, language: "fr" }); + expect(asked.every((x) => x === "/fake/fr.gguf")).toBe(true); + expect(cached).toHaveBeenCalledOnce(); + expect(ensure).not.toHaveBeenCalled(); + }); + + it("aborts the download on cancel and on shutdown, then retries on a later run", async () => { + const { ensure } = await mocks(); + const signals: AbortSignal[] = []; + ensure.mockImplementation( + ({ signal }: { signal?: AbortSignal }) => + new Promise((_, reject) => { + if (signal) signals.push(signal); + signal?.addEventListener("abort", () => reject(signal.reason)); + }), + ); + recordAligner(); + const mgr = new SttManager(); + await mgr.init({ modelsBaseDir: "/tmp/fake-stt-models" }); + await mgr.transcribe({ samples: long, language: "fr" }); + expect(signals).toHaveLength(1); + + mgr.cancel(); + expect(signals[0].aborted).toBe(true); + await new Promise((resolve) => setTimeout(resolve, 0)); + + // Not cached as failed for good: the next transcription starts it again. + await mgr.transcribe({ samples: long, language: "fr" }); + expect(signals).toHaveLength(2); + expect(signals[1].aborted).toBe(false); + await mgr.shutdown(); + expect(signals[1].aborted).toBe(true); + }); + + it("gives a language without an aligner nothing, and downloads nothing", async () => { + const { ensure } = await mocks(); + const mm = await import("./modelManager"); + vi.mocked(mm.alignerPath).mockReturnValueOnce(null); + const asked = recordAligner(); + const mgr = new SttManager(); + await mgr.init({ modelsBaseDir: "/tmp/fake-stt-models" }); + await mgr.transcribe({ samples: new Float32Array(16_000), language: "fr" }); + expect(asked).toEqual([null]); + expect(ensure).not.toHaveBeenCalled(); + }); + }); + it("fans status out to every sink, and detaching one leaves the others", async () => { const mgr = new SttManager(); const a = vi.fn<(e: SttStatusEvent) => void>(); diff --git a/electron/stt/index.ts b/electron/stt/index.ts index 4ee2b24c9..11c760f0b 100644 --- a/electron/stt/index.ts +++ b/electron/stt/index.ts @@ -2,7 +2,13 @@ import path from "node:path"; import { app, type IpcMain } from "electron"; import { planChunks } from "./chunking"; import { extractMono16kPcm } from "./extractAudio"; -import { ensureModels, modelPaths } from "./modelManager"; +import { + alignerPath, + cachedAligner, + ensureAligner, + ensureModels, + modelPaths, +} from "./modelManager"; import type { SttPhraseSegment, SttStatusEvent, @@ -105,6 +111,60 @@ export class SttManager { */ private cancelEpoch = 0; + /** + * The word aligner per language (ctcAlign.ts). A chunk never waits for a + * download: it is an improvement, and a 350 MB fetch inside a chunk would hold + * the helper's queue for as long as the network takes, past any request + * ceiling and out of Cancel's reach. So: + * - a copy already on disk is verified in place the first time a language is + * asked for in a session (a local read, no network), then reused; + * - otherwise the download starts in the background and the chunks keep + * whisper's times until it lands; later chunks and transcriptions use it; + * - Cancel and quit abort it, and so does a stall (`ensureAligner`). A failed + * one is not retried for the rest of the run, only on the next transcription. + */ + private readonly alignersReady = new Map(); + private readonly alignersChecked = new Set(); + private readonly alignerDownloads = new Map(); + private readonly alignersFailed = new Set(); + + private readonly alignerFor = async (language: string): Promise => { + const ready = this.alignersReady.get(language); + if (ready) return ready; + if (!alignerPath(this.getModelsDir(), language)) return null; + if (!this.alignersChecked.has(language)) { + this.alignersChecked.add(language); + const cached = await cachedAligner(this.getModelsDir(), language); + if (cached) { + this.alignersReady.set(language, cached); + return cached; + } + } + if (!this.alignersFailed.has(language) && !this.alignerDownloads.has(language)) { + this.downloadAligner(language); + } + return null; + }; + + private downloadAligner(language: string): void { + const controller = new AbortController(); + this.alignerDownloads.set(language, controller); + console.info(`[stt] downloading the "${language}" word aligner in the background`); + ensureAligner({ baseDir: this.getModelsDir(), language, signal: controller.signal }) + .then((file) => { + if (file) this.alignersReady.set(language, file); + console.info(`[stt] "${language}" word aligner ready`); + }) + .catch((error: unknown) => { + this.alignersFailed.add(language); + console.warn( + `[stt] no word aligner for "${language}", keeping whisper's word times: ` + + `${error instanceof Error ? error.message : String(error)}`, + ); + }) + .finally(() => this.alignerDownloads.delete(language)); + } + /** * The extraction in flight, if any. `cancelEpoch` alone stops the CHUNK loop, which is * checked between chunks — so a cancel during the decode left ffmpeg running to @@ -144,6 +204,7 @@ export class SttManager { cancel(): void { this.cancelEpoch++; this.extraction?.abort(); + for (const download of this.alignerDownloads.values()) download.abort(cancelledError()); } /** @@ -235,7 +296,7 @@ export class SttManager { for (let attempt = 1; attempt <= CHUNK_ATTEMPTS; attempt++) { if (this.shuttingDown) throw cancelledError(); try { - return await this.server.transcribe({ samples, language }); + return await this.server.transcribe({ samples, language, alignerFor: this.alignerFor }); } catch (error) { lastError = error; if (this.shuttingDown) throw cancelledError(); @@ -273,6 +334,7 @@ export class SttManager { await this.init(); const epoch = this.cancelEpoch; + this.alignersFailed.clear(); // Extraction is part of the run, and on a long file it is the part the user used // to watch the editor freeze through. Doing it here means the renderer hands over // a path and gets segments back, holding none of the audio. No new status phase: @@ -437,7 +499,7 @@ export class SttManager { async shutdown(): Promise { if (this.shuttingDown) return; this.shuttingDown = true; - this.cancelEpoch++; + this.cancel(); await this.server.shutdown(); } } diff --git a/electron/stt/modelManager.test.ts b/electron/stt/modelManager.test.ts index 441906234..dcee6a24f 100644 --- a/electron/stt/modelManager.test.ts +++ b/electron/stt/modelManager.test.ts @@ -4,7 +4,16 @@ import { mkdir, mkdtemp, readFile, rm, stat, writeFile } from "node:fs/promises" import { tmpdir } from "node:os"; import path from "node:path"; import { afterEach, beforeEach, describe, expect, it } from "vitest"; -import { areModelsPresent, ensureModels, modelPaths, STT_MODELS } from "./modelManager"; +import { + alignerPath, + areModelsPresent, + CTC_ALIGNERS, + cachedAligner, + ensureAligner, + ensureModels, + modelPaths, + STT_MODELS, +} from "./modelManager"; describe("modelManager", () => { let dir: string; @@ -254,4 +263,116 @@ describe("modelManager", () => { STT_MODELS.whisper.files[0].expectedSha256 = originalSha; } }); + + describe("CTC aligners", () => { + it("pins every aligner to a digest and names its upstream revision", () => { + for (const file of Object.values(CTC_ALIGNERS)) { + expect(file.expectedSha256).toMatch(/^[0-9a-f]{64}$/); + expect(file.source).toMatch(/@[0-9a-f]{40}$/); + // Hosted under a `v0.0.0-*` release tag, never a moving one. + expect(file.url).toMatch(/\/releases\/download\/v0\.0\.0-[^/]+\//); + } + }); + + it("has none for a language it does not cover, prototype keys included", async () => { + expect(alignerPath(dir, "de")).toBeNull(); + expect(alignerPath(dir, "constructor")).toBeNull(); + expect(await ensureAligner({ baseDir: dir, language: "de" })).toBeNull(); + }); + + it("downloads, verifies and returns the aligner of a covered language", async () => { + const bytes = Buffer.from("gguf weights"); + const original = CTC_ALIGNERS.fr.expectedSha256; + CTC_ALIGNERS.fr.expectedSha256 = createHash("sha256").update(bytes).digest("hex"); + try { + const file = await ensureAligner({ + baseDir: dir, + language: "fr", + fetcher: async () => new Response(bytes, { status: 200 }), + }); + expect(file).toBe(alignerPath(dir, "fr")); + expect(await readFile(file as string)).toEqual(bytes); + } finally { + CTC_ALIGNERS.fr.expectedSha256 = original; + } + }); + + it("refuses a download whose digest does not match", async () => { + await expect( + ensureAligner({ + baseDir: dir, + language: "en", + fetcher: async () => new Response("tampered", { status: 200 }), + }), + ).rejects.toThrow(/SHA-256 mismatch/); + expect(existsSync(alignerPath(dir, "en") as string)).toBe(false); + }); + + /** A body that sends a few bytes, then nothing, ever. */ + const stalledBody = () => + new Response( + new ReadableStream({ + start(controller) { + controller.enqueue(new Uint8Array(8)); + }, + }), + { status: 200 }, + ); + + it("aborts a download that stops sending, and leaves no partial file", async () => { + const file = alignerPath(dir, "fr") as string; + await expect( + ensureAligner({ + baseDir: dir, + language: "fr", + stallMs: 50, + fetcher: async () => stalledBody(), + }), + ).rejects.toThrow(/stalled/); + expect(existsSync(file)).toBe(false); + expect(existsSync(`${file}.partial`)).toBe(false); + }); + + it("stops when its signal aborts, mid-body or between attempts", async () => { + const controller = new AbortController(); + const download = ensureAligner({ + baseDir: dir, + language: "fr", + signal: controller.signal, + fetcher: async () => stalledBody(), + }); + setTimeout(() => controller.abort(new Error("cancelled")), 20); + await expect(download).rejects.toThrow("cancelled"); + + // A 503 puts it in a backoff of seconds: the abort must not wait for it. + const again = new AbortController(); + const started = Date.now(); + const retrying = ensureAligner({ + baseDir: dir, + language: "fr", + signal: again.signal, + fetcher: async () => new Response("busy", { status: 503 }), + }); + setTimeout(() => again.abort(new Error("quit")), 20); + await expect(retrying).rejects.toThrow("quit"); + expect(Date.now() - started).toBeLessThan(1000); + }); + + it("verifies a cached copy without the network", async () => { + expect(await cachedAligner(dir, "fr")).toBeNull(); + const file = alignerPath(dir, "fr") as string; + const bytes = Buffer.from("gguf weights"); + await mkdir(path.dirname(file), { recursive: true }); + await writeFile(file, bytes); + // Present but not the pinned bytes: not usable. + expect(await cachedAligner(dir, "fr")).toBeNull(); + const original = CTC_ALIGNERS.fr.expectedSha256; + CTC_ALIGNERS.fr.expectedSha256 = createHash("sha256").update(bytes).digest("hex"); + try { + expect(await cachedAligner(dir, "fr")).toBe(file); + } finally { + CTC_ALIGNERS.fr.expectedSha256 = original; + } + }); + }); }); diff --git a/electron/stt/modelManager.ts b/electron/stt/modelManager.ts index bc37a0359..d9d4aef32 100644 --- a/electron/stt/modelManager.ts +++ b/electron/stt/modelManager.ts @@ -22,7 +22,9 @@ import { pipeline } from "node:stream/promises"; * * Word timestamps come from whisper.cpp's native DTW token timestamps. The * Silero VAD model only decides which audio whisper decodes, and gives the - * speech edges phrases are anchored on. See `technical-documentation/architecture/transcription-and-captions.md`. + * speech edges phrases are anchored on. The CTC aligners (`CTC_ALIGNERS`) re-time + * the words for the languages that have one, and are fetched only when a + * transcription detects such a language. See `technical-documentation/architecture/transcription-and-captions.md`. */ export type SttModelId = "whisper" | "silero-vad"; @@ -94,6 +96,89 @@ export const STT_MODELS: Record = { }, }; +/** + * The CTC aligners that re-time whisper's words (electron/stt/ctcAlign.ts), per + * language whisper reports. A language missing here keeps whisper's DTW times. + * + * Each file is a wav2vec2 CTC model converted to GGUF by + * `scripts/convert-wav2vec2-gguf.mjs` (deterministic: re-running it on the source + * named in `source` reproduces the digest below), then published under a `v0.0.0-*` + * release tag like the other binaries that need a permanent URL but are not a + * product version (see scripts/fetch-onnxruntime.mjs). Apache-2.0, both of them. + */ +const ALIGNER_RELEASE = + "https://github.com/getopenscreen/openscreen/releases/download/v0.0.0-ctc-aligners-1"; + +export interface CtcAlignerFile extends SttModelFile { + /** HuggingFace repo and revision the file was converted from. */ + source: string; +} + +export const CTC_ALIGNERS: Record = { + en: { + name: "w2v-en-base-q8_0.gguf", + url: `${ALIGNER_RELEASE}/w2v-en-base-q8_0.gguf`, + expectedSha256: "b7f21a97208f368d3505bd9a7bc9ff3795b1169d8e3b036028082eb25eb464ae", + approximateBytes: 109_040_064, + source: "facebook/wav2vec2-base-960h@22aad52d435eb6dbaf354bdad9b0da84ce7d6156", + }, + fr: { + name: "w2v-fr-large-q8_0.gguf", + url: `${ALIGNER_RELEASE}/w2v-fr-large-q8_0.gguf`, + expectedSha256: "e3c284da3e27564db07ac226bf6402a4d7856b806455b3e0f5283f29f9495f48", + approximateBytes: 348_037_120, + // The repo's safetensors conversion PR: main only has pytorch_model.bin. + source: "jonatasgrosman/wav2vec2-large-xlsr-53-french@70db24a266633ffcc8edce4e72f3a5cb69d602d6", + }, +}; + +const ALIGNER_DIR = "ctc-aligner"; + +/** Where the aligner for `language` lives, or null when there is none for it. */ +export function alignerPath(baseDir: string, language: string): string | null { + const file = Object.keys(CTC_ALIGNERS).includes(language) ? CTC_ALIGNERS[language] : null; + return file ? path.join(baseDir, ALIGNER_DIR, file.name) : null; +} + +/** + * The aligner for `language` when it is already on disk and intact, without + * touching the network: null when it is missing, corrupt, or the language has + * none. A local read, so a transcription can afford to wait for it. + */ +export async function cachedAligner(baseDir: string, language: string): Promise { + const filePath = alignerPath(baseDir, language); + if (!filePath || !existsSync(filePath)) return null; + const actual = await sha256OfFile(filePath).catch(() => ""); + return actual === CTC_ALIGNERS[language].expectedSha256 ? filePath : null; +} + +/** Abort an aligner download that receives no bytes for this long. */ +export const ALIGNER_STALL_MS = 30_000; + +/** + * Make sure the aligner for `language` is on disk, downloading and verifying it + * like the whisper model. Null when the language has none; throws when the + * download fails, stalls for `stallMs` or is aborted through `signal`, which the + * caller turns into "keep whisper's times". + */ +export async function ensureAligner(opts: { + baseDir: string; + language: string; + signal?: AbortSignal; + stallMs?: number; + fetcher?: typeof fetch; +}): Promise { + const filePath = alignerPath(opts.baseDir, opts.language); + if (!filePath) return null; + const file = CTC_ALIGNERS[opts.language]; + await ensureFile(filePath, file.url, file.expectedSha256, { + fetcher: opts.fetcher, + signal: opts.signal, + stallMs: opts.stallMs ?? ALIGNER_STALL_MS, + }); + return filePath; +} + export function modelPaths(baseDir: string): Record { return { whisper: path.join(baseDir, STT_MODELS.whisper.cacheDir, MODEL_FILE), @@ -132,8 +217,19 @@ export async function sha256OfFile(filePath: string): Promise { const MAX_ATTEMPTS = 6; const RETRYABLE_STATUS = new Set([408, 425, 429, 500, 502, 503, 504]); -function sleep(ms: number): Promise { - return new Promise((resolve) => setTimeout(resolve, ms)); +function sleep(ms: number, signal?: AbortSignal): Promise { + return new Promise((resolve, reject) => { + if (signal?.aborted) return reject(signal.reason); + const timer = setTimeout(resolve, ms); + signal?.addEventListener( + "abort", + () => { + clearTimeout(timer); + reject(signal.reason); + }, + { once: true }, + ); + }); } function backoffMs(attempt: number, retryAfter: string | null): number { @@ -146,19 +242,34 @@ function backoffMs(attempt: number, retryAfter: string | null): number { return Math.min(60_000, 2_000 * 2 ** (attempt - 1)) + Math.floor(Math.random() * 1000); } -async function fetchWithRetry(url: string, fetcher: typeof fetch): Promise { +/** + * `signal` aborts the request and the waits between attempts; `waiting` is told + * when a backoff starts and ends, so a stall timer does not count it. + */ +async function fetchWithRetry( + url: string, + fetcher: typeof fetch, + signal?: AbortSignal, + waiting: (yes: boolean) => void = () => undefined, +): Promise { let lastErr: unknown; + const backoff = async (ms: number) => { + waiting(true); + await sleep(ms, signal); + waiting(false); + }; for (let attempt = 1; attempt <= MAX_ATTEMPTS; attempt++) { try { const res = await fetcher(url, { headers: { "user-agent": "openscreen-stt" }, + signal, }); if (res.ok && res.body) return res; if (res.status >= 400 && res.status < 500 && !RETRYABLE_STATUS.has(res.status)) { throw new Error(`Failed to download ${url}: HTTP ${res.status} ${res.statusText}`); } if (RETRYABLE_STATUS.has(res.status) && attempt < MAX_ATTEMPTS) { - await sleep(backoffMs(attempt, res.headers.get("retry-after"))); + await backoff(backoffMs(attempt, res.headers.get("retry-after"))); continue; } throw new Error(`Failed to download ${url}: HTTP ${res.status} ${res.statusText}`); @@ -167,8 +278,8 @@ async function fetchWithRetry(url: string, fetcher: typeof fetch): Promise= MAX_ATTEMPTS) throw err; - await sleep(backoffMs(attempt, null)); + if (signal?.aborted || attempt >= MAX_ATTEMPTS) throw signal?.reason ?? err; + await backoff(backoffMs(attempt, null)); } } throw lastErr; @@ -179,6 +290,10 @@ export interface DownloadOptions { onProgress?: (bytes: number) => void; /** Override fetch (for tests); defaults to `globalThis.fetch`. */ fetcher?: typeof fetch; + /** Aborts the download (and the waits between attempts). */ + signal?: AbortSignal; + /** Abort when no byte arrives for this long, backoffs aside. Off when unset. */ + stallMs?: number; } /** @@ -212,17 +327,44 @@ async function ensureFile( await mkdir(path.dirname(filePath), { recursive: true }); const fetcher = options.fetcher ?? fetch; - const res = await fetchWithRetry(fileUrl, fetcher); const tmp = `${filePath}.partial`; - let downloaded = 0; - - const source = Readable.fromWeb(res.body as never); - source.on("data", (chunk: Buffer | Uint8Array) => { - downloaded += chunk.length; - options.onProgress?.(downloaded); - }); - const { createWriteStream } = await import("node:fs"); - await pipeline(source, createWriteStream(tmp)); + // One controller for the caller's abort and the stall timer: a connection that + // stops sending leaves `pipeline` pending forever, with nothing else to end it. + const controller = new AbortController(); + const onAbort = () => controller.abort(options.signal?.reason); + if (options.signal?.aborted) onAbort(); + options.signal?.addEventListener("abort", onAbort, { once: true }); + let stall: ReturnType | undefined; + const arm = (on = true) => { + clearTimeout(stall); + if (!on || !options.stallMs) return; + const ms = options.stallMs; + stall = setTimeout( + () => controller.abort(new Error(`download stalled: no data for ${ms / 1000}s`)), + ms, + ); + }; + try { + arm(); + const res = await fetchWithRetry(fileUrl, fetcher, controller.signal, (waiting) => + arm(!waiting), + ); + let downloaded = 0; + const source = Readable.fromWeb(res.body as never); + source.on("data", (chunk: Buffer | Uint8Array) => { + arm(); + downloaded += chunk.length; + options.onProgress?.(downloaded); + }); + const { createWriteStream } = await import("node:fs"); + await pipeline(source, createWriteStream(tmp), { signal: controller.signal }); + } catch (error) { + await rm(tmp, { force: true }).catch(() => undefined); + throw controller.signal.aborted ? (controller.signal.reason ?? error) : error; + } finally { + arm(false); + options.signal?.removeEventListener("abort", onAbort); + } if (expectedSha256) { const actual = await sha256OfFile(tmp); diff --git a/electron/stt/whisperServer.test.ts b/electron/stt/whisperServer.test.ts index 62dd7442d..773235577 100644 --- a/electron/stt/whisperServer.test.ts +++ b/electron/stt/whisperServer.test.ts @@ -382,6 +382,115 @@ describe("WhisperServerManager", () => { } }); + describe("CTC word aligner", () => { + const inference = { + segments: [ + { + text: " ab ba", + start: 0.2, + end: 1, + words: [ + { word: " ab", start: 0.3, end: 0.5, anchor: 0.4 }, + { word: " ba", start: 0.5, end: 0.8, anchor: 0.6 }, + ], + }, + ], + speech: [{ start: 0.2, end: 1 }], + detected_language: "fr", + backend: "whispercpp-vulkan", + timing: { elapsed_s: 0.5, audio_s: 2, rtf: 0.25 }, + }; + // "ab" spoken on frames 20-21, the word delimiter on 22, "ba" on 24-25. + const vocab = ["", "|", "a", "b"]; + const frames = Array(60).fill(0); + frames[20] = 2; + frames[21] = 3; + frames[22] = 1; + frames[24] = 3; + frames[25] = 2; + const logprobs = new Float32Array(60 * vocab.length).fill(-12); + frames.forEach((tok, f) => { + logprobs[f * vocab.length + tok] = -0.01; + }); + const emissions = { + vocab, + blank: 0, + stride_s: 0.02, + elapsed_s: 0.1, + regions: [ + { start: 0, frames: 60, logprobs: Buffer.from(logprobs.buffer).toString("base64") }, + ], + }; + + async function run(emissionsReply: Response, reply: object = inference) { + const fetchMock = vi.fn(async (url: string, _init?: RequestInit) => + url.endsWith("/emissions") + ? emissionsReply + : new Response(JSON.stringify(reply), { status: 200 }), + ); + vi.stubGlobal("fetch", fetchMock); + try { + const mgr = new WhisperServerManager(); + (mgr as unknown as { process: unknown; port: number }).process = {}; + (mgr as unknown as { process: unknown; port: number }).port = 9999; + const alignerFor = vi.fn(async (language: string) => `/models/${language}.gguf`); + const result = await mgr.transcribe({ + samples: new Float32Array(16_000 * 2), + alignerFor, + }); + return { result, alignerFor, fetchMock }; + } finally { + vi.unstubAllGlobals(); + } + } + + it("re-times the words on the aligner of the detected language, and counts its cost", async () => { + const { result, alignerFor, fetchMock } = await run( + new Response(JSON.stringify(emissions), { status: 200 }), + ); + expect(alignerFor).toHaveBeenCalledWith("fr"); + const body = fetchMock.mock.calls[1][1]?.body as FormData; + expect(body.get("model")).toBe("/models/fr.gguf"); + expect(JSON.parse(String(body.get("regions")))).toEqual([[0, 1.3]]); + // The phrase edges stay on the speech; the inner boundary moves to the letters. + const ms = (x: number) => Math.round(x * 1000); + expect(result.wordSegments.map((w) => [w.word, ms(w.startSec), ms(w.endSec)])).toEqual([ + ["ab", 200, 450], + ["ba", 450, 1000], + ]); + expect(result.timing?.elapsedSec).toBeCloseTo(0.6); + expect(result.timing?.rtf).toBeCloseTo(0.3); + }); + + it("leaves the aligner out on the CPU, without even asking for it", async () => { + const { result, alignerFor, fetchMock } = await run( + new Response(JSON.stringify(emissions), { status: 200 }), + { ...inference, backend: "whispercpp-cpu" }, + ); + expect(alignerFor).not.toHaveBeenCalled(); + expect(fetchMock).toHaveBeenCalledOnce(); + expect(result.wordSegments.map((w) => [w.word, w.startSec, w.endSec])).toEqual([ + ["ab", 0.2, 0.5], + ["ba", 0.5, 1], + ]); + }); + + it("keeps whisper's times when the aligner fails", async () => { + const warn = vi.spyOn(console, "warn").mockImplementation(() => undefined); + try { + const { result } = await run(new Response("not found", { status: 404 })); + expect(result.wordSegments.map((w) => [w.word, w.startSec, w.endSec])).toEqual([ + ["ab", 0.2, 0.5], + ["ba", 0.5, 1], + ]); + expect(result.timing?.elapsedSec).toBe(0.5); + expect(warn).toHaveBeenCalledWith(expect.stringMatching(/word aligner failed/)); + } finally { + warn.mockRestore(); + } + }); + }); + it("spawns whisper-stt-server with --model", async () => { const fs = await import("node:fs/promises"); const { spawn } = await import("node:child_process"); diff --git a/electron/stt/whisperServer.ts b/electron/stt/whisperServer.ts index a03c05449..2543a79bb 100644 --- a/electron/stt/whisperServer.ts +++ b/electron/stt/whisperServer.ts @@ -6,6 +6,12 @@ import os from "node:os"; import path from "node:path"; import type { Readable } from "node:stream"; +import { + alignWordsOnEmissions, + type EmissionsJson, + emissionRegions, + parseEmissions, +} from "./ctcAlign"; import { resolveBinaryPath } from "./gpuDetector"; import { anchorWordsOnSpeech } from "./snapWordBoundaries"; import type { @@ -76,6 +82,13 @@ export interface WhisperServerStatus { lastError: string | null; } +export interface TranscribeOptions { + samples: Float32Array; + language?: string; + /** The aligner model for a language, or null for none (see `transcribe`). */ + alignerFor?: (language: string) => Promise; +} + /** Per-word entry inside a whisper-stt-server `/inference` JSON segment. */ interface WhisperJsonWord { word?: string; @@ -441,6 +454,33 @@ export class WhisperServerManager { } } + /** + * The CTC aligner's letter scores for `regions` of the same upload + * (`/emissions`, see ctcAlign.ts). Throws on any failure, including a helper + * built before the endpoint existed (404): the caller keeps whisper's times. + */ + private async runEmissions(opts: { + wavPath: string; + modelPath: string; + regions: Array<[number, number]>; + }): Promise { + const form = new FormData(); + const fileBuffer = await readFile(opts.wavPath); + form.set("file", new Blob([fileBuffer], { type: "audio/wav" }), path.basename(opts.wavPath)); + form.set("model", opts.modelPath); + form.set("regions", JSON.stringify(opts.regions)); + const res = await fetch(`${this.baseUrl()}/emissions`, { + method: "POST", + body: form, + signal: AbortSignal.timeout(REQUEST_TIMEOUT_MS), + }); + if (!res.ok) { + const text = await res.text().catch(() => ""); + throw new Error(`whisper-stt-server /emissions HTTP ${res.status}: ${text.slice(0, 512)}`); + } + return (await res.json()) as EmissionsJson; + } + private async runMultipartInfer(opts: { wavPath: string; language?: string; @@ -541,8 +581,16 @@ export class WhisperServerManager { return { elapsedSec, audioSec, rtf }; } - /** Run one transcription; serializes concurrent callers. */ - async transcribe(opts: { samples: Float32Array; language?: string }): Promise<{ + /** + * Run one transcription; serializes concurrent callers. + * + * `alignerFor` resolves the CTC aligner model for the language whisper + * detected, or null when there is none or it is not ready yet (it must not + * wait on the network: no download runs inside a chunk, see SttManager). With one, + * the words are re-timed on it before their phrase edges go on the speech + * (ctcAlign.ts); any failure there keeps whisper's own times. + */ + async transcribe(opts: TranscribeOptions): Promise<{ segments: SttPhraseSegment[]; wordSegments: SttWordSegment[]; detectedLanguage: string; @@ -555,7 +603,7 @@ export class WhisperServerManager { return task; } - private async transcribeImpl(opts: { samples: Float32Array; language?: string }): Promise<{ + private async transcribeImpl(opts: TranscribeOptions): Promise<{ segments: SttPhraseSegment[]; wordSegments: SttWordSegment[]; detectedLanguage: string; @@ -578,33 +626,62 @@ export class WhisperServerManager { startSec: this.toSec(s.start, 0), endSec: this.toSec(s.end, 0), })); - // The helper's word times are right inside a phrase but not on its edges, - // which the transcript editor turns into imprecise trims: put those on the - // helper's speech intervals (see snapWordBoundaries.ts). - const wordSegments: SttWordSegment[] = anchorWordsOnSpeech( - raw - .flatMap((seg) => - (seg.words ?? []).map((w) => { - const word = (w.word ?? "").trim(); - const startSec = this.toSec(w.start, 0); - const endSec = this.toSec(w.end, startSec + 0.05); - const anchorSec = this.toSec(w.anchor, startSec); - const confidence = typeof w.probability === "number" ? w.probability : undefined; - return { - word, - startSec, - endSec: Math.max(startSec + 0.02, endSec), - anchorSec, - confidence, - }; - }), - ) - .filter((w) => w.word.length > 0), - speech, - ); const detectedLanguage = json.detected_language ?? json.language ?? "auto"; + let words = raw + .flatMap((seg) => + (seg.words ?? []).map((w) => { + const word = (w.word ?? "").trim(); + const startSec = this.toSec(w.start, 0); + const endSec = this.toSec(w.end, startSec + 0.05); + const anchorSec = this.toSec(w.anchor, startSec); + const confidence = typeof w.probability === "number" ? w.probability : undefined; + return { + word, + startSec, + endSec: Math.max(startSec + 0.02, endSec), + anchorSec, + confidence, + }; + }), + ) + .filter((w) => w.word.length > 0); + let alignSec = 0; const backend = this.toBackend(json.backend); + // GPU only. On the CPU the aligner's forward pass costs +29% (English) to + // +58% (French) of a transcription, against +9% to +14% on Vulkan, and + // character-level DTW alone is already within 16 ms median there + // (tools/stt-eval/word-timing, issue #948). Not asking also means a + // CPU-only machine never downloads the model. + if (speech?.length && words.length && opts.alignerFor && backend !== "whispercpp-cpu") { + try { + const modelPath = await opts.alignerFor(detectedLanguage); + if (modelPath) { + const reply = await this.runEmissions({ + wavPath, + modelPath, + regions: emissionRegions(speech, opts.samples.length / 16_000), + }); + const emissions = parseEmissions(reply); + if (!emissions) throw new Error("unreadable /emissions reply"); + words = alignWordsOnEmissions(words, speech, emissions); + alignSec = this.toSec(reply.elapsed_s, 0); + } + } catch (error) { + console.warn( + `[stt] word aligner failed, keeping whisper's word times: ${error instanceof Error ? error.message : String(error)}`, + ); + } + } + // Phrase edges go on the helper's speech intervals, aligner or not: the + // VAD is the better judge of where speech starts after a pause (see + // snapWordBoundaries.ts). + const wordSegments: SttWordSegment[] = anchorWordsOnSpeech(words, speech); const timing = this.toTiming(json.timing); + // The aligner is part of what this chunk cost. + if (timing && alignSec > 0) { + timing.elapsedSec += alignSec; + timing.rtf = timing.elapsedSec / timing.audioSec; + } return { segments, wordSegments, detectedLanguage, backend, timing }; } finally { await cleanupWav(wavPath); diff --git a/scripts/convert-wav2vec2-gguf.mjs b/scripts/convert-wav2vec2-gguf.mjs new file mode 100644 index 000000000..55e102d75 --- /dev/null +++ b/scripts/convert-wav2vec2-gguf.mjs @@ -0,0 +1,225 @@ +// Converts a HuggingFace Wav2Vec2ForCTC checkpoint (model.safetensors + +// config.json + vocab.json) into the GGUF file whisper-stt-server's CTC aligner +// loads (electron/native/whisper-stt/src/ctc_aligner.cpp). +// +// Usage: +// node scripts/convert-wav2vec2-gguf.mjs \ +// --languages en[,xx] [--type q8_0|f16] +// +// Node stdlib only. What it writes, and why: +// - Linear weights (attention, feed-forward, projection, lm_head) as Q8_0, the +// rest of the matrices as F16, norms and biases as F32. Q8_0 is what keeps the +// French large model under the 350 MB download budget (issue #948). The +// convolutions must stay F16: the helper multiplies them with F16 columns. +// - Sources, as pinned in electron/stt/modelManager.ts (CTC_ALIGNERS[].source): +// en: https://huggingface.co/facebook/wav2vec2-base-960h (model.safetensors) +// fr: https://huggingface.co/jonatasgrosman/wav2vec2-large-xlsr-53-french, +// model.safetensors from its refs/pr/2 (main only has pytorch_model.bin) +// with the config.json and vocab.json of the same revision. +// - The positional convolution's weight norm is folded (weight = g * v / |v|), +// so the helper never sees `weight_g` / `weight_v`. +// - Tensor names lose the `wav2vec2.` prefix; `masked_spec_embed` is dropped. +// - The config and the vocabulary go in as `w2v.*` metadata. +// The output is byte-for-byte deterministic for a given input, so its SHA-256 can +// be pinned in electron/stt/modelManager.ts. +import fs from "node:fs"; + +const [stPath, configPath, vocabPath, outPath, ...rest] = process.argv.slice(2); +if (!outPath) { + console.error( + "usage: node convert-wav2vec2-gguf.mjs --languages en [--type q8_0|f16]", + ); + process.exit(2); +} +const flag = (name, dflt) => { + const i = rest.indexOf(name); + return i >= 0 ? rest[i + 1] : dflt; +}; +const linearType = flag("--type", "q8_0"); +const languages = flag("--languages", "").split(",").filter(Boolean); +if (!languages.length) throw new Error("--languages is required (e.g. --languages en)"); + +const config = JSON.parse(fs.readFileSync(configPath, "utf8")); +const vocabMap = JSON.parse(fs.readFileSync(vocabPath, "utf8")); +const vocab = []; +for (const [tok, id] of Object.entries(vocabMap)) vocab[id] = tok; +if (vocab.length !== config.vocab_size || vocab.includes(undefined)) + throw new Error(`vocab.json has ${vocab.length} ids, config says ${config.vocab_size}`); + +// ---- safetensors ---- +const fd = fs.openSync(stPath); +const lenBuf = Buffer.alloc(8); +fs.readSync(fd, lenBuf, 0, 8, 0); +const headerLen = Number(lenBuf.readBigUInt64LE()); +const headerBuf = Buffer.alloc(headerLen); +fs.readSync(fd, headerBuf, 0, headerLen, 8); +const header = JSON.parse(headerBuf.toString("utf8")); +delete header.__metadata__; +const readF32 = (name) => { + const t = header[name]; + if (!t) throw new Error(`missing tensor ${name}`); + if (t.dtype !== "F32") throw new Error(`${name}: expected F32, got ${t.dtype}`); + const [b, e] = t.data_offsets; + const buf = Buffer.alloc(e - b); + fs.readSync(fd, buf, 0, e - b, 8 + headerLen + b); + return { shape: t.shape, data: new Float32Array(buf.buffer, buf.byteOffset, (e - b) / 4) }; +}; + +// ---- f16 / q8_0 ---- +const f32 = new Float32Array(1); +const u32 = new Uint32Array(f32.buffer); +function toHalf(v) { + f32[0] = v; + const x = u32[0]; + const sign = (x >>> 16) & 0x8000; + let exp = ((x >>> 23) & 0xff) - 127 + 15; + let mant = x & 0x7fffff; + if (exp >= 31) return sign | 0x7c00; // overflow -> inf (never for these weights) + if (exp <= 0) { + if (exp < -10) return sign; + mant |= 0x800000; + const shift = 14 - exp; + let h = mant >> shift; + if ((mant >> (shift - 1)) & 1 && (mant & ((1 << (shift - 1)) - 1) || h & 1)) h++; + return sign | h; + } + let h = sign | (exp << 10) | (mant >> 13); + // round to nearest even + if (mant & 0x1000 && (mant & 0xfff || h & 1)) h++; + return h; +} +function encodeF16(data) { + const out = Buffer.alloc(data.length * 2); + for (let i = 0; i < data.length; i++) out.writeUInt16LE(toHalf(data[i]), i * 2); + return out; +} +// ggml block_q8_0: f16 scale + 32 int8, scale = amax / 127 (ggml's quantize_row_q8_0_ref). +function encodeQ8_0(data) { + if (data.length % 32) throw new Error("q8_0 needs a multiple of 32 elements"); + const nb = data.length / 32; + const out = Buffer.alloc(nb * 34); + for (let b = 0; b < nb; b++) { + let amax = 0; + for (let i = 0; i < 32; i++) amax = Math.max(amax, Math.abs(data[b * 32 + i])); + const d = amax / 127; + const id = d ? 1 / d : 0; + out.writeUInt16LE(toHalf(d), b * 34); + for (let i = 0; i < 32; i++) out.writeInt8(Math.round(data[b * 32 + i] * id), b * 34 + 2 + i); + } + return out; +} + +const GGML = { f32: 0, f16: 1, q8_0: 8 }; +const tensors = []; // { name, ne (ggml order), type, bytes } +function add(name, shape, data, type) { + const ne = [...shape].reverse(); + if (type === "q8_0" && ne[0] % 32) type = "f16"; + const bytes = + type === "f32" + ? Buffer.from(data.buffer, data.byteOffset, data.byteLength) + : type === "f16" + ? encodeF16(data) + : encodeQ8_0(data); + tensors.push({ name, ne, type: GGML[type], bytes }); +} + +const isLinear = (n) => + /(attention\.(q|k|v|out)_proj|feed_forward\.(intermediate|output)_dense|feature_projection\.projection|lm_head)\.weight$/.test( + n, + ); +for (const name of Object.keys(header).sort()) { + if (name.endsWith("masked_spec_embed") || name.includes("pos_conv_embed.conv.weight_")) continue; + const { shape, data } = readF32(name); + const short = name.replace(/^wav2vec2\./, ""); + if (shape.length === 1) add(short, shape, data, "f32"); + else add(short, shape, data, isLinear(name) ? linearType : "f16"); +} + +// Weight norm with dim=2: one norm per kernel tap k, over (out, in/groups). +{ + const g = readF32("wav2vec2.encoder.pos_conv_embed.conv.weight_g").data; + const v = readF32("wav2vec2.encoder.pos_conv_embed.conv.weight_v"); + const [O, I, K] = v.shape; + const w = new Float32Array(v.data.length); + for (let k = 0; k < K; k++) { + let s = 0; + for (let o = 0; o < O; o++) for (let i = 0; i < I; i++) s += v.data[(o * I + i) * K + k] ** 2; + const scale = g[k] / Math.sqrt(s); + for (let o = 0; o < O; o++) + for (let i = 0; i < I; i++) w[(o * I + i) * K + k] = v.data[(o * I + i) * K + k] * scale; + } + add("encoder.pos_conv_embed.conv.weight", v.shape, w, "f16"); +} + +// ---- GGUF v3 ---- +const chunks = []; +const u32b = (v) => { + const b = Buffer.alloc(4); + b.writeUInt32LE(v); + return b; +}; +const u64b = (v) => { + const b = Buffer.alloc(8); + b.writeBigUInt64LE(BigInt(v)); + return b; +}; +const strb = (s) => { + const b = Buffer.from(s, "utf8"); + return Buffer.concat([u64b(b.length), b]); +}; +const T = { u32: 4, i32: 5, f32: 6, bool: 7, string: 8, array: 9 }; +const kvs = []; +const kv = (key, type, value) => kvs.push({ key, type, value }); +kv("general.architecture", "string", "wav2vec2"); +kv("general.alignment", "u32", 32); +kv("w2v.languages", "array:string", languages); +kv("w2v.vocab", "array:string", vocab); +kv("w2v.blank_id", "u32", config.pad_token_id ?? 0); +kv("w2v.word_delimiter", "string", "|"); +kv("w2v.hidden_size", "u32", config.hidden_size); +kv("w2v.intermediate_size", "u32", config.intermediate_size); +kv("w2v.num_hidden_layers", "u32", config.num_hidden_layers); +kv("w2v.num_attention_heads", "u32", config.num_attention_heads); +kv("w2v.layer_norm_eps", "f32", config.layer_norm_eps); +kv("w2v.conv_kernel", "array:u32", config.conv_kernel); +kv("w2v.conv_stride", "array:u32", config.conv_stride); +kv("w2v.conv_bias", "bool", !!config.conv_bias); +kv("w2v.feat_extract_norm", "string", config.feat_extract_norm); +kv("w2v.stable_layer_norm", "bool", !!config.do_stable_layer_norm); +kv("w2v.num_conv_pos_embeddings", "u32", config.num_conv_pos_embeddings); +kv("w2v.num_conv_pos_embedding_groups", "u32", config.num_conv_pos_embedding_groups); + +chunks.push(Buffer.from("GGUF"), u32b(3), u64b(tensors.length), u64b(kvs.length)); +for (const { key, type, value } of kvs) { + chunks.push(strb(key)); + if (type.startsWith("array:")) { + const et = type.slice(6); + chunks.push(u32b(T.array), u32b(T[et]), u64b(value.length)); + for (const v of value) chunks.push(et === "string" ? strb(v) : u32b(v)); + continue; + } + chunks.push(u32b(T[type])); + if (type === "string") chunks.push(strb(value)); + else if (type === "u32") chunks.push(u32b(value)); + else if (type === "bool") chunks.push(Buffer.from([value ? 1 : 0])); + else if (type === "f32") { + const b = Buffer.alloc(4); + b.writeFloatLE(value); + chunks.push(b); + } +} +const ALIGN = 32; +const pad = (n) => (ALIGN - (n % ALIGN)) % ALIGN; +let offset = 0; +for (const t of tensors) { + chunks.push(strb(t.name), u32b(t.ne.length)); + for (const d of t.ne) chunks.push(u64b(d)); + chunks.push(u32b(t.type), u64b(offset)); + offset += t.bytes.length + pad(t.bytes.length); +} +let headLen = chunks.reduce((s, c) => s + c.length, 0); +chunks.push(Buffer.alloc(pad(headLen))); +for (const t of tensors) chunks.push(t.bytes, Buffer.alloc(pad(t.bytes.length))); +fs.writeFileSync(outPath, Buffer.concat(chunks)); +headLen = fs.statSync(outPath).size; +console.log(`${outPath}: ${tensors.length} tensors, ${(headLen / 1e6).toFixed(1)} MB`); diff --git a/scripts/test-whisper-stt.mjs b/scripts/test-whisper-stt.mjs index 2a0e11223..8ea9f2a34 100644 --- a/scripts/test-whisper-stt.mjs +++ b/scripts/test-whisper-stt.mjs @@ -21,12 +21,20 @@ // OPENSCREEN_WHISPER_MODEL GGML model (default: the userData cache location) // OPENSCREEN_VAD_MODEL Silero VAD model (default: next to the GGML model; // the speech checks run only when it is there, as in the app) +// OPENSCREEN_ALIGNER_MODEL CTC word aligner (default: the cached one for the detected +// language; the aligner checks run only when it is there) import { spawn, spawnSync } from "node:child_process"; import fs from "node:fs"; import os from "node:os"; import path from "node:path"; import { fileURLToPath } from "node:url"; +import { + alignWordsOnEmissions, + emissionRegions, + parseEmissions, +} from "../electron/stt/ctcAlign.ts"; +import { CTC_ALIGNERS } from "../electron/stt/modelManager.ts"; const __dirname = path.dirname(fileURLToPath(import.meta.url)); const ROOT = path.join(__dirname, ".."); @@ -244,6 +252,9 @@ async function main() { }); let json; + let emissions = null; + let alignerModel = null; + let alignerSkipped = ""; try { await waitForReady(`http://127.0.0.1:${port}/`); const form = new FormData(); @@ -256,6 +267,32 @@ async function main() { }); if (!res.ok) throw new Error(`/inference returned ${res.status}: ${await res.text()}`); json = await res.json(); + // The word aligner's acoustic pass, on the same helper, as SttManager runs it. + const cached = CTC_ALIGNERS[json.detected_language]; + alignerModel = + process.env.OPENSCREEN_ALIGNER_MODEL ?? + (cached ? path.join(path.dirname(path.dirname(MODEL)), "ctc-aligner", cached.name) : null); + if (WITH_VAD && alignerModel && fs.existsSync(alignerModel)) { + const speech = (json.speech ?? []).map((x) => ({ startSec: x.start, endSec: x.end })); + const aform = new FormData(); + aform.append("file", new Blob([fs.readFileSync(wavPath)]), path.basename(wavPath)); + aform.append("model", alignerModel); + aform.append("regions", JSON.stringify(emissionRegions(speech, meta.durationSec))); + const ares = await fetch(`http://127.0.0.1:${port}/emissions`, { + method: "POST", + body: aform, + }); + if (!ares.ok) throw new Error(`/emissions returned ${ares.status}: ${await ares.text()}`); + emissions = await ares.json(); + } else { + // Say which file was looked for: a custom OPENSCREEN_WHISPER_MODEL moves + // the cache this path is derived from, and would read as "no aligner". + alignerSkipped = !WITH_VAD + ? "no VAD model, so no speech to align" + : alignerModel + ? `no aligner at ${alignerModel} (set OPENSCREEN_ALIGNER_MODEL)` + : `no aligner for "${json.detected_language}"`; + } } finally { child.kill(); cleanup(); @@ -361,6 +398,76 @@ async function main() { check(rate <= 0.15, "WER within tolerance", `${rate.toFixed(4)}`); } + // The aligner (electron/stt/ctcAlign.ts) only gets letter scores from the + // helper; a forward pass gone wrong on a backend (a missing op, a wrong + // layout) still answers 200 with well-formed garbage. Reading its best letter + // per frame back as text is what tells the two apart. + if (emissions) { + const parsed = parseEmissions(emissions); + check(parsed !== null, "aligner answers well-formed letter scores", alignerModel); + if (parsed) { + const V = parsed.vocab.length; + let worst = 0; + let greedy = ""; + for (const r of parsed.regions) { + let prev = -1; + for (let t = 0; t < r.frames; t++) { + let sum = 0; + let best = 0; + for (let v = 0; v < V; v++) { + sum += Math.exp(r.logprobs[t * V + v]); + if (r.logprobs[t * V + v] > r.logprobs[t * V + best]) best = v; + } + worst = Math.max(worst, Math.abs(sum - 1)); + if (best !== prev && best !== parsed.blank) + greedy += parsed.vocab[best] === "|" ? " " : parsed.vocab[best]; + prev = best; + } + greedy += " "; + } + check( + worst < 1e-3, + "aligner frames are log-probabilities", + `max |sum - 1| = ${worst.toExponential(1)}`, + ); + if (gpuExpected) { + check( + emissions.device !== "CPU", + "aligner runs on the GPU on a GPU-capable host", + emissions.device, + ); + } + // Letters, not words: the aligner spells "OpenScreen" as it hears it. + // Measured: 0.31 on a real French take, 0.11 on English TTS; a forward + // pass gone wrong decodes to noise, near 1. + const letters = (x) => normalize(x).join("").split("").join(" "); + const heard = wer(letters(text), letters(greedy)); + console.log(`\naligner: ${greedy.trim().slice(0, 200)}`); + check( + heard <= 0.5, + "aligner hears what whisper heard", + `letter error rate ${heard.toFixed(2)} on ${emissions.device}`, + ); + const speech = (json.speech ?? []).map((x) => ({ startSec: x.start, endSec: x.end })); + const plain = words.map((w) => ({ + word: w.word.trim(), + startSec: w.start, + endSec: w.end, + anchorSec: w.anchor ?? w.start, + })); + const aligned = alignWordsOnEmissions(plain, speech, parsed); + check( + aligned.every( + (w, i) => w.endSec >= w.startSec && (i === 0 || w.startSec >= aligned[i - 1].startSec), + ), + "aligned words stay ordered and non-empty", + ); + console.log(`aligner: ${emissions.elapsed_s.toFixed(2)}s on ${emissions.device}`); + } + } else { + console.log(`\n(aligner checks skipped: ${alignerSkipped})`); + } + if (json.timing) { console.log( `\ntiming : ${json.timing.elapsed_s?.toFixed?.(2)}s for ` + diff --git a/technical-documentation/architecture/transcription-and-captions.md b/technical-documentation/architecture/transcription-and-captions.md index 41645e23a..d3b1dfcaf 100644 --- a/technical-documentation/architecture/transcription-and-captions.md +++ b/technical-documentation/architecture/transcription-and-captions.md @@ -112,6 +112,10 @@ renderer imposes a timeout that a slow download could trip: the preload does a bare `ipcRenderer.invoke`, `fetchWithRetry` has no per-request deadline, and whisper-server's 30 s readiness budget only starts once the download resolved. +The word aligner (§ Word-level alignment, step 5) is not part of that phase: it +is fetched in the background the first time a transcription on a GPU detects English or +French (109 or 348 MB), and no chunk ever waits for it. + Three edges make that promise hold, and each is load-bearing: - **A failed setup is not cached.** `SttManager.init` used to memoise the @@ -184,14 +188,14 @@ it verbatim in the response. > `ggml_backend_dev_type()` being GPU/IGPU, which also drops ggml-blas > (Accelerate on macOS, device type ACCEL) from consideration. -### Word-level alignment (DTW token timestamps) +### Word-level alignment (DTW token timestamps, then a CTC aligner) 0. **Speech only** — when the Silero VAD model is on disk (`ggml-silero-v6.2.0.bin`, downloaded next to the whisper model), the helper cuts the silence out of the upload before whisper sees it, the way whisper.cpp's own VAD does: each speech stretch plus 0.1 s of tail, 0.1 s of silence between stretches. It keeps the map, and every time below goes back through it onto the upload's - clock. The stretches are returned as `speech` for step 5. + clock. The stretches are returned as `speech` for steps 5 and 6. > **Not whisper_full's `vad` param.** whisper.cpp 1.9.1 maps its *segment* > times back onto the original audio but not its token times, `t_dtw` > included, so every word came back early by all the silence removed before @@ -248,7 +252,7 @@ it verbatim in the response. first word has none and falls back to its own first token. - `word.end` = `t_dtw` of the word's **last** token. - `word.anchor` = `t_dtw` of the word's first token. It always lies inside - the word, and only decides which stretch of speech owns it (step 5). + the word, and only decides which stretch of speech owns it (steps 5 and 6). The result is a monotonic, gap-free timeline of word ranges on the upload's clock. Taking the first token's `t_dtw` as the start, as the helper did @@ -256,9 +260,44 @@ it verbatim in the response. below, and a median 200 ms on a real French take. The repo's DTW POC saw the same thing as whisper.cpp's `t_dtw` matching faster-whisper's word **end** (`tools/stt-eval/whispercpp-dtw-poc/REPORT.md`). -5. **Anchor phrase edges on the speech** — +5. **Re-time the words on a CTC aligner** (issue #948, phase 3) — when the + chunk ran on a GPU and the language whisper detected has an aligner, a wav2vec2 model fine-tuned for CTC scores + every 20 ms frame of the speech against every letter, and + [`electron/stt/ctcAlign.ts`](../../electron/stt/ctcAlign.ts) forces + whisper's own words through those scores: + - **The helper only scores.** `POST /emissions` (same WAV, the aligner's + GGUF path, the regions to score: each stretch of speech ±0.3 s, merged) + answers the log-probabilities per frame + ([`ctc_aligner.cpp`](../../electron/native/whisper-stt/src/ctc_aligner.cpp), + a wav2vec2 forward pass on ggml, on the GPU whisper uses). It is a + separate request so the stage stays separate: it reads whisper's words, + whatever produced their times. + - **Spelling.** Each word is lowercased (uppercased for the English model), + its apostrophes normalised, its punctuation dropped, and an accent the + vocabulary lacks falls back to the base letter. A word that cannot be + spelled (digits, `€`, another script) becomes one wildcard token that + scores as the frame's best letter, so it holds its place and its + neighbours stay aligned. The vocabulary's word delimiter `|` goes between + words. + - **Viterbi** over the frames of each stretch (the one that owns the words + by `anchor`, as in step 6), then calibration: CTC is sure of a letter a + little after the sound starts and before it ends, so a word starts 45 ms + before its first letter's frame and ends 25 ms after its last one. Where + the two estimates cross inside continuous speech, the boundary is their + midpoint. Next to a pause (a gap of 100 ms or more between letters) the + edges move further out, 60 ms before and 150 ms after, both capped at + mid-pause: early there is silence, late is an audible attack. + - **GPU only.** On the CPU the forward pass would add 29% (English) to 58% + (French) to a transcription, against 9 to 14% on Vulkan, for a gain over + step 2's character-level DTW that is not worth that there. A chunk whisper + ran on the CPU skips this step, and a CPU-only machine never downloads the + model. + - **Fallback.** CPU, no aligner for the language, a download not landed yet or + failed, a helper without `/emissions`, a stretch whose letters do not fit in + its frames: the words keep the times of step 4. Step 6 runs either way. +6. **Anchor phrase edges on the speech** — [`electron/stt/snapWordBoundaries.ts`](../../electron/stt/snapWordBoundaries.ts) - leaves boundaries inside a phrase alone: they are within ~16 ms of the audio + leaves boundaries inside a phrase alone: they are within ~15 ms of the audio already, and an energy snap on top of them drags correct boundaries early. The edges of a phrase are different. The token before a phrase's first word is the previous phrase's last one, so that word starts in the pause, as @@ -287,15 +326,44 @@ scores the helper and this post-pass against a Windows-TTS corpus with exact word times (27 min of French and English narration, clean and noisy): boundary error by position, and the audible residue and clipping of deleting one word or one phrase. Issue #948 has the method, the baseline and the plan. +The same harness scores real read speech: 46 min of LibriSpeech test-clean +(40 speakers, English) against Montreal Forced Aligner word times, which are +themselves about 10 to 20 ms from a human's. -| Pipeline (clean + noisy) | Inner start: median / P90 / within 50 ms | Phrase-initial start within 50 ms | One-word delete: clean cuts | Phrase delete: clean | +| Pipeline | Inner start: median / P90 / within 50 ms | Phrase-initial start within 50 ms | One-word delete: clean cuts | Phrase delete: clean | |---|---|---|---|---| -| First-token start + 150 ms RMS snap + VAD edges (before #948) | 105 / 275 ms / 32% | 83% | 5% | 84% | -| Previous-token start + two-way VAD edges | 31 / 125 ms / 64% | 89% | 19% | 88% | -| Same, character-level DTW | 16 / 60 ms / 86% | 89% | 39% | 93% | +| First-token start + 150 ms RMS snap + VAD edges (before #948), TTS | 105 / 275 ms / 32% | 83% | 5% | 84% | +| Previous-token start + two-way VAD edges (phase 1), TTS | 31 / 125 ms / 64% | 89% | 19% | 88% | +| Same, character-level DTW (phase 2), TTS clean | 17 / 60 ms / 86% | 89% | 39% | 94% | +| + CTC aligner (phase 3), TTS clean | 14 / 35 ms / 96% | 95% | 50% | 95% | +| Phase 2, TTS noisy | 16 / 60 ms / 86% | 88% | 40% | 92% | +| + CTC aligner, TTS noisy | 15 / 39 ms / 95% | 94% | 48% | 93% | +| Phase 2, LibriSpeech | 20 / 65 ms / 82% | 63% | 35% | 81% | +| + CTC aligner, LibriSpeech | 15 / 45 ms / 92% | 73% | 44% | 82% | Per language, character-level DTW gives French 19 ms median and 83% within -50 ms (was 26 ms, 70%) and English 15 ms and 89% (was 40 ms, 57%). +50 ms (was 26 ms, 70%) and English 15 ms and 89% (was 40 ms, 57%). The aligner +takes French to 14 ms and 95% (P90 66 → 36 ms, clean cuts 33% → 47%) and +English to 15 ms and 96% (P90 50 → 35 ms, clean cuts 45% → 52%). Deleting one +word then leaves 17 ms of it audible and cuts 15 ms into its neighbours (TTS +clean; 20 and 12 ms on LibriSpeech), against 19 and 24 ms with phase 2 alone. +The misses left on phrase-initial starts are mostly the reference's: the French +TTS voices start a word on its silent stop closure, which no aligner hears. + +On a real French take (25 s, no reference), every boundary the aligner moved +by 40 ms or more was checked against the audio's energy and zero crossings at +10 ms: each now sits on the acoustic boundary. "c'est" starts on its /s/ +(phase 2 was 90 ms early, in the end of the word before); "quoi" starts where +the /k/ closure begins, after the vowel of *c'est* (phase 2 was on the burst, +70 to 80 ms later, which is also a clean cut since the closure is silent); +in "en fait on va" the aligner is on the /f/, the dip before *on* and the /v/, +where phase 2 was 30 to 50 ms off each. + +Dropping French silent final consonants before aligning, as step 2 does, was +tried and left out: it changed none of the take's boundaries by more than +15 ms, and on the corpus it traded 3 points of French clean cuts for 1 to 2 +points of French phrase deletes (a phrase-final *fois* and *Windows* lost the +letter that held their end). A +15 ms calibration offset on every boundary gained 4 points of inner boundaries within 50 ms but dropped noisy phrase deletes to 82%, so it is not @@ -340,6 +408,29 @@ once instead of merely breaking new installs; pinning makes the recorded digest an invariant. Bumping the model therefore means bumping the revision and the digest together. +**The word aligners** (step 5) are fetched only for a language that has one, +into `stt-models/ctc-aligner/`, and never inside a chunk. The first time a +session meets a language, a copy already on disk is verified in place (a local +read). Without one, the download starts in the background +(`SttManager.alignerFor`): the chunks that run before it lands keep whisper's +times, the later ones and later transcriptions use it. A download stalled for +30 s, cancelled with the transcription, or interrupted by quitting is +abandoned, and retried on the next transcription rather than per chunk. + +| Language | Model | Download | +|---|---|---| +| English | `facebook/wav2vec2-base-960h` (12 layers, letters) | 109 MB | +| French | `jonatasgrosman/wav2vec2-large-xlsr-53-french` (24 layers, letters and accents) | 348 MB | + +Both are Apache-2.0. `scripts/convert-wav2vec2-gguf.mjs` turns the upstream +`model.safetensors` into the GGUF the helper loads: linear weights Q8_0 (which +is what keeps the French model under 350 MB), convolutions F16, the positional +convolution's weight norm folded. The conversion is deterministic, so the +SHA-256 pinned in `CTC_ALIGNERS` (`electron/stt/modelManager.ts`) can be +reproduced from the upstream revision it names. The files are served from a +`v0.0.0-ctc-aligners-1` release, the convention this repo uses for binaries +that need a permanent URL but are not a product version. + > The HuggingFace identifier is intentionally `ggerganov/whisper.cpp`, > **not** `ggml-org/whisper.cpp`. The latter matches the GitHub org the > engine itself now lives under, but on HuggingFace it is a separate, @@ -352,13 +443,15 @@ and the digest together. | Module | Role | |---|---| | `electron/stt/whisperServer.ts` | Server lifecycle; `POST /inference` client; verbose_json parser. | -| `electron/stt/snapWordBoundaries.ts` | Anchors phrase edges on the VAD's speech (see step 5 above). | +| `electron/stt/ctcAlign.ts` | Re-times words on the CTC aligner's letter scores (step 5): spelling, Viterbi, calibration, fallback. | +| `electron/stt/snapWordBoundaries.ts` | Anchors phrase edges on the VAD's speech (see step 6 above). | | `electron/stt/wav.ts` | WAV write + temp-file cleanup helpers. | | `electron/stt/gpuDetector.ts` | Per-platform binary resolver (no GPU probing). | -| `electron/stt/modelManager.ts` | Single GGML file download, SHA-256 verify, atomic write. | +| `electron/stt/modelManager.ts` | Model downloads (whisper, Silero VAD, the per-language aligners), SHA-256 verify, atomic write. | | `electron/stt/transcriptionContract.ts` | Shared IPC types (`SttBackend`, `SttWordSegment`, `SttPhraseSegment`, `SttStatusEvent`). | | `electron/stt/index.ts` | `SttManager` — IPC entry point; wires the pieces together. | -| `electron/native/whisper-stt/src/main.cpp` | httplib HTTP server; cuts the silence out with Silero VAD and maps times back (step 0); calls `whisper_full()` with DTW; reports the device it bound via `ggml_backend_dev_name()`. | +| `electron/native/whisper-stt/src/main.cpp` | httplib HTTP server; cuts the silence out with Silero VAD and maps times back (step 0); calls `whisper_full()` with DTW; reports the device it bound via `ggml_backend_dev_name()`; `POST /emissions` for the aligner. | +| `electron/native/whisper-stt/src/ctc_aligner.cpp` | wav2vec2-for-CTC forward pass on ggml (base and large layouts), in 20 s windows; loads the GGUF written by `scripts/convert-wav2vec2-gguf.mjs`. | | `electron/native/whisper-stt/CMakeLists.txt` | Pulls whisper.cpp via FetchContent; enables Metal (macOS arm64), Vulkan (Windows/Linux x64), CPU fallback everywhere; static backend linking into `whisper.dll`/`ggml.dll`. | The helper is one executable per platform; backends are baked in at build @@ -392,7 +485,11 @@ timings present, the DTW guardrail passed, `backend` reporting GPU offload on a GPU-capable host, `detected_language` resolved rather than echoed, word times monotonic and inside the clip, and WER against a reference. On macOS it synthesizes its own clip with `say`, so it needs no fixture; elsewhere pass -`--wav ` (and optionally `--ref ""`). This is the check +`--wav ` (and optionally `--ref ""`). When the detected +language's aligner is cached (or `OPENSCREEN_ALIGNER_MODEL` names one), it also +calls `/emissions` and checks that the answer is well-formed log-probabilities, +computed on the GPU on a GPU-capable host, decoding to roughly what whisper +heard (letter error rate), with the aligned words still ordered. This is the check that the unit tests structurally cannot make: they mock `fetch`, so they assert against a hand-written fixture rather than the binary. @@ -818,9 +915,13 @@ it deletes data - **No C++ unit tests.** The WAV reader and the DTW-inactive guardrail in `electron/native/whisper-stt/src/main.cpp` are exercised only at runtime. -- **Word timing inside a phrase.** Phrase edges sit on the VAD (step 5), but - a word in the middle of continuous speech is only as good as whisper-small's - character-level DTW: about 16 ms median and 60 ms P90 on synthetic speech, - so about two single-word deletes in five are clean. A CTC forced aligner is - the upgrade path (issue #948 Phase 3); `tools/stt-eval/word-timing` - measures it. +- **Word timing without an aligner.** Phrase edges sit on the VAD (step 6), + and on a GPU English and French words are re-timed on a CTC aligner (step 5). + Every other language, and every language on the CPU, keeps whisper-small's + character-level DTW: about 16 ms median and 60 ms P90 on synthetic speech, so + about two single-word deletes in five are clean. A multilingual CTC model + (`facebook/omniASR-CTC-300M`) would cover the other languages, at a cost per + language family nobody has measured yet (issue #948). +- **The aligner is GPU-only.** On the CPU fallback (16 threads, Ryzen 7 5800X) + it would add 29% in English and 58% in French: the French model is a 24-layer + wav2vec2 large, about half of whisper-small's cost on its own. diff --git a/tools/stt-eval/word-timing/README.md b/tools/stt-eval/word-timing/README.md index b49938d6f..064e3c9bb 100644 --- a/tools/stt-eval/word-timing/README.md +++ b/tools/stt-eval/word-timing/README.md @@ -33,7 +33,9 @@ node validate-ref.mjs # checks the reference against the audio's energy ```sh node run-helper.mjs [--cpu] -node evaluate.mjs [--snap ] [--out ] +node run-align.mjs -ctc \ + --model en= --model fr= [--cpu] +node evaluate.mjs -ctc [--snap ] [--ctc ] [--out ] node summarize.mjs "Before=" "After=" ``` @@ -45,20 +47,43 @@ node summarize.mjs "Before=" "After=" - `evaluate.mjs` parses the responses as `whisperServer.ts` does and runs the post-pass: the repo's `electron/stt/snapWordBoundaries.ts` by default, or the file given with `--snap`. It reports three stages: `raw`, `post` (no speech - intervals) and `post+vad` (what the app ships). + intervals) and `post+vad` (what the app ships without an aligner). +- `run-align.mjs` adds the CTC aligner's letter scores (`/emissions`) to saved + responses, without running whisper again, for the language each clip was + detected in. On those, `evaluate.mjs` adds `ctc` (the aligner, + `electron/stt/ctcAlign.ts` or `--ctc`) and `ctc+vad` (what the app ships with + an aligner), and the aligner's share of the runtime. +- The aligner GGUFs come from `scripts/convert-wav2vec2-gguf.mjs` (sources in + its header), or from the app's cache: `stt-models/ctc-aligner/`. - A baseline is the same two commands on an older helper and post-pass: build the helper at that revision, and pass `git show :electron/stt/snapWordBoundaries.ts` saved to a file as `--snap`. ## Real speech +**LibriSpeech, against forced-aligned word times** (English read speech; the +Montreal Forced Aligner reference is itself 10 to 20 ms from a human's): + +```sh +# test-clean.tar.gz from https://www.openslr.org/12, librispeech_alignments.zip +# from https://zenodo.org/records/2619474 (both CC-BY-4.0, dev-time only) +export OSC_WORD_TIMING_DATA= +node make-librispeech.mjs # 2 clips x 40 speakers, ~46 min +node run-helper.mjs ls --condition clean +node run-align.mjs ls ls-ctc --model en= +node evaluate.mjs ls-ctc +``` + +**A take of your own**, with no ground truth: + ```sh -node real-check.mjs take.wav # 16 kHz mono s16 +node real-check.mjs take.wav [--align ] # 16 kHz mono s16 ``` -No ground truth there: it prints how far each word's first-token time lies after -its start, and where each phrase's first word lands against the VAD onset, raw -and after the post-pass. +It prints how far each word's first-token time lies after its start, and where +each phrase's first word lands against the VAD onset, raw and after the +post-pass. With `--align`, every word is listed with its DTW and aligner +times, so the boundaries that moved can be checked by ear or on a spectrogram. ## Reading the numbers @@ -68,3 +93,6 @@ and after the post-pass. at most 20 ms of its neighbours is cut. - Synthetic speech flatters every method. Use the harness to rank approaches, and check a winner on real speech. +- The aligner's calibration (`START_OFFSET_SEC` and the others in + `ctcAlign.ts`) was fitted on the TTS corpus and checked on LibriSpeech; refit + both before changing a model. diff --git a/tools/stt-eval/word-timing/evaluate.mjs b/tools/stt-eval/word-timing/evaluate.mjs index b923c5d6a..8a08b718f 100644 --- a/tools/stt-eval/word-timing/evaluate.mjs +++ b/tools/stt-eval/word-timing/evaluate.mjs @@ -1,11 +1,13 @@ // Scores helper output against the TTS reference, for each pipeline stage. -// Usage: node evaluate.mjs [--snap ] [--out ] [--quiet] +// Usage: node evaluate.mjs [--snap ] [--ctc ] [--out ] [--quiet] // reads /results/raw//*.json, writes /results/.json and // .txt, and prints the table. // Stages: // raw helper words, parsed as whisperServer.ts transcribeImpl does // post the post-pass without speech intervals // post+vad the post-pass with the helper's `speech`, as the app runs it +// ctc the CTC aligner's times (when run-align.mjs saved emissions) +// ctc+vad ctc, then the post-pass with `speech`, as the app runs it with an aligner // `--snap` swaps the post-pass (default: the repo's snapWordBoundaries.ts), so a // candidate is scored exactly like the shipped code. It takes the current // `anchorWordsOnSpeech(words, speech)` or the pre-#948 @@ -24,6 +26,8 @@ const flag = (name, dflt) => { const snapPath = path.resolve( flag("--snap", path.join(REPO, "electron/stt/snapWordBoundaries.ts")), ); +const ctcPath = path.resolve(flag("--ctc", path.join(REPO, "electron/stt/ctcAlign.ts"))); +const ctc = await import(pathToFileURL(ctcPath).href); const outName = flag("--out", tag); const quiet = rest.includes("--quiet"); const mod = await import(pathToFileURL(snapPath).href); @@ -96,7 +100,7 @@ const inside = (env, r0, r1, c0, c1) => { return b > a ? { sec: b - a, aud: audible(env, a, b) } : { sec: 0, aud: 0 }; }; -const STAGES = ["raw", "post", "post+vad"]; +let STAGES = ["raw", "post", "post+vad"]; const acc = {}; const push = (k, v) => { acc[k] ??= []; @@ -105,6 +109,7 @@ const push = (k, v) => { const counts = {}; const inc = (k, v = 1) => (counts[k] = (counts[k] ?? 0) + v); const perClip = []; +let ctcMs = 0; const rawDir = path.join(RESULTS, "raw", tag); for (const file of readdirSync(rawDir) @@ -121,6 +126,17 @@ for (const file of readdirSync(rawDir) const full = postPass(rawWords, samples, speech); const tsMs = performance.now() - t0; const variants = { raw: rawWords, post, "post+vad": full }; + const emissions = json.emissions ? ctc.parseEmissions(json.emissions) : null; + if (json.emissions && !emissions) + console.warn(`${file}: unreadable emissions, CTC stages skipped`); + if (emissions && speech) { + const t1 = performance.now(); + const aligned = ctc.alignWordsOnEmissions(rawWords, speech, emissions); + variants.ctc = aligned; + variants["ctc+vad"] = postPass(aligned, samples, speech); + ctcMs += performance.now() - t1; + STAGES = ["raw", "post", "post+vad", "ctc", "ctc+vad"]; + } const R = ref.words.map((w, i) => ({ ...w, i, n: norm(w.text) })).filter((w) => w.n); const H = rawWords.map((w, j) => ({ j, n: norm(w.word) })).filter((w) => w.n); @@ -143,6 +159,7 @@ for (const file of readdirSync(rawDir) audioSec: ref.durationSec, helperSec: json.timing?.elapsed_s, wallMs: json.wallMs, + alignSec: json.emissions?.elapsed_s, tsMs, refWords: R.length, matched: pairs.length, @@ -291,11 +308,15 @@ const rt = perClip.reduce( helper: a.helper + (c.helperSec ?? 0), wall: a.wall + c.wallMs / MS, ts: a.ts + c.tsMs / MS, + align: a.align + (c.alignSec ?? 0), }), - { audio: 0, helper: 0, wall: 0, ts: 0 }, + { audio: 0, helper: 0, wall: 0, ts: 0, align: 0 }, ); lines.push( - `\nruntime: ${perClip.length} clips, ${(rt.audio / 60).toFixed(1)} min audio; helper ${rt.helper.toFixed(1)} s (RTF ${(rt.helper / rt.audio).toFixed(3)}), per clip mean ${(rt.helper / perClip.length).toFixed(2)} s; TS post-pass ${(rt.ts * MS).toFixed(0)} ms total (${((rt.ts * MS) / perClip.length).toFixed(1)} ms/clip, both stages)`, + `\nruntime: ${perClip.length} clips, ${(rt.audio / 60).toFixed(1)} min audio; helper ${rt.helper.toFixed(1)} s (RTF ${(rt.helper / rt.audio).toFixed(3)}), per clip mean ${(rt.helper / perClip.length).toFixed(2)} s; TS post-pass ${(rt.ts * MS).toFixed(0)} ms total (${((rt.ts * MS) / perClip.length).toFixed(1)} ms/clip, both stages)` + + (rt.align + ? `; aligner emissions ${rt.align.toFixed(1)} s (+${((100 * rt.align) / rt.helper).toFixed(0)}% of the helper), CTC Viterbi ${ctcMs.toFixed(0)} ms` + : ""), ); writeFileSync(path.join(RESULTS, `${outName}.txt`), lines.join("\n") + "\n"); if (!quiet) console.log(lines.join("\n")); diff --git a/tools/stt-eval/word-timing/make-librispeech.mjs b/tools/stt-eval/word-timing/make-librispeech.mjs new file mode 100644 index 000000000..b5e51ea91 --- /dev/null +++ b/tools/stt-eval/word-timing/make-librispeech.mjs @@ -0,0 +1,125 @@ +// Real-speech corpus for the harness: LibriSpeech test-clean read speech with +// Montreal Forced Aligner word times as the reference (silver labels, about +// 10-20 ms from a human's). +// Usage: OSC_WORD_TIMING_DATA= node make-librispeech.mjs \ +// [--max-min 60] +// -> the same layout as make-corpus.mjs (clips/.wav + .ref.json, +// corpus-manifest.json), clean condition only, so run-helper.mjs +// --condition clean, run-align.mjs and evaluate.mjs work unchanged. +// Consecutive utterances of a chapter are joined into clips of up to 40 s, so a +// clip reads like narration: phrases separated by the speaker's own pauses. +// Two clips per speaker (the first chapter), so the corpus spans all 40 voices. +// Sources (CC-BY-4.0, dev-time only, never shipped): https://www.openslr.org/12 +// (test-clean.tar.gz) and https://zenodo.org/records/2619474. +import { execFileSync } from "node:child_process"; +import { mkdirSync, readdirSync, readFileSync, writeFileSync } from "node:fs"; +import path from "node:path"; +import { CLIPS, DATA, FFMPEG, SR } from "./lib.mjs"; + +const [audioRoot, alignRoot, ...rest] = process.argv.slice(2); +if (!alignRoot) + throw new Error( + "usage: node make-librispeech.mjs [--max-min 60]", + ); +const i = rest.indexOf("--max-min"); +const maxSec = (i >= 0 ? Number(rest[i + 1]) : 60) * 60; +const CLIP_SEC = 40; +const PER_SPEAKER = 2; + +/** Word intervals of a TextGrid's "words" tier; empty text is silence. */ +function textGridWords(file) { + const src = readFileSync(file, "utf8"); + const tier = src.slice(src.indexOf('name = "words"')); + const end = tier.indexOf("item [2]"); + const body = end > 0 ? tier.slice(0, end) : tier; + const words = []; + for (const m of body.matchAll(/xmin = ([\d.]+)\s+xmax = ([\d.]+)\s+text = "([^"]*)"/g)) { + if (m[3].trim()) words.push({ text: m[3].trim(), start: Number(m[1]), end: Number(m[2]) }); + } + return words; +} + +mkdirSync(CLIPS, { recursive: true }); +const clips = []; +let total = 0; +outer: for (const spk of readdirSync(alignRoot).sort()) { + for (const chap of readdirSync(path.join(alignRoot, spk)).sort()) { + const utts = readdirSync(path.join(alignRoot, spk, chap)) + .filter((f) => f.endsWith(".TextGrid")) + .map((f) => f.replace(/\.TextGrid$/, "")) + .sort(); + let group = []; + let groupSec = 0; + const flush = () => { + if (!group.length) return; + const id = `ls-${group[0]}`; + const pcm = []; + const words = []; + let at = 0; + for (const u of group) { + const raw = execFileSync( + FFMPEG, + [ + "-v", + "error", + "-i", + path.join(audioRoot, spk, chap, `${u}.flac`), + "-ac", + "1", + "-ar", + String(SR), + "-f", + "s16le", + "-", + ], + { maxBuffer: 1 << 28 }, + ); + for (const w of textGridWords(path.join(alignRoot, spk, chap, `${u}.TextGrid`))) + words.push({ text: w.text, start: at + w.start, end: at + w.end }); + pcm.push(raw); + at += raw.length / 2 / SR; + } + const data = Buffer.concat(pcm); + const header = Buffer.alloc(44); + header.write("RIFF", 0); + header.writeUInt32LE(36 + data.length, 4); + header.write("WAVEfmt ", 8); + header.writeUInt32LE(16, 16); + header.writeUInt16LE(1, 20); + header.writeUInt16LE(1, 22); + header.writeUInt32LE(SR, 24); + header.writeUInt32LE(SR * 2, 28); + header.writeUInt16LE(2, 32); + header.writeUInt16LE(16, 34); + header.write("data", 36); + header.writeUInt32LE(data.length, 40); + writeFileSync(path.join(CLIPS, `${id}.wav`), Buffer.concat([header, data])); + writeFileSync( + path.join(CLIPS, `${id}.ref.json`), + JSON.stringify({ id, lang: "en", engine: "mfa", durationSec: at, words }, null, 1), + ); + clips.push({ id, lang: "en", durationSec: at, words: words.length }); + total += at; + console.log(id, at.toFixed(1), "s", words.length, "words"); + group = []; + groupSec = 0; + }; + const before = clips.length; + for (const u of utts) { + if (clips.length - before >= PER_SPEAKER) break; + const w = textGridWords(path.join(alignRoot, spk, chap, `${u}.TextGrid`)); + const sec = w.length ? w[w.length - 1].end + 0.3 : 0; + if (groupSec + sec > CLIP_SEC) flush(); + group.push(u); + groupSec += sec; + } + if (clips.length - before < PER_SPEAKER) flush(); + if (total >= maxSec) break outer; + break; + } +} +writeFileSync( + path.join(DATA, "corpus-manifest.json"), + JSON.stringify({ totalSec: total, source: "LibriSpeech test-clean + MFA", clips }, null, 1), +); +console.log(`${clips.length} clips, ${(total / 60).toFixed(1)} min`); diff --git a/tools/stt-eval/word-timing/real-check.mjs b/tools/stt-eval/word-timing/real-check.mjs index 09a3d7865..e60788f7f 100644 --- a/tools/stt-eval/word-timing/real-check.mjs +++ b/tools/stt-eval/word-timing/real-check.mjs @@ -1,36 +1,58 @@ // Sanity check on real speech, which has no ground truth: the VAD onset is the // only boundary we can trust there. // Usage: node real-check.mjs [--cpu] [--snap ] +// [--align ] // (16 kHz mono s16 WAV, e.g. ffmpeg -i take.webm -ar 16000 -ac 1 take.wav) // Prints how far each word's first-token time (`anchor`) lies after its start // (the one-token lag the helper corrects), and, per VAD stretch, where its first // word starts relative to the onset: as the helper reports it, and after the -// post-pass. Writes the raw response to /results/real-.json. -import { mkdirSync, writeFileSync } from "node:fs"; +// post-pass. With `--align`, the CTC aligner (electron/stt/ctcAlign.ts) runs too, +// and every word is listed with both times so the boundaries can be checked by +// ear or on a spectrogram. Writes the raw response (and the aligner's words) to +// /results/real-.json. +import { mkdirSync, readFileSync, writeFileSync } from "node:fs"; import path from "node:path"; import { pathToFileURL } from "node:url"; -import { REPO, RESULTS, startHelper, transcribe } from "./lib.mjs"; +import { + alignWordsOnEmissions, + emissionRegions, + parseEmissions, +} from "../../../electron/stt/ctcAlign.ts"; +import { REPO, RESULTS, readWav, SR, startHelper, transcribe } from "./lib.mjs"; const [exe, wav, ...rest] = process.argv.slice(2); if (!exe || !wav) throw new Error( - "usage: node real-check.mjs [--cpu] [--snap ]", + "usage: node real-check.mjs [--cpu] [--snap ] [--align ]", ); -const snapAt = rest.indexOf("--snap"); +const flag = (name) => { + const i = rest.indexOf(name); + return i >= 0 ? rest[i + 1] : undefined; +}; const snapPath = path.resolve( - snapAt >= 0 ? rest[snapAt + 1] : path.join(REPO, "electron/stt/snapWordBoundaries.ts"), + flag("--snap") ?? path.join(REPO, "electron/stt/snapWordBoundaries.ts"), ); +const alignModel = flag("--align") && path.resolve(flag("--align")); const { anchorWordsOnSpeech } = await import(pathToFileURL(snapPath).href); const { base, stop } = await startHelper(exe, { cpu: rest.includes("--cpu") }); let json; +let emissions; try { ({ json } = await transcribe(base, wav)); + if (alignModel) { + const speech = json.speech.map((s) => ({ startSec: s.start, endSec: s.end })); + const form = new FormData(); + form.set("file", new Blob([readFileSync(wav)], { type: "audio/wav" }), path.basename(wav)); + form.set("model", alignModel); + form.set("regions", JSON.stringify(emissionRegions(speech, readWav(wav).length / SR))); + const res = await fetch(`${base}/emissions`, { method: "POST", body: form }); + emissions = await res.json(); + if (!res.ok) throw new Error(`/emissions: HTTP ${res.status} ${JSON.stringify(emissions)}`); + } } finally { stop(); } -mkdirSync(RESULTS, { recursive: true }); -writeFileSync(path.join(RESULTS, `real-${path.basename(wav, ".wav")}.json`), JSON.stringify(json)); const isP = (w) => /^[\p{P}\p{S}]+$/u.test(w.word); const raw = json.segments @@ -44,6 +66,18 @@ const raw = json.segments .filter((w) => w.word); const speech = json.speech.map((s) => ({ startSec: s.start, endSec: s.end })); const post = anchorWordsOnSpeech(raw, speech); +const parsed = emissions ? parseEmissions(emissions) : null; +if (emissions && !parsed) + throw new Error("/emissions answered something that is not letter scores (vocab/regions/sizes)"); +const aligned = parsed + ? anchorWordsOnSpeech(alignWordsOnEmissions(raw, speech, parsed), speech) + : null; +mkdirSync(RESULTS, { recursive: true }); +writeFileSync( + path.join(RESULTS, `real-${path.basename(wav, ".wav")}.json`), + JSON.stringify({ ...json, phase1: post, aligned }), +); + const lag = raw .filter((w) => !isP(w)) .map((w) => w.anchorSec - w.startSec) @@ -63,3 +97,23 @@ for (const [i, s] of speech.entries()) { ); while (k < raw.length && raw[k].anchorSec < tail) k++; } +if (aligned) { + const moved = []; + console.log("\nword DTW (s) aligner (s) start / end moved"); + raw.forEach((w, j) => { + const a = post[j]; + const b = aligned[j]; + if (!isP(w)) moved.push(Math.abs(b.startSec - a.startSec), Math.abs(b.endSec - a.endSec)); + console.log( + `${w.word.padEnd(20)} ${a.startSec.toFixed(3)}-${a.endSec.toFixed(3)} ${b.startSec.toFixed(3)}-${b.endSec.toFixed(3)} ${ms(b.startSec - a.startSec)} / ${ms(b.endSec - a.endSec)}`, + ); + }); + moved.sort((x, y) => x - y); + const m = (p) => (moved[Math.floor(moved.length * p)] * 1000).toFixed(0); + const spread = moved.length + ? `median ${m(0.5)} ms, p90 ${m(0.9)} ms` + : "no word to compare (punctuation only)"; + console.log( + `\naligner vs DTW, |moved| per boundary: ${spread} (${emissions.elapsed_s.toFixed(2)} s on ${emissions.device})`, + ); +} diff --git a/tools/stt-eval/word-timing/run-align.mjs b/tools/stt-eval/word-timing/run-align.mjs new file mode 100644 index 000000000..9f372aa51 --- /dev/null +++ b/tools/stt-eval/word-timing/run-align.mjs @@ -0,0 +1,64 @@ +// Adds the CTC aligner's emissions to helper responses already saved by +// run-helper.mjs, so evaluate.mjs can score the aligner without re-running whisper. +// Usage: node run-align.mjs \ +// --model en= --model fr= [--cpu] [--only ] +// reads /results/raw//*.json (needs `speech`: run with the VAD model) +// writes /results/raw//*.json with `emissions` (the /emissions answer +// + wallMs) next to the /inference answer +// The helper is started once and asked like the app does (electron/stt/index.ts): +// the regions come from emissionRegions(), the model from the detected language. +import { existsSync, mkdirSync, readdirSync, readFileSync, writeFileSync } from "node:fs"; +import path from "node:path"; +import { emissionRegions } from "../../../electron/stt/ctcAlign.ts"; +import { CLIPS, RESULTS, readWav, SR, startHelper } from "./lib.mjs"; + +const [fromTag, toTag, exe, ...rest] = process.argv.slice(2); +if (!exe) + throw new Error( + "usage: node run-align.mjs --model en= [--model fr=] [--cpu] [--only ]", + ); +const models = {}; +let only = ""; +for (let i = 0; i < rest.length; i++) { + if (rest[i] === "--model") { + const [lang, file] = rest[++i].split("="); + models[lang] = path.resolve(file); + } + if (rest[i] === "--only") only = rest[++i]; +} + +const { base, stop } = await startHelper(exe, { cpu: rest.includes("--cpu") }); +try { + const inDir = path.join(RESULTS, "raw", fromTag); + const outDir = path.join(RESULTS, "raw", toTag); + mkdirSync(outDir, { recursive: true }); + for (const file of readdirSync(inDir).filter((f) => f.endsWith(".json"))) { + if (only && !file.includes(only)) continue; + const out = path.join(outDir, file); + if (existsSync(out)) continue; + const json = JSON.parse(readFileSync(path.join(inDir, file), "utf8")); + const model = models[json.detected_language]; + if (!model || !json.speech) { + writeFileSync(out, JSON.stringify(json)); + continue; + } + const [id, cond] = file.replace(/\.json$/, "").split("."); + const wav = path.join(CLIPS, cond === "clean" ? `${id}.wav` : `${id}.noisy.wav`); + const bytes = readFileSync(wav); + const durationSec = readWav(wav).length / SR; + const speech = json.speech.map((s) => ({ startSec: s.start, endSec: s.end })); + const form = new FormData(); + form.set("file", new Blob([bytes], { type: "audio/wav" }), path.basename(wav)); + form.set("model", model); + form.set("regions", JSON.stringify(emissionRegions(speech, durationSec))); + const t0 = performance.now(); + const res = await fetch(`${base}/emissions`, { method: "POST", body: form }); + const wallMs = performance.now() - t0; + const emissions = await res.json(); + if (!res.ok) throw new Error(`${file}: HTTP ${res.status} ${JSON.stringify(emissions)}`); + writeFileSync(out, JSON.stringify({ ...json, emissions: { ...emissions, wallMs } })); + console.log(file, `${emissions.elapsed_s.toFixed(3)} s on ${emissions.device}`); + } +} finally { + stop(); +}