Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion electron/native/whisper-stt/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
433 changes: 433 additions & 0 deletions electron/native/whisper-stt/src/ctc_aligner.cpp

Large diffs are not rendered by default.

37 changes: 37 additions & 0 deletions electron/native/whisper-stt/src/ctc_aligner.h
Original file line number Diff line number Diff line change
@@ -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 <memory>
#include <string>
#include <vector>

struct CtcModel;

struct CtcModelInfo {
std::vector<std::string> vocab; // token id -> text; "|" is the word delimiter
std::vector<std::string> 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<CtcModel, CtcModelDeleter>;

// 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<float>& out,
int& frames, std::string& err);
107 changes: 101 additions & 6 deletions electron/native/whisper-stt/src/main.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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 <algorithm>
#include <atomic>
Expand Down Expand Up @@ -259,6 +260,21 @@ double to_original_sec(int64_t cs, const std::vector<Kept>& 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<const unsigned char*>(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) {
Expand Down Expand Up @@ -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<float>& 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;

Expand All @@ -404,23 +422,29 @@ int main(int argc, char** argv) {
out.write(file_entry.content.data(),
static_cast<std::streamsize>(file_entry.content.size()));
}
std::vector<float> pcm;
int sample_rate = 0, channels = 0;
const bool ok = read_wav_pcm16(tmp_wav, pcm, sample_rate, channels);
std::error_code ec;
std::filesystem::remove(tmp_wav, ec);
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<float> pcm;
if (!read_upload(req, res, pcm)) return;

// language param
std::string language = "auto";
Expand Down Expand Up @@ -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<float> 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<std::mutex> 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<int64_t>(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<int64_t>(std::llround(r[0].get<double>() * 16000.0), 0, n_pcm);
const int64_t to = std::clamp<int64_t>(std::llround(r[1].get<double>() * 16000.0), from, n_pcm);
std::vector<float> lp;
int frames = 0;
std::string err;
if (!ctc_emissions(*aligner, pcm.data() + from, static_cast<size_t>(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<double>(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) {
Expand Down
Loading
Loading