mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-09-29 19:11:11 +02:00
whisper : add abort_callback on lang detection (#4077)
This commit is contained in:
+20
-5
@@ -4149,12 +4149,14 @@ const char * whisper_lang_str_full(int id) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
int whisper_lang_auto_detect_with_state(
|
||||
static int whisper_lang_auto_detect_internal(
|
||||
struct whisper_context * ctx,
|
||||
struct whisper_state * state,
|
||||
int offset_ms,
|
||||
int n_threads,
|
||||
float * lang_probs) {
|
||||
float * lang_probs,
|
||||
ggml_abort_callback abort_callback,
|
||||
void * abort_callback_data) {
|
||||
const int seek = offset_ms/10;
|
||||
|
||||
if (seek < 0) {
|
||||
@@ -4168,14 +4170,18 @@ int whisper_lang_auto_detect_with_state(
|
||||
}
|
||||
|
||||
// run the encoder
|
||||
if (whisper_encode_with_state(ctx, state, seek, n_threads) != 0) {
|
||||
if (!whisper_encode_internal(*ctx, *state, seek, n_threads, abort_callback, abort_callback_data)) {
|
||||
WHISPER_LOG_ERROR("%s: failed to encode\n", __func__);
|
||||
return -6;
|
||||
}
|
||||
|
||||
const std::vector<whisper_token> prompt = { whisper_token_sot(ctx) };
|
||||
|
||||
if (whisper_decode_with_state(ctx, state, prompt.data(), prompt.size(), 0, n_threads) != 0) {
|
||||
whisper_batch_prep_legacy(state->batch, prompt.data(), prompt.size(), 0, 0);
|
||||
|
||||
whisper_kv_cache_seq_rm(state->kv_self, 0, 0, -1);
|
||||
|
||||
if (!whisper_decode_internal(*ctx, *state, state->batch, n_threads, false, abort_callback, abort_callback_data)) {
|
||||
WHISPER_LOG_ERROR("%s: failed to decode\n", __func__);
|
||||
return -7;
|
||||
}
|
||||
@@ -4224,6 +4230,15 @@ int whisper_lang_auto_detect_with_state(
|
||||
return logits_id[0].second;
|
||||
}
|
||||
|
||||
int whisper_lang_auto_detect_with_state(
|
||||
struct whisper_context * ctx,
|
||||
struct whisper_state * state,
|
||||
int offset_ms,
|
||||
int n_threads,
|
||||
float * lang_probs) {
|
||||
return whisper_lang_auto_detect_internal(ctx, state, offset_ms, n_threads, lang_probs, nullptr, nullptr);
|
||||
}
|
||||
|
||||
int whisper_lang_auto_detect(
|
||||
struct whisper_context * ctx,
|
||||
int offset_ms,
|
||||
@@ -6977,7 +6992,7 @@ int whisper_full_with_state(
|
||||
}
|
||||
}
|
||||
|
||||
const auto lang_id = whisper_lang_auto_detect_with_state(ctx, state, 0, params.n_threads, probs.data());
|
||||
const auto lang_id = whisper_lang_auto_detect_internal(ctx, state, 0, params.n_threads, probs.data(), params.abort_callback, params.abort_callback_user_data);
|
||||
if (lang_id < 0) {
|
||||
WHISPER_LOG_ERROR("%s: failed to auto-detect language\n", __func__);
|
||||
return -3;
|
||||
|
||||
@@ -113,6 +113,16 @@ target_compile_definitions(${ZERO_SAMPLES_TEST} PRIVATE
|
||||
add_test(NAME ${ZERO_SAMPLES_TEST} COMMAND ${ZERO_SAMPLES_TEST})
|
||||
set_tests_properties(${ZERO_SAMPLES_TEST} PROPERTIES LABELS "tiny;gh")
|
||||
|
||||
# abort_callback must be honored during language auto-detection (#3888)
|
||||
set(LANG_DETECT_ABORT_TEST test-whisper-lang-detect-abort)
|
||||
add_executable(${LANG_DETECT_ABORT_TEST} ${LANG_DETECT_ABORT_TEST}.cpp)
|
||||
target_include_directories(${LANG_DETECT_ABORT_TEST} PRIVATE ../include ../ggml/include ../examples)
|
||||
target_link_libraries(${LANG_DETECT_ABORT_TEST} PRIVATE common)
|
||||
target_compile_definitions(${LANG_DETECT_ABORT_TEST} PRIVATE
|
||||
WHISPER_MODEL_PATH="${PROJECT_SOURCE_DIR}/models/for-tests-ggml-tiny.bin")
|
||||
add_test(NAME ${LANG_DETECT_ABORT_TEST} COMMAND ${LANG_DETECT_ABORT_TEST})
|
||||
set_tests_properties(${LANG_DETECT_ABORT_TEST} PROPERTIES LABELS "tiny;gh")
|
||||
|
||||
# VAD test tests VAD in isolation
|
||||
set(VAD_TEST test-vad)
|
||||
add_executable(${VAD_TEST} ${VAD_TEST}.cpp)
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
#include "whisper.h"
|
||||
|
||||
#include <cstdio>
|
||||
#include <vector>
|
||||
|
||||
#ifdef NDEBUG
|
||||
#undef NDEBUG
|
||||
#endif
|
||||
#include <cassert>
|
||||
|
||||
static int n_encoder_begin = 0;
|
||||
|
||||
static bool encoder_begin_cb(struct whisper_context *, struct whisper_state *, void *) {
|
||||
n_encoder_begin++;
|
||||
return true; // don't block anything, just count
|
||||
}
|
||||
|
||||
static bool abort_cb(void *) {
|
||||
return true; // abort immediately
|
||||
}
|
||||
|
||||
int main() {
|
||||
ggml_backend_load_all();
|
||||
|
||||
struct whisper_context_params cparams = whisper_context_default_params();
|
||||
cparams.use_gpu = false;
|
||||
|
||||
struct whisper_context * ctx = whisper_init_from_file_with_params(WHISPER_MODEL_PATH, cparams);
|
||||
assert(ctx != nullptr);
|
||||
|
||||
std::vector<float> pcmf32(2*WHISPER_SAMPLE_RATE, 0.0f); // 2 s of silence
|
||||
|
||||
struct whisper_full_params params = whisper_full_default_params(WHISPER_SAMPLING_GREEDY);
|
||||
params.language = "auto";
|
||||
params.print_progress = false;
|
||||
params.print_realtime = false;
|
||||
|
||||
params.encoder_begin_callback = encoder_begin_cb;
|
||||
params.abort_callback = abort_cb;
|
||||
|
||||
const int rc = whisper_full(ctx, params, pcmf32.data(), pcmf32.size());
|
||||
|
||||
assert(rc != 0); // the call must fail because we aborted
|
||||
assert(n_encoder_begin == 1); // aborted during auto-detect: main-loop encoder never started
|
||||
|
||||
whisper_free(ctx);
|
||||
|
||||
printf("test-whisper-lang-detect-abort: OK\n");
|
||||
return 0;
|
||||
}
|
||||
Reference in New Issue
Block a user