diff --git a/src/whisper.cpp b/src/whisper.cpp index 771c1149c..dd8bd9378 100644 --- a/src/whisper.cpp +++ b/src/whisper.cpp @@ -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 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; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 57bfb51b6..ab4acfac7 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -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) diff --git a/tests/test-whisper-lang-detect-abort.cpp b/tests/test-whisper-lang-detect-abort.cpp new file mode 100644 index 000000000..ba362062a --- /dev/null +++ b/tests/test-whisper-lang-detect-abort.cpp @@ -0,0 +1,50 @@ +#include "whisper.h" + +#include +#include + +#ifdef NDEBUG +#undef NDEBUG +#endif +#include + +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 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; +}