Xenobd/whisper.cpp
0
1#include "common.h"2#include "common-sdl.h"3#include "whisper.h"4#include "json.hpp"5 6#include <cassert>7#include <chrono>8#include <cstdio>9#include <deque>10#include <iostream>11#include <set>12#include <string>13#include <thread>14#include <vector>15 16using json = nlohmann::json;17 18// command-line parameters19struct whisper_params {20 int32_t n_threads = std::min(4, (int32_t) std::thread::hardware_concurrency());21 int32_t prompt_ms = 5000;22 int32_t command_ms = 8000;23 int32_t capture_id = -1;24 int32_t max_tokens = 32;25 int32_t audio_ctx = 0;26 27 float vad_thold = 0.6f;28 float freq_thold = 100.0f;29 30 bool translate = false;31 bool print_special = false;32 bool print_energy = false;33 bool use_gpu = true;34 bool flash_attn = false;35 36 std::string language = "en";37 std::string model = "models/ggml-base.en.bin";38};39struct command {40 std::vector<whisper_token> tokens;41 std::string plaintext;42};43struct commandset {44 std::vector<struct command> commands;45 std::vector<whisper_token> prompt_tokens;46 // TODO: Store longest command?47 // Multi-token commands should have probabilities of subsequent logits48 // given that the prior logit is correct.49 // In this case, all commands must be iterated.50 // This however, is likely highly involved as different tokens51 // almost certainly have different spoken lengths52 // It would also have performance implications equivalent to a beam search53};54 55void whisper_print_usage(int argc, char ** argv, const whisper_params & params);56 57static bool whisper_params_parse(int argc, char ** argv, whisper_params & params) {58 for (int i = 1; i < argc; i++) {59 std::string arg = argv[i];60 61 if (arg == "-h" || arg == "--help") {62 whisper_print_usage(argc, argv, params);63 exit(0);64 }65 else if (arg == "-t" || arg == "--threads") { params.n_threads = std::stoi(argv[++i]); }66 else if (arg == "-pms" || arg == "--prompt-ms") { params.prompt_ms = std::stoi(argv[++i]); }67 else if (arg == "-cms" || arg == "--command-ms") { params.command_ms = std::stoi(argv[++i]); }68 else if (arg == "-c" || arg == "--capture") { params.capture_id = std::stoi(argv[++i]); }69 else if (arg == "-mt" || arg == "--max-tokens") { params.max_tokens = std::stoi(argv[++i]); }70 else if (arg == "-ac" || arg == "--audio-ctx") { params.audio_ctx = std::stoi(argv[++i]); }71 else if (arg == "-vth" || arg == "--vad-thold") { params.vad_thold = std::stof(argv[++i]); }72 else if (arg == "-fth" || arg == "--freq-thold") { params.freq_thold = std::stof(argv[++i]); }73 else if (arg == "-tr" || arg == "--translate") { params.translate = true; }74 else if (arg == "-ps" || arg == "--print-special") { params.print_special = true; }75 else if (arg == "-pe" || arg == "--print-energy") { params.print_energy = true; }76 else if (arg == "-ng" || arg == "--no-gpu") { params.use_gpu = false; }77 else if (arg == "-fa" || arg == "--flash-attn") { params.flash_attn = true; }78 else if (arg == "-l" || arg == "--language") { params.language = argv[++i]; }79 else if (arg == "-m" || arg == "--model") { params.model = argv[++i]; }80 else {81 fprintf(stderr, "error: unknown argument: %s\n", arg.c_str());82 whisper_print_usage(argc, argv, params);83 exit(0);84 }85 }86 87 return true;88}89 90void whisper_print_usage(int /*argc*/, char ** argv, const whisper_params & params) {91 fprintf(stderr, "\n");92 fprintf(stderr, "usage: %s [options]\n", argv[0]);93 fprintf(stderr, "\n");94 fprintf(stderr, "options:\n");95 fprintf(stderr, " -h, --help [default] show this help message and exit\n");96 fprintf(stderr, " -t N, --threads N [%-7d] number of threads to use during computation\n", params.n_threads);97 fprintf(stderr, " -pms N, --prompt-ms N [%-7d] prompt duration in milliseconds\n", params.prompt_ms);98 fprintf(stderr, " -cms N, --command-ms N [%-7d] command duration in milliseconds\n", params.command_ms);99 fprintf(stderr, " -c ID, --capture ID [%-7d] capture device ID\n", params.capture_id);100 fprintf(stderr, " -mt N, --max-tokens N [%-7d] maximum number of tokens per audio chunk\n", params.max_tokens);101 fprintf(stderr, " -ac N, --audio-ctx N [%-7d] audio context size (0 - all)\n", params.audio_ctx);102 fprintf(stderr, " -vth N, --vad-thold N [%-7.2f] voice activity detection threshold\n", params.vad_thold);103 fprintf(stderr, " -fth N, --freq-thold N [%-7.2f] high-pass frequency cutoff\n", params.freq_thold);104 fprintf(stderr, " -tr, --translate [%-7s] translate from source language to english\n", params.translate ? "true" : "false");105 fprintf(stderr, " -ps, --print-special [%-7s] print special tokens\n", params.print_special ? "true" : "false");106 fprintf(stderr, " -pe, --print-energy [%-7s] print sound energy (for debugging)\n", params.print_energy ? "true" : "false");107 fprintf(stderr, " -ng, --no-gpu [%-7s] disable GPU\n", params.use_gpu ? "false" : "true");108 fprintf(stderr, " -fa, --flash-attn [%-7s] flash attention\n", params.flash_attn ? "true" : "false");109 fprintf(stderr, " -l LANG, --language LANG [%-7s] spoken language\n", params.language.c_str());110 fprintf(stderr, " -m FNAME, --model FNAME [%-7s] model path\n", params.model.c_str());111 fprintf(stderr, "\n");112}113static uint64_t wait_for_vad(audio_async & audio, json jparams, const whisper_params & params, uint64_t maxlength_ms, std::vector<float> & pcmf32) {114 using namespace std::chrono;115 uint64_t time_now = time_point_cast<milliseconds>(system_clock::now()).time_since_epoch().count();116 uint64_t start_time = time_now;117 if (jparams.contains("timestamp")) {118 start_time = jparams.at("timestamp");119 }120 if(time_now - start_time < 500) {121 //wait for a backlog of audio122 std::this_thread::sleep_for(milliseconds(500 - (time_now - start_time)));123 time_now = time_point_cast<milliseconds>(system_clock::now()).time_since_epoch().count();124 } else if (time_now - start_time > 1000) {125 audio.get(time_now-start_time, pcmf32);126 size_t max_offset = pcmf32.size() - WHISPER_SAMPLE_RATE;127 for(size_t offset=0;offset < max_offset;offset+=WHISPER_SAMPLE_RATE/10) {128 std::vector<float> audio_chunk(&pcmf32[offset], &pcmf32[offset+WHISPER_SAMPLE_RATE]);129 if(::vad_simple(audio_chunk, WHISPER_SAMPLE_RATE, 1000, params.vad_thold, params.freq_thold, params.print_energy)) {130 pcmf32.resize(offset+WHISPER_SAMPLE_RATE);131 if (offset*1000/WHISPER_SAMPLE_RATE+1000 > maxlength_ms) {132 //remove samples from the beginning133 pcmf32.erase(pcmf32.begin(),pcmf32.end()-(maxlength_ms*WHISPER_SAMPLE_RATE/1000));134 fprintf(stderr, "Shortened samples");135 }136 return start_time + offset*1000/WHISPER_SAMPLE_RATE+1000;137 }138 }139 }140 size_t window_duration = std::max((uint64_t)1000, time_now-start_time);141 audio.get(window_duration, pcmf32);142 while (!::vad_simple(pcmf32, WHISPER_SAMPLE_RATE, 1000, params.vad_thold, params.freq_thold, params.print_energy)) {143 std::this_thread::sleep_for(milliseconds(100));144 time_now = time_point_cast<milliseconds>(system_clock::now()).time_since_epoch().count();145 window_duration = std::max((uint64_t)1000,time_now-start_time);146 audio.get(window_duration, pcmf32);147 }148 if (time_now - start_time > maxlength_ms) {149 audio.get(maxlength_ms, pcmf32);150 } else {151 audio.get(time_now - start_time, pcmf32);152 }153 154 return time_now;155}156 157static json unguided_transcription(struct whisper_context * ctx, audio_async &audio, json jparams, const whisper_params ¶ms) {158 std::vector<whisper_token> prompt_tokens;159 std::vector<float> pcmf32;160 uint64_t unprocessed_audio_timestamp = wait_for_vad(audio, jparams, params, 10000U, pcmf32);161 162 whisper_full_params wparams = whisper_full_default_params(WHISPER_SAMPLING_GREEDY);163 if (jparams.contains("prompt")) {164 // unlikely to see much use. Under normal circumstances, no_context would be set to false165 std::string prompt = jparams.at("prompt");166 prompt_tokens.resize(1024);167 int n = whisper_tokenize(ctx, prompt.c_str(), prompt_tokens.data(), 1024);168 prompt_tokens.resize(n);169 170 wparams.prompt_tokens = prompt_tokens.data();171 wparams.prompt_n_tokens = prompt_tokens.size();172 }173 wparams.print_progress = false;174 wparams.print_special = params.print_special;175 wparams.print_realtime = false;176 wparams.print_timestamps = false;177 wparams.translate = params.translate;178 wparams.no_context = jparams.value("no_context", true);179 wparams.single_segment = true;180 wparams.max_tokens = params.max_tokens;181 wparams.language = params.language.c_str();182 wparams.n_threads = params.n_threads;183 184 wparams.audio_ctx = params.audio_ctx;185 wparams.suppress_nst = true;186 // run the transformer and a single decoding pass187 if (whisper_full(ctx, wparams, pcmf32.data(), pcmf32.size()) != 0) {188 fprintf(stderr, "%s: ERROR: whisper_full() failed\n", __func__);189 throw json{190 {"code", -32803},191 {"message", "ERROR: whisper_full() failed"}192 };193 }194 std::string result = whisper_full_get_segment_text(ctx,0);195 return json {196 {"transcription", result},197 {"timestamp", unprocessed_audio_timestamp}198 };199}200 201// command-list mode202// guide the transcription to match the most likely command from a provided list203static json guided_transcription(struct whisper_context * ctx, audio_async &audio, const whisper_params ¶ms, json jparams, std::vector<struct commandset> commandset_list) {204 struct commandset cs = commandset_list[jparams.value("commandset_index", commandset_list.size()-1)];205 std::vector<float> pcmf32;206 uint64_t unprocessed_audio_timestamp = wait_for_vad(audio, jparams, params, 2000U, pcmf32);207 208 fprintf(stderr, "%s: Speech detected! Processing ...\n", __func__);209 whisper_full_params wparams = whisper_full_default_params(WHISPER_SAMPLING_GREEDY);210 211 wparams.print_progress = false;212 wparams.print_special = params.print_special;213 wparams.print_realtime = false;214 wparams.print_timestamps = false;215 wparams.translate = params.translate;216 wparams.no_context = true;217 wparams.single_segment = true;218 wparams.max_tokens = 1;219 wparams.language = params.language.c_str();220 wparams.n_threads = params.n_threads;221 222 wparams.audio_ctx = params.audio_ctx;223 224 // TODO: Do some time testing. Does an overly long prompt slow down processing?225 // Set up command sets/precompute prompts226 wparams.prompt_tokens = cs.prompt_tokens.data();227 wparams.prompt_n_tokens = cs.prompt_tokens.size();228 // TODO: properly expose as option229 wparams.suppress_nst = true;230 231 // run the transformer and a single decoding pass232 if (whisper_full(ctx, wparams, pcmf32.data(), pcmf32.size()) != 0) {233 fprintf(stderr, "%s: ERROR: whisper_full() failed\n", __func__);234 throw json{235 {"code", -32803},236 {"message", "ERROR: whisper_full() failed"}//TODO: format string (sprintf?)237 };238 }239 240 // estimate command probability241 // NOTE: not optimal242 {243 const auto * logits = whisper_get_logits(ctx);244 245 std::vector<float> probs(whisper_n_vocab(ctx), 0.0f);246 247 // compute probs from logits via softmax248 {249 float max = -1e9;250 for (int i = 0; i < (int) probs.size(); ++i) {251 max = std::max(max, logits[i]);252 }253 254 float sum = 0.0f;255 for (int i = 0; i < (int) probs.size(); ++i) {256 probs[i] = expf(logits[i] - max);257 sum += probs[i];258 }259 260 for (int i = 0; i < (int) probs.size(); ++i) {261 probs[i] /= sum;262 }263 }264 265 std::vector<std::pair<float, int>> probs_id;266 267 // In my testing, the most verbose token is always the desired.268 // TODO: Trim commandset struct once efficacy has been verified269 for (int i = 0; i < (int) cs.commands.size(); ++i) {270 probs_id.emplace_back(probs[cs.commands[i].tokens[0]], i);271 }272 273 // sort descending274 {275 using pair_type = decltype(probs_id)::value_type;276 std::sort(probs_id.begin(), probs_id.end(), [](const pair_type & a, const pair_type & b) {277 return a.first > b.first;278 });279 }280 int id = probs_id[0].second;281 return json{282 {"command_index", id},283 {"command_text", cs.commands[id].plaintext},284 {"timestamp", unprocessed_audio_timestamp},285 };286 }287}288 289static json register_commandset(struct whisper_context * ctx, json jparams, std::vector<struct commandset> &commandset_list) {290 // TODO: check for token collision291 struct commandset cs;292 293 std::string k_prompt = " select one from the available words: ";294 std::set<whisper_token> token_set;295 whisper_token tokens[32];296 for (std::string s : jparams) {297 std::vector<whisper_token> token_vec;298 // The existing command implementation uses a nested for loop to tokenize single characters299 // I fail to see the purpose of this when ' a' has a wholly different pronunciation than the start of ' apple'300 const int n = whisper_tokenize(ctx, (" " + s).c_str(), tokens, 32);301 if (n < 0) {302 fprintf(stderr, "%s: error: failed to tokenize command '%s'\n", __func__, s.c_str());303 return 3;304 }305 token_vec.push_back(tokens[0]);306 if (!token_set.insert(tokens[0]).second) {307 fprintf(stderr, "%s: warning: %s is a duplicate of an existing token\n", __func__, s.c_str());308 throw json{309 {"code",-31000},310 {"message", "Duplicate token in token set: " + s}311 };312 }313 if (n > 1) {// empty string if n=0? Should never occur314 fprintf(stderr, "%s: error: command is more than a single token: %s\n", __func__, s.c_str());315 }316 struct command command = {token_vec, s};317 cs.commands.push_back(command);318 k_prompt += s;319 }320 k_prompt = k_prompt.substr(0,k_prompt.length()-2) + ". Selected word:";321 cs.prompt_tokens.resize(1024);322 int n = whisper_tokenize(ctx, k_prompt.c_str(), cs.prompt_tokens.data(), 1024);323 cs.prompt_tokens.resize(n);324 // prepare response325 int index = commandset_list.size();326 commandset_list.push_back(cs);327 return json{{"index",index}};328}329 330static json seek(struct whisper_context * /*ctx*/, audio_async & /*audio*/, json /*params*/) {331 // whisper_state has the pertinent offsets, but there also seem to be a large332 // number of scratch buffers that would prevent rewinding context in a manner similar to llama333 // I'll give this a another pass once everything else is implemented,334 // but for now, it's unsupported335 throw json {336 {"code", -32601},337 {"message", "Seeking is not yet supported."}338 };339}340 341static json parse_job(const json &body, struct whisper_context * ctx, audio_async &audio, const whisper_params ¶ms, std::vector<struct commandset> &commandset_list) {342 // See: https://www.jsonrpc.org/specification343 json id = body.at("id");344 try {345 std::string version = body.at("jsonrpc");346 if (version != "2.0") {347 // unsupported version348 throw json{349 {"code", -3260},350 {"message", "invalid jsonrpc version"}351 };352 }353 std::string method = body.at("method");354 json jparams = json{{"dummy", "dummy"}};355 if (body.contains("params"))356 jparams = body.at("params");357 json res;358 // TODO: be consistent about argument order359 fprintf(stderr, "Dispatching a job\n");360 if (method == "unguided") { res = unguided_transcription(ctx, audio, jparams, params); }361 else if (method == "guided") { res = guided_transcription(ctx, audio, params, jparams, commandset_list); }362 else if (method == "seek") { res = seek(ctx, audio, jparams); }363 else if (method == "registerCommandset") { res = register_commandset(ctx, jparams, commandset_list); }364 else if (method == "echo") { res = jparams; }365 366 367 return json{368 {"jsonrpc", "2.0"},369 {"result", res},370 {"id", id}371 };372 } catch(json ex) {373 return json {374 {"jsonrpc", "2.0"},375 {"error", ex},376 {"id", id}377 };378 }379}380 381static void process_loop(struct whisper_context * ctx, audio_async &audio, const whisper_params ¶ms) {382 std::deque<json> jobqueue;383 std::vector<struct commandset> commandset_list;384 while (true) {385 // For eventual cancellation support, shouldn't block if job exists386 if (std::cin.rdbuf()->in_avail() > 22 || jobqueue.size() == 0) {387 int content_length;388 if (scanf("Content-Length: %d", &content_length) != 1) {389 fprintf(stderr, "Could not read input: %d", std::cin.peek());390 return;391 }392 // scanf leaves the new lines intact393 std::cin.ignore(2);394 if (std::cin.peek() != 13) {395 // Content-Type. jsonrpc necessitates utf8.396 std::cin.ignore(200,10);397 }398 std::cin.ignore(2);399 // A message is being sent and blocking is acceptable400 std::string content(content_length,'\0');401 std::cin.read(&content[0], content_length);402 json job = json::parse(content);403 // TODO: Some messages(cancellation) should skip queue here404 if (job.is_array()) {405 // response must also be batched. Will implement later406 // for (subjob : job.begin())407 // TODO: At the very least respond with an unsupported error.408 } else {409 jobqueue.push_back(job);410 }411 }412 assert(jobqueue.size() > 0);413 json job = jobqueue.front();414 json resp = parse_job(job, ctx, audio, params, commandset_list);415 if (resp != "unfinished") {416 jobqueue.pop_front();417 // send response418 std::string data = resp.dump(-1, ' ', false, json::error_handler_t::replace);419 fprintf(stdout, "Content-Length: %d\r\n\r\n%s\n", (int)data.length()+1, data.c_str());420 std::cout.flush();421 422 }423 }424}425 426int main(int argc, char ** argv) {427 ggml_backend_load_all();428 429 whisper_params params;430 if (whisper_params_parse(argc, argv, params) == false) {431 return 1;432 }433 434 if (whisper_lang_id(params.language.c_str()) == -1) {435 fprintf(stderr, "error: unknown language '%s'\n", params.language.c_str());436 whisper_print_usage(argc, argv, params);437 exit(0);438 }439 440 // whisper init441 struct whisper_context_params cparams = whisper_context_default_params();442 443 cparams.use_gpu = params.use_gpu;444 cparams.flash_attn = params.flash_attn;445 446 struct whisper_context * ctx = whisper_init_from_file_with_params(params.model.c_str(), cparams);447 // init audio448 449 audio_async audio(30*1000);450 if (!audio.init(params.capture_id, WHISPER_SAMPLE_RATE)) {451 fprintf(stderr, "%s: audio.init() failed!\n", __func__);452 return 1;453 }454 455 audio.resume();456 // TODO: Investigate why this is required. An extra second of startup latency is not great457 // wait for 1 second to avoid any buffered noise458 std::this_thread::sleep_for(std::chrono::milliseconds(1000));459 audio.clear();460 // TODO: consider some sort of indicator to designate loading has finished?461 // Potentially better for the client to just start with a non-blocking message (register commands)462 process_loop(ctx, audio, params);463 464 audio.pause();465 whisper_print_timings(ctx);466 whisper_free(ctx);467 468 return 0;469}470 