Skip to content

Commit 2d048a1

Browse files
committed
feat(stt): re-time whisper's words on a CTC forced aligner
A second alignment pass (issue #948, phase 3): a wav2vec2 CTC model scores every 20 ms frame of the speech against every letter, and a Viterbi pass forces whisper's own words through those scores. English uses wav2vec2-base-960h (109 MB), French wav2vec2-large-xlsr-53-french (348 MB), both Apache-2.0, converted to Q8_0 GGUF and run on ggml inside whisper-stt-server, on the GPU whisper uses. The helper only scores (POST /emissions); spelling, Viterbi, calibration and fallback live in electron/stt/ctcAlign.ts. Other languages, a failed download or an older helper keep the phase 1 times. TTS corpus, clean: inner start median 31 -> 14 ms, P90 125 -> 35 ms, within 50 ms 64% -> 96%; clean single-word cuts 19% -> 50%; phrase delete 89% -> 96%. LibriSpeech vs MFA: inner median 40 -> 15 ms, within 50 ms 55% -> 92%, clean cuts 14% -> 44%. Extra runtime about +10% on Vulkan.
1 parent c9c5642 commit 2d048a1

22 files changed

Lines changed: 2194 additions & 78 deletions

‎electron/native/whisper-stt/CMakeLists.txt‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -144,9 +144,10 @@ endif()
144144

145145
FetchContent_MakeAvailable(whisper httplib json)
146146

147-
add_executable(whisper-stt-server src/main.cpp)
147+
add_executable(whisper-stt-server src/main.cpp src/ctc_aligner.cpp)
148148
target_link_libraries(whisper-stt-server PRIVATE
149149
whisper
150+
ggml
150151
httplib::httplib
151152
nlohmann_json::nlohmann_json
152153
)

‎electron/native/whisper-stt/src/ctc_aligner.cpp‎

Lines changed: 433 additions & 0 deletions
Large diffs are not rendered by default.
Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
// CTC acoustic model (wav2vec2 fine-tuned for CTC) run on ggml, for the word
2+
// aligner's second pass. The helper only computes the per-frame log-probs; the
3+
// forced alignment of whisper's words over them lives on the Node side
4+
// (electron/stt/ctcAlign.ts). See transcription-and-captions.md § Word-level alignment.
5+
#pragma once
6+
7+
#include <memory>
8+
#include <string>
9+
#include <vector>
10+
11+
struct CtcModel;
12+
13+
struct CtcModelInfo {
14+
std::vector<std::string> vocab; // token id -> text; "|" is the word delimiter
15+
std::vector<std::string> languages; // what the model was fine-tuned on
16+
int blank = 0;
17+
int stride = 320; // samples per output frame (20 ms at 16 kHz)
18+
int receptive_field = 400; // samples the first frame sees
19+
};
20+
21+
struct CtcModelDeleter { void operator()(CtcModel* m) const; };
22+
using CtcModelPtr = std::unique_ptr<CtcModel, CtcModelDeleter>;
23+
24+
// Loads a GGUF written by scripts/convert-wav2vec2-gguf.mjs. On the GPU when
25+
// `use_gpu` and a GPU device is registered, else on the CPU. Null + `err` on failure.
26+
CtcModelPtr ctc_load(const std::string& path, bool use_gpu, int threads, std::string& err);
27+
28+
const CtcModelInfo& ctc_info(const CtcModel& model);
29+
30+
// Name of the device the weights live on ("Vulkan0", "MTL0", "CPU").
31+
std::string ctc_device(const CtcModel& model);
32+
33+
// Log-softmax emissions for 16 kHz mono `pcm`: `frames` rows of vocab-size
34+
// floats. Frame i sees samples [i * stride, i * stride + receptive_field).
35+
// Long inputs run in overlapping windows, so memory stays bounded.
36+
bool ctc_emissions(CtcModel& model, const float* pcm, size_t n, std::vector<float>& out,
37+
int& frames, std::string& err);

‎electron/native/whisper-stt/src/main.cpp‎

Lines changed: 101 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
// this is a belt-and-braces guarantee against a future bug or parallel invoker).
3737

3838
#include "whisper.h"
39+
#include "ctc_aligner.h"
3940

4041
#include <algorithm>
4142
#include <atomic>
@@ -259,6 +260,21 @@ double to_original_sec(int64_t cs, const std::vector<Kept>& kept) {
259260
return (kept.back().from + kept.back().len) / 16000.0;
260261
}
261262

263+
std::string base64(const void* data, size_t n) {
264+
static const char* abc = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
265+
const auto* p = static_cast<const unsigned char*>(data);
266+
std::string out;
267+
out.reserve((n + 2) / 3 * 4);
268+
for (size_t i = 0; i < n; i += 3) {
269+
const uint32_t v = (p[i] << 16) | (i + 1 < n ? p[i + 1] << 8 : 0) | (i + 2 < n ? p[i + 2] : 0);
270+
out += abc[(v >> 18) & 63];
271+
out += abc[(v >> 12) & 63];
272+
out += i + 1 < n ? abc[(v >> 6) & 63] : '=';
273+
out += i + 2 < n ? abc[v & 63] : '=';
274+
}
275+
return out;
276+
}
277+
262278
} // namespace
263279

