Xenobd/whisper.cpp
0
1// Voice assistant example2//3// Speak short text commands to the microphone.4// This program will detect your voice command and convert them to text.5//6// ref: https://github.com/ggml-org/whisper.cpp/issues/1717//8 9#include "common-sdl.h"10#include "common.h"11#include "whisper.h"12#include "grammar-parser.h"13 14#include <algorithm>15#include <chrono>16#include <cstdio>17#include <fstream>18#include <map>19#include <sstream>20#include <string>21#include <thread>22#include <vector>23 24// command-line parameters25struct whisper_params {26 int32_t n_threads = std::min(4, (int32_t) std::thread::hardware_concurrency());27 int32_t prompt_ms = 5000;28 int32_t command_ms = 8000;29 int32_t capture_id = -1;30 int32_t max_tokens = 32;31 int32_t audio_ctx = 0;32 33 float vad_thold = 0.6f;34 float freq_thold = 100.0f;35 36 float grammar_penalty = 100.0f;37 38 grammar_parser::parse_state grammar_parsed;39 40 bool translate = false;41 bool print_special = false;42 bool print_energy = false;43 bool no_timestamps = true;44 bool use_gpu = true;45 bool flash_attn = false;46 47 std::string language = "en";48 std::string model = "models/ggml-base.en.bin";49 std::string fname_out;50 std::string commands;51 std::string prompt;52 std::string context;53 std::string grammar;54 55 // A regular expression that matches tokens to suppress56 std::string suppress_regex;57};58 59void whisper_print_usage(int argc, char ** argv, const whisper_params & params);60 61static bool whisper_params_parse(int argc, char ** argv, whisper_params & params) {62 for (int i = 1; i < argc; i++) {63 std::string arg = argv[i];64 65 if (arg == "-h" || arg == "--help") {66 whisper_print_usage(argc, argv, params);67 exit(0);68 }69 else if (arg == "-t" || arg == "--threads") { params.n_threads = std::stoi(argv[++i]); }70 else if (arg == "-pms" || arg == "--prompt-ms") { params.prompt_ms = std::stoi(argv[++i]); }71 else if (arg == "-cms" || arg == "--command-ms") { params.command_ms = std::stoi(argv[++i]); }72 else if (arg == "-c" || arg == "--capture") { params.capture_id = std::stoi(argv[++i]); }73 else if (arg == "-mt" || arg == "--max-tokens") { params.max_tokens = std::stoi(argv[++i]); }74 else if (arg == "-ac" || arg == "--audio-ctx") { params.audio_ctx = std::stoi(argv[++i]); }75 else if (arg == "-vth" || arg == "--vad-thold") { params.vad_thold = std::stof(argv[++i]); }76 else if (arg == "-fth" || arg == "--freq-thold") { params.freq_thold = std::stof(argv[++i]); }77 else if (arg == "-tr" || arg == "--translate") { params.translate = true; }78 else if (arg == "-ps" || arg == "--print-special") { params.print_special = true; }79 else if (arg == "-pe" || arg == "--print-energy") { params.print_energy = true; }80 else if (arg == "-ng" || arg == "--no-gpu") { params.use_gpu = false; }81 else if (arg == "-fa" || arg == "--flash-attn") { params.flash_attn = true; }82 else if (arg == "-l" || arg == "--language") { params.language = argv[++i]; }83 else if (arg == "-m" || arg == "--model") { params.model = argv[++i]; }84 else if (arg == "-f" || arg == "--file") { params.fname_out = argv[++i]; }85 else if (arg == "-cmd" || arg == "--commands") { params.commands = argv[++i]; }86 else if (arg == "-p" || arg == "--prompt") { params.prompt = argv[++i]; }87 else if (arg == "-ctx" || arg == "--context") { params.context = argv[++i]; }88 else if ( arg == "--grammar") { params.grammar = argv[++i]; }89 else if ( arg == "--grammar-penalty") { params.grammar_penalty = std::stof(argv[++i]); }90 else if ( arg == "--suppress-regex") { params.suppress_regex = argv[++i]; }91 else {92 fprintf(stderr, "error: unknown argument: %s\n", arg.c_str());93 whisper_print_usage(argc, argv, params);94 exit(0);95 }96 }97 98 return true;99}100 101void whisper_print_usage(int /*argc*/, char ** argv, const whisper_params & params) {102 fprintf(stderr, "\n");103 fprintf(stderr, "usage: %s [options]\n", argv[0]);104 fprintf(stderr, "\n");105 fprintf(stderr, "options:\n");106 fprintf(stderr, " -h, --help [default] show this help message and exit\n");107 fprintf(stderr, " -t N, --threads N [%-7d] number of threads to use during computation\n", params.n_threads);108 fprintf(stderr, " -pms N, --prompt-ms N [%-7d] prompt duration in milliseconds\n", params.prompt_ms);109 fprintf(stderr, " -cms N, --command-ms N [%-7d] command duration in milliseconds\n", params.command_ms);110 fprintf(stderr, " -c ID, --capture ID [%-7d] capture device ID\n", params.capture_id);111 fprintf(stderr, " -mt N, --max-tokens N [%-7d] maximum number of tokens per audio chunk\n", params.max_tokens);112 fprintf(stderr, " -ac N, --audio-ctx N [%-7d] audio context size (0 - all)\n", params.audio_ctx);113 fprintf(stderr, " -vth N, --vad-thold N [%-7.2f] voice activity detection threshold\n", params.vad_thold);114 fprintf(stderr, " -fth N, --freq-thold N [%-7.2f] high-pass frequency cutoff\n", params.freq_thold);115 fprintf(stderr, " -tr, --translate [%-7s] translate from source language to english\n", params.translate ? "true" : "false");116 fprintf(stderr, " -ps, --print-special [%-7s] print special tokens\n", params.print_special ? "true" : "false");117 fprintf(stderr, " -pe, --print-energy [%-7s] print sound energy (for debugging)\n", params.print_energy ? "true" : "false");118 fprintf(stderr, " -ng, --no-gpu [%-7s] disable GPU\n", params.use_gpu ? "false" : "true");119 fprintf(stderr, " -fa, --flash-attn [%-7s] flash attention\n", params.flash_attn ? "true" : "false");120 fprintf(stderr, " -l LANG, --language LANG [%-7s] spoken language\n", params.language.c_str());121 fprintf(stderr, " -m FNAME, --model FNAME [%-7s] model path\n", params.model.c_str());122 fprintf(stderr, " -f FNAME, --file FNAME [%-7s] text output file name\n", params.fname_out.c_str());123 fprintf(stderr, " -cmd FNAME, --commands FNAME [%-7s] text file with allowed commands\n", params.commands.c_str());124 fprintf(stderr, " -p, --prompt [%-7s] the required activation prompt\n", params.prompt.c_str());125 fprintf(stderr, " -ctx, --context [%-7s] sample text to help the transcription\n", params.context.c_str());126 fprintf(stderr, " --grammar GRAMMAR [%-7s] GBNF grammar to guide decoding\n", params.grammar.c_str());127 fprintf(stderr, " --grammar-penalty N [%-7.1f] scales down logits of nongrammar tokens\n", params.grammar_penalty);128 fprintf(stderr, " --suppress-regex REGEX [%-7s] regular expression matching tokens to suppress\n", params.suppress_regex.c_str());129 fprintf(stderr, "\n");130}131 132static std::string transcribe(133 whisper_context * ctx,134 const whisper_params & params,135 const std::vector<float> & pcmf32,136 const std::string & grammar_rule,137 float & logprob_min,138 float & logprob_sum,139 int & n_tokens,140 int64_t & t_ms) {141 const auto t_start = std::chrono::high_resolution_clock::now();142 143 logprob_min = 0.0f;144 logprob_sum = 0.0f;145 n_tokens = 0;146 t_ms = 0;147 148 //whisper_full_params wparams = whisper_full_default_params(WHISPER_SAMPLING_GREEDY);149 whisper_full_params wparams = whisper_full_default_params(WHISPER_SAMPLING_BEAM_SEARCH);150 151 wparams.print_progress = false;152 wparams.print_special = params.print_special;153 wparams.print_realtime = false;154 wparams.print_timestamps = !params.no_timestamps;155 wparams.translate = params.translate;156 wparams.no_context = true;157 wparams.no_timestamps = params.no_timestamps;158 wparams.single_segment = true;159 wparams.max_tokens = params.max_tokens;160 wparams.language = params.language.c_str();161 wparams.n_threads = params.n_threads;162 163 wparams.audio_ctx = params.audio_ctx;164 165 wparams.temperature = 0.4f;166 wparams.temperature_inc = 1.0f;167 wparams.greedy.best_of = 5;168 169 wparams.beam_search.beam_size = 5;170 171 wparams.initial_prompt = params.context.data();172 173 wparams.suppress_regex = params.suppress_regex.c_str();174 175 const auto & grammar_parsed = params.grammar_parsed;176 auto grammar_rules = grammar_parsed.c_rules();177 178 if (!params.grammar_parsed.rules.empty() && !grammar_rule.empty()) {179 if (grammar_parsed.symbol_ids.find(grammar_rule) == grammar_parsed.symbol_ids.end()) {180 fprintf(stderr, "%s: warning: grammar rule '%s' not found - skipping grammar sampling\n", __func__, grammar_rule.c_str());181 } else {182 wparams.grammar_rules = grammar_rules.data();183 wparams.n_grammar_rules = grammar_rules.size();184 wparams.i_start_rule = grammar_parsed.symbol_ids.at(grammar_rule);185 wparams.grammar_penalty = params.grammar_penalty;186 }187 }188 189 if (whisper_full(ctx, wparams, pcmf32.data(), pcmf32.size()) != 0) {190 return "";191 }192 193 std::string result;194 195 const int n_segments = whisper_full_n_segments(ctx);196 for (int i = 0; i < n_segments; ++i) {197 const char * text = whisper_full_get_segment_text(ctx, i);198 199 result += text;200 201 const int n = whisper_full_n_tokens(ctx, i);202 for (int j = 0; j < n; ++j) {203 const auto token = whisper_full_get_token_data(ctx, i, j);204 205 if(token.plog > 0.0f) exit(0);206 logprob_min = std::min(logprob_min, token.plog);207 logprob_sum += token.plog;208 ++n_tokens;209 }210 }211 212 const auto t_end = std::chrono::high_resolution_clock::now();213 t_ms = std::chrono::duration_cast<std::chrono::milliseconds>(t_end - t_start).count();214 215 return result;216}217 218static std::vector<std::string> read_allowed_commands(const std::string & fname) {219 std::vector<std::string> allowed_commands;220 221 std::ifstream ifs(fname);222 if (!ifs.is_open()) {223 return allowed_commands;224 }225 226 std::string line;227 while (std::getline(ifs, line)) {228 line = ::trim(line);229 if (line.empty()) {230 continue;231 }232 233 std::transform(line.begin(), line.end(),line.begin(), ::tolower);234 allowed_commands.push_back(std::move(line));235 }236 237 return allowed_commands;238}239 240static std::vector<std::string> get_words(const std::string &txt) {241 std::vector<std::string> words;242 243 std::istringstream iss(txt);244 std::string word;245 while (iss >> word) {246 words.push_back(word);247 }248 249 return words;250}251 252// command-list mode253// guide the transcription to match the most likely command from a provided list254static int process_command_list(struct whisper_context * ctx, audio_async &audio, const whisper_params ¶ms, std::ofstream &fout) {255 fprintf(stderr, "\n");256 fprintf(stderr, "%s: guided mode\n", __func__);257 258 std::vector<std::string> allowed_commands = read_allowed_commands(params.commands);259 260 if (allowed_commands.empty()) {261 fprintf(stderr, "%s: error: failed to read allowed commands from '%s'\n", __func__, params.commands.c_str());262 return 2;263 }264 265 int max_len = 0;266 267 std::vector<std::vector<whisper_token>> allowed_tokens;268 269 for (const auto & cmd : allowed_commands) {270 whisper_token tokens[1024];271 allowed_tokens.emplace_back();272 273 for (int l = 0; l < (int) cmd.size(); ++l) {274 // NOTE: very important to add the whitespace !275 // the reason is that the first decoded token starts with a whitespace too!276 std::string ss = std::string(" ") + cmd.substr(0, l + 1);277 278 const int n = whisper_tokenize(ctx, ss.c_str(), tokens, 1024);279 if (n < 0) {280 fprintf(stderr, "%s: error: failed to tokenize command '%s'\n", __func__, cmd.c_str());281 return 3;282 }283 284 if (n == 1) {285 allowed_tokens.back().push_back(tokens[0]);286 }287 }288 289 max_len = std::max(max_len, (int) cmd.size());290 }291 292 fprintf(stderr, "%s: allowed commands [ tokens ]:\n", __func__);293 fprintf(stderr, "\n");294 for (int i = 0; i < (int) allowed_commands.size(); ++i) {295 fprintf(stderr, " - \033[1m%-*s\033[0m = [", max_len, allowed_commands[i].c_str());296 for (const auto & token : allowed_tokens[i]) {297 fprintf(stderr, " %5d", token);298 }299 fprintf(stderr, " ]\n");300 }301 302 std::string k_prompt = "select one from the available words: ";303 for (int i = 0; i < (int) allowed_commands.size(); ++i) {304 if (i > 0) {305 k_prompt += ", ";306 }307 k_prompt += allowed_commands[i];308 }309 k_prompt += ". selected word: ";310 311 // tokenize prompt312 std::vector<whisper_token> k_tokens;313 {314 k_tokens.resize(1024);315 const int n = whisper_tokenize(ctx, k_prompt.c_str(), k_tokens.data(), 1024);316 if (n < 0) {317 fprintf(stderr, "%s: error: failed to tokenize prompt '%s'\n", __func__, k_prompt.c_str());318 return 4;319 }320 k_tokens.resize(n);321 }322 323 fprintf(stderr, "\n");324 fprintf(stderr, "%s: prompt: '%s'\n", __func__, k_prompt.c_str());325 fprintf(stderr, "%s: tokens: [", __func__);326 for (const auto & token : k_tokens) {327 fprintf(stderr, " %d", token);328 }329 fprintf(stderr, " ]\n");330 331 fprintf(stderr, "\n");332 fprintf(stderr, "%s: listening for a command ...\n", __func__);333 fprintf(stderr, "\n");334 335 bool is_running = true;336 337 std::vector<float> pcmf32_cur;338 std::vector<float> pcmf32_prompt;339 340 // main loop341 while (is_running) {342 // handle Ctrl + C343 is_running = sdl_poll_events();344 345 // delay346 std::this_thread::sleep_for(std::chrono::milliseconds(100));347 348 audio.get(2000, pcmf32_cur);349 350 if (::vad_simple(pcmf32_cur, WHISPER_SAMPLE_RATE, 1000, params.vad_thold, params.freq_thold, params.print_energy)) {351 fprintf(stdout, "%s: Speech detected! Processing ...\n", __func__);352 353 const auto t_start = std::chrono::high_resolution_clock::now();354 355 whisper_full_params wparams = whisper_full_default_params(WHISPER_SAMPLING_GREEDY);356 357 wparams.print_progress = false;358 wparams.print_special = params.print_special;359 wparams.print_realtime = false;360 wparams.print_timestamps = !params.no_timestamps;361 wparams.translate = params.translate;362 wparams.no_context = true;363 wparams.single_segment = true;364 wparams.max_tokens = 1;365 wparams.language = params.language.c_str();366 wparams.n_threads = params.n_threads;367 368 wparams.audio_ctx = params.audio_ctx;369 370 wparams.prompt_tokens = k_tokens.data();371 wparams.prompt_n_tokens = k_tokens.size();372 373 // run the transformer and a single decoding pass374 if (whisper_full(ctx, wparams, pcmf32_cur.data(), pcmf32_cur.size()) != 0) {375 fprintf(stderr, "%s: ERROR: whisper_full() failed\n", __func__);376 break;377 }378 379 // estimate command probability380 // NOTE: not optimal381 {382 const auto * logits = whisper_get_logits(ctx);383 384 std::vector<float> probs(whisper_n_vocab(ctx), 0.0f);385 386 // compute probs from logits via softmax387 {388 float max = -1e9;389 for (int i = 0; i < (int) probs.size(); ++i) {390 max = std::max(max, logits[i]);391 }392 393 float sum = 0.0f;394 for (int i = 0; i < (int) probs.size(); ++i) {395 probs[i] = expf(logits[i] - max);396 sum += probs[i];397 }398 399 for (int i = 0; i < (int) probs.size(); ++i) {400 probs[i] /= sum;401 }402 }403 404 std::vector<std::pair<float, int>> probs_id;405 406 double psum = 0.0;407 for (int i = 0; i < (int) allowed_commands.size(); ++i) {408 probs_id.emplace_back(probs[allowed_tokens[i][0]], i);409 for (int j = 1; j < (int) allowed_tokens[i].size(); ++j) {410 probs_id.back().first += probs[allowed_tokens[i][j]];411 }412 probs_id.back().first /= allowed_tokens[i].size();413 psum += probs_id.back().first;414 }415 416 // normalize417 for (auto & p : probs_id) {418 p.first /= psum;419 }420 421 // sort descending422 {423 using pair_type = decltype(probs_id)::value_type;424 std::sort(probs_id.begin(), probs_id.end(), [](const pair_type & a, const pair_type & b) {425 return a.first > b.first;426 });427 }428 429 // print the commands and the respective probabilities430 {431 fprintf(stdout, "\n");432 for (const auto & cmd : probs_id) {433 fprintf(stdout, "%s: %s%-*s%s = %f | ", __func__, "\033[1m", max_len, allowed_commands[cmd.second].c_str(), "\033[0m", cmd.first);434 for (int token : allowed_tokens[cmd.second]) {435 fprintf(stdout, "'%4s' %f ", whisper_token_to_str(ctx, token), probs[token]);436 }437 fprintf(stdout, "\n");438 }439 }440 441 // best command442 {443 const auto t_end = std::chrono::high_resolution_clock::now();444 445 const float prob = probs_id[0].first;446 const int index = probs_id[0].second;447 const char * best_command = allowed_commands[index].c_str();448 449 fprintf(stdout, "\n");450 fprintf(stdout, "%s: detected command: %s%s%s | p = %f | t = %d ms\n", __func__,451 "\033[1m", best_command, "\033[0m", prob,452 (int) std::chrono::duration_cast<std::chrono::milliseconds>(t_end - t_start).count());453 fprintf(stdout, "\n");454 if (fout.is_open()) {455 fout << best_command << std::endl;456 }457 }458 }459 460 audio.clear();461 }462 }463 464 return 0;465}466 467// always-prompt mode468// transcribe the voice into text after valid prompt469static int always_prompt_transcription(struct whisper_context * ctx, audio_async & audio, const whisper_params & params, std::ofstream & fout) {470 bool is_running = true;471 bool ask_prompt = true;472 473 float logprob_min = 0.0f;474 float logprob_sum = 0.0f;475 int n_tokens = 0;476 477 std::vector<float> pcmf32_cur;478 479 const std::string k_prompt = params.prompt;480 481 const int k_prompt_length = get_words(k_prompt).size();482 483 fprintf(stderr, "\n");484 fprintf(stderr, "%s: always-prompt mode\n", __func__);485 486 // main loop487 while (is_running) {488 // handle Ctrl + C489 is_running = sdl_poll_events();490 491 // delay492 std::this_thread::sleep_for(std::chrono::milliseconds(100));493 494 if (ask_prompt) {495 fprintf(stdout, "\n");496 fprintf(stdout, "%s: The prompt is: '%s%s%s'\n", __func__, "\033[1m", k_prompt.c_str(), "\033[0m");497 fprintf(stdout, "\n");498 499 ask_prompt = false;500 }501 502 {503 audio.get(2000, pcmf32_cur);504 505 if (::vad_simple(pcmf32_cur, WHISPER_SAMPLE_RATE, 1000, params.vad_thold, params.freq_thold, params.print_energy)) {506 fprintf(stdout, "%s: Speech detected! Processing ...\n", __func__);507 508 int64_t t_ms = 0;509 510 // detect the commands511 audio.get(params.command_ms, pcmf32_cur);512 513 const auto txt = ::trim(::transcribe(ctx, params, pcmf32_cur, "", logprob_min, logprob_sum, n_tokens, t_ms));514 515 const auto words = get_words(txt);516 517 std::string prompt;518 std::string command;519 520 for (int i = 0; i < (int) words.size(); ++i) {521 if (i < k_prompt_length) {522 prompt += words[i] + " ";523 } else {524 command += words[i] + " ";525 }526 }527 528 const float sim = similarity(prompt, k_prompt);529 530 //debug531 //fprintf(stdout, "command size: %i\n", command_length);532 533 if ((sim > 0.7f) && (command.size() > 0)) {534 fprintf(stdout, "%s: Command '%s%s%s', (t = %d ms)\n", __func__, "\033[1m", command.c_str(), "\033[0m", (int) t_ms);535 if (fout.is_open()) {536 fout << command << std::endl;537 }538 }539 540 fprintf(stdout, "\n");541 542 audio.clear();543 }544 }545 }546 547 return 0;548}549 550// general-purpose mode551// freely transcribe the voice into text552static int process_general_transcription(struct whisper_context * ctx, audio_async & audio, const whisper_params & params, std::ofstream & fout) {553 bool is_running = true;554 bool have_prompt = false;555 bool ask_prompt = true;556 557 float logprob_min0 = 0.0f;558 float logprob_min = 0.0f;559 560 float logprob_sum0 = 0.0f;561 float logprob_sum = 0.0f;562 563 int n_tokens0 = 0;564 int n_tokens = 0;565 566 std::vector<float> pcmf32_cur;567 std::vector<float> pcmf32_prompt;568 569 std::string k_prompt = "Ok Whisper, start listening for commands.";570 if (!params.prompt.empty()) {571 k_prompt = params.prompt;572 }573 574 fprintf(stderr, "\n");575 fprintf(stderr, "%s: general-purpose mode\n", __func__);576 577 // main loop578 while (is_running) {579 // handle Ctrl + C580 is_running = sdl_poll_events();581 582 // delay583 std::this_thread::sleep_for(std::chrono::milliseconds(100));584 585 if (ask_prompt) {586 fprintf(stdout, "\n");587 fprintf(stdout, "%s: Say the following phrase: '%s%s%s'\n", __func__, "\033[1m", k_prompt.c_str(), "\033[0m");588 fprintf(stdout, "\n");589 590 ask_prompt = false;591 }592 593 {594 audio.get(2000, pcmf32_cur);595 596 if (::vad_simple(pcmf32_cur, WHISPER_SAMPLE_RATE, 1000, params.vad_thold, params.freq_thold, params.print_energy)) {597 fprintf(stdout, "%s: Speech detected! Processing ...\n", __func__);598 599 int64_t t_ms = 0;600 601 if (!have_prompt) {602 // wait for activation phrase603 audio.get(params.prompt_ms, pcmf32_cur);604 605 const auto txt = ::trim(::transcribe(ctx, params, pcmf32_cur, "prompt", logprob_min0, logprob_sum0, n_tokens0, t_ms));606 607 const float p = 100.0f * std::exp(logprob_min0);608 609 fprintf(stdout, "%s: Heard '%s%s%s', (t = %d ms, p = %.2f%%)\n", __func__, "\033[1m", txt.c_str(), "\033[0m", (int) t_ms, p);610 611 const float sim = similarity(txt, k_prompt);612 613 if (txt.length() < 0.8*k_prompt.length() || txt.length() > 1.2*k_prompt.length() || sim < 0.8f) {614 fprintf(stdout, "%s: WARNING: prompt not recognized, try again\n", __func__);615 ask_prompt = true;616 } else {617 fprintf(stdout, "\n");618 fprintf(stdout, "%s: The prompt has been recognized!\n", __func__);619 fprintf(stdout, "%s: Waiting for voice commands ...\n", __func__);620 fprintf(stdout, "\n");621 622 // save the audio for the prompt623 pcmf32_prompt = pcmf32_cur;624 have_prompt = true;625 }626 } else {627 // we have heard the activation phrase, now detect the commands628 audio.get(params.command_ms, pcmf32_cur);629 630 //printf("len prompt: %.4f\n", pcmf32_prompt.size() / (float) WHISPER_SAMPLE_RATE);631 //printf("len command: %.4f\n", pcmf32_cur.size() / (float) WHISPER_SAMPLE_RATE);632 633 // prepend 3 second of silence634 pcmf32_cur.insert(pcmf32_cur.begin(), 3.0f*WHISPER_SAMPLE_RATE, 0.0f);635 636 // prepend the prompt audio637 pcmf32_cur.insert(pcmf32_cur.begin(), pcmf32_prompt.begin(), pcmf32_prompt.end());638 639 const auto txt = ::trim(::transcribe(ctx, params, pcmf32_cur, "root", logprob_min, logprob_sum, n_tokens, t_ms));640 641 //const float p = 100.0f * std::exp((logprob - logprob0) / (n_tokens - n_tokens0));642 const float p = 100.0f * std::exp(logprob_min);643 644 //fprintf(stdout, "%s: heard '%s'\n", __func__, txt.c_str());645 646 // find the prompt in the text647 float best_sim = 0.0f;648 size_t best_len = 0;649 for (size_t n = 0.8*k_prompt.size(); n <= 1.2*k_prompt.size(); ++n) {650 if (n >= txt.size()) {651 break;652 }653 654 const auto prompt = txt.substr(0, n);655 656 const float sim = similarity(prompt, k_prompt);657 658 //fprintf(stderr, "%s: prompt = '%s', sim = %f\n", __func__, prompt.c_str(), sim);659 660 if (sim > best_sim) {661 best_sim = sim;662 best_len = n;663 }664 }665 666 fprintf(stdout, "%s: DEBUG: txt = '%s', prob = %.2f%%\n", __func__, txt.c_str(), p);667 if (best_len == 0) {668 fprintf(stdout, "%s: WARNING: command not recognized, try again\n", __func__);669 } else {670 // cut the prompt from the decoded text671 const std::string command = ::trim(txt.substr(best_len));672 fprintf(stdout, "%s: Command '%s%s%s', (t = %d ms)\n", __func__, "\033[1m", command.c_str(), "\033[0m", (int) t_ms);673 if (fout.is_open()) {674 fout << command << std::endl;675 }676 }677 678 fprintf(stdout, "\n");679 }680 681 audio.clear();682 }683 }684 }685 686 return 0;687}688 689int main(int argc, char ** argv) {690 ggml_backend_load_all();691 692 whisper_params params;693 694 if (whisper_params_parse(argc, argv, params) == false) {695 return 1;696 }697 698 if (whisper_lang_id(params.language.c_str()) == -1) {699 fprintf(stderr, "error: unknown language '%s'\n", params.language.c_str());700 whisper_print_usage(argc, argv, params);701 exit(0);702 }703 704 // whisper init705 706 struct whisper_context_params cparams = whisper_context_default_params();707 708 cparams.use_gpu = params.use_gpu;709 cparams.flash_attn = params.flash_attn;710 711 struct whisper_context * ctx = whisper_init_from_file_with_params(params.model.c_str(), cparams);712 if (ctx == nullptr) {713 fprintf(stderr, "error: failed to initialize whisper context\n");714 return 2;715 }716 717 // print some info about the processing718 {719 fprintf(stderr, "\n");720 if (!whisper_is_multilingual(ctx)) {721 if (params.language != "en" || params.translate) {722 params.language = "en";723 params.translate = false;724 fprintf(stderr, "%s: WARNING: model is not multilingual, ignoring language and translation options\n", __func__);725 }726 }727 fprintf(stderr, "%s: processing, %d threads, lang = %s, task = %s, timestamps = %d ...\n",728 __func__,729 params.n_threads,730 params.language.c_str(),731 params.translate ? "translate" : "transcribe",732 params.no_timestamps ? 0 : 1);733 734 fprintf(stderr, "\n");735 }736 737 // init audio738 739 audio_async audio(30*1000);740 if (!audio.init(params.capture_id, WHISPER_SAMPLE_RATE)) {741 fprintf(stderr, "%s: audio.init() failed!\n", __func__);742 return 1;743 }744 745 audio.resume();746 747 // wait for 1 second to avoid any buffered noise748 std::this_thread::sleep_for(std::chrono::milliseconds(1000));749 audio.clear();750 751 int ret_val = 0;752 753 if (!params.grammar.empty()) {754 auto & grammar = params.grammar_parsed;755 if (is_file_exist(params.grammar.c_str())) {756 // read grammar from file757 std::ifstream ifs(params.grammar.c_str());758 const std::string txt = std::string((std::istreambuf_iterator<char>(ifs)), std::istreambuf_iterator<char>());759 grammar = grammar_parser::parse(txt.c_str());760 } else {761 // read grammar from string762 grammar = grammar_parser::parse(params.grammar.c_str());763 }764 765 // will be empty (default) if there are parse errors766 if (grammar.rules.empty()) {767 ret_val = 1;768 } else {769 fprintf(stderr, "%s: grammar:\n", __func__);770 grammar_parser::print_grammar(stderr, grammar);771 fprintf(stderr, "\n");772 }773 }774 775 std::ofstream fout;776 if (params.fname_out.length() > 0) {777 fout.open(params.fname_out);778 if (!fout.is_open()) {779 fprintf(stderr, "%s: failed to open output file '%s'!\n", __func__, params.fname_out.c_str());780 return 1;781 }782 }783 784 if (ret_val == 0) {785 if (!params.commands.empty()) {786 ret_val = process_command_list(ctx, audio, params, fout);787 } else if (!params.prompt.empty() && params.grammar_parsed.rules.empty()) {788 ret_val = always_prompt_transcription(ctx, audio, params, fout);789 } else {790 ret_val = process_general_transcription(ctx, audio, params, fout);791 }792 }793 794 audio.pause();795 796 whisper_print_timings(ctx);797 whisper_free(ctx);798 799 return ret_val;800}801 