264280
int main(int argc, char** argv) {
@@ -376,13 +392,15 @@ int main(int argc, char** argv) {
376392
res.set_content(probe.dump(), "application/json");
377393
});
378394

379-
// POST /inference — multipart form with `file` (WAV) + `language` + `response_format`.
380-
svr.Post("/inference", [&](const httplib::Request& req, httplib::Response& res) {
395+
// The upload of /inference and /emissions: a 16 kHz mono PCM16 WAV in the
396+
// multipart field `file`. False after answering 400.
397+
const auto read_upload = [](const httplib::Request& req, httplib::Response& res,
398+
std::vector<float>& pcm) -> bool {
381399
auto it = req.files.find("file");
382400
if (it == req.files.end()) {
383401
res.status = 400;
384402
res.set_content(R"({"error":"missing 'file' form field"})", "application/json");
385-
return;
403+
return false;
386404
}
387405
const auto& file_entry = it->second;
388406

@@ -404,23 +422,29 @@ int main(int argc, char** argv) {
404422
out.write(file_entry.content.data(),
405423
static_cast<std::streamsize>(file_entry.content.size()));
406424
}
407-
std::vector<float> pcm;
408425
int sample_rate = 0, channels = 0;
409426
const bool ok = read_wav_pcm16(tmp_wav, pcm, sample_rate, channels);
410427
std::error_code ec;
411428
std::filesystem::remove(tmp_wav, ec);
412429
if (!ok) {
413430
res.status = 400;
414431
res.set_content(R"({"error":"failed to parse WAV"})", "application/json");
415-
return;
432+
return false;
416433
}
417434
if (sample_rate != 16000 || channels != 1) {
418435
res.status = 400;
419436
res.set_content(
420437
R"({"error":"expected 16 kHz mono PCM16 WAV"})",
421438
"application/json");
422-
return;
439+
return false;
423440
}
441+
return true;
442+
};
443+
444+
// POST /inference — multipart form with `file` (WAV) + `language` + `response_format`.
445+
svr.Post("/inference", [&](const httplib::Request& req, httplib::Response& res) {
446+
std::vector<float> pcm;
447+
if (!read_upload(req, res, pcm)) return;
424448

425449
// language param
426450
std::string language = "auto";
@@ -669,6 +693,77 @@ int main(int argc, char** argv) {
669693
res.set_content(reply.dump(), "application/json");
670694
});
671695

696+
// POST /emissions — the CTC aligner's acoustic pass (issue #948, phase 3).
697+
// Multipart form: `file` (the same WAV as /inference), `model` (path of a
698+
// wav2vec2 GGUF, see ctc_aligner.h) and `regions` (JSON [[start_s, end_s], ...]).
699+
// Answers the model's vocabulary and, per region, base64 float32 log-probs
700+
// [frames x vocab]; frame i of a region sees the audio from
701+
// `start + i * stride_s` for `receptive_s`. The forced alignment itself runs
702+
// on the Node side (electron/stt/ctcAlign.ts). The model stays loaded until a
703+
// request names another one.
704+
CtcModelPtr aligner;
705+
std::string aligner_path;
706+
svr.Post("/emissions", [&](const httplib::Request& req, httplib::Response& res) {
707+
std::vector<float> pcm;
708+
if (!read_upload(req, res, pcm)) return;
709+
const std::string model = req.get_file_value("model").content;
710+
nlohmann::json regions = nlohmann::json::parse(req.get_file_value("regions").content, nullptr, false);
711+
if (model.empty() || !regions.is_array()) {
712+
res.status = 400;
713+
res.set_content(R"({"error":"need 'model' and a JSON 'regions' array"})", "application/json");
714+
return;
715+
}
716+
const std::lock_guard<std::mutex> lk(infer_mu);
717+
const auto t0 = std::chrono::steady_clock::now();
718+
if (!aligner || aligner_path != model) {
719+
aligner.reset();
720+
std::string err;
721+
aligner = ctc_load(model, cparams.use_gpu, threads, err);
722+
if (!aligner) {
723+
log("aligner: " + err);
724+
res.status = 500;
725+
res.set_content(nlohmann::json{{"error", "aligner: " + err}}.dump(), "application/json");
726+
return;
727+
}
728+
aligner_path = model;
729+
log("aligner loaded on " + ctc_device(*aligner) + ": " + model);
730+
}
731+
const CtcModelInfo& info = ctc_info(*aligner);
732+
nlohmann::json out_regions = nlohmann::json::array();
733+
const int64_t n_pcm = static_cast<int64_t>(pcm.size());
734+
for (const auto& r : regions) {
735+
if (!r.is_array() || r.size() != 2 || !r[0].is_number() || !r[1].is_number()) continue;
736+
const int64_t from = std::clamp<int64_t>(std::llround(r[0].get<double>() * 16000.0), 0, n_pcm);
737+
const int64_t to = std::clamp<int64_t>(std::llround(r[1].get<double>() * 16000.0), from, n_pcm);
738+
std::vector<float> lp;
739+
int frames = 0;
740+
std::string err;
741+
if (!ctc_emissions(*aligner, pcm.data() + from, static_cast<size_t>(to - from), lp, frames, err)) {
742+
log("aligner: " + err);
743+
res.status = 500;
744+
res.set_content(nlohmann::json{{"error", "aligner: " + err}}.dump(), "application/json");
745+
return;
746+
}
747+
out_regions.push_back({
748+
{"start", from / 16000.0},
749+
{"frames", frames},
750+
{"logprobs", base64(lp.data(), lp.size() * sizeof(float))},
751+
});
752+
}
753+
const double elapsed_s = std::chrono::duration<double>(std::chrono::steady_clock::now() - t0).count();
754+
nlohmann::json reply = {
755+
{"vocab", info.vocab},
756+
{"blank", info.blank},
757+
{"languages", info.languages},
758+
{"stride_s", info.stride / 16000.0},
759+
{"receptive_s", info.receptive_field / 16000.0},
760+
{"device", ctc_device(*aligner)},
761+
{"elapsed_s", elapsed_s},
762+
{"regions", std::move(out_regions)},
763+
};
764+
res.set_content(reply.dump(), "application/json");
765+
});
766+
672767
// ---- bind + listen ----
673768
int bound_port = port;
674769
if (bound_port == 0) {

0 commit comments

Comments
 (0)