diff --git a/src/ggml_graph.cpp b/src/ggml_graph.cpp index 7bef2a6..cbcfcb8 100644 --- a/src/ggml_graph.cpp +++ b/src/ggml_graph.cpp @@ -81,6 +81,11 @@ Backend& global_backend() { return global_backend_locked(); } +int backend_thread_count() { + std::lock_guard lock(g_backend_mutex); + return g_backend_threads; +} + void shutdown_backend() { // Free the process-global backend explicitly. Required for GPU backends: the // backend (and its gallocr's device buffer) must be released while the CUDA/ @@ -107,16 +112,26 @@ bool run_graph(size_t /*mem_bytes*/, int n_threads, std::lock_guard lock(g_backend_mutex); Backend& be = global_backend_locked(); // When no global override is set, honor the caller's per-call n_threads (the - // historical behavior, used by the unit tests). A positive global override - // already pinned the backend's thread count in global_backend(). + // historical behavior, used by the unit tests) for THIS call only. The + // backend is process-global, so a count left behind would slow every later + // graph: a Silero VAD pass (1 thread per tiny chunk) used to leave the ASR + // decode single-threaded. The previous count is restored after the compute. + // A positive global override already pinned the backend's thread count in + // global_backend_locked() and wins over the per-call value. const int g = g_num_threads.load(std::memory_order_relaxed); - if (g <= 0 && n_threads > 0 && n_threads != g_backend_threads) { + const int prev_threads = g_backend_threads; + const bool scoped = g <= 0 && n_threads > 0 && n_threads != prev_threads; + if (scoped) { be.set_n_threads(n_threads); g_backend_threads = n_threads; } // Backend::compute builds the graph in a no_alloc context, allocates via the // persistent gallocr, pushes inputs AFTER alloc, computes, and reads back. + struct Restore { + Backend& be; bool on; int prev; + ~Restore() { if (on) { be.set_n_threads(prev); g_backend_threads = prev; } } + } restore{be, scoped, prev_threads}; return be.compute(build, out); } diff --git a/src/ggml_graph.hpp b/src/ggml_graph.hpp index c3e9fef..a0cb33d 100644 --- a/src/ggml_graph.hpp +++ b/src/ggml_graph.hpp @@ -36,6 +36,11 @@ int num_threads(); // current override (0 == unset) // if set, else the built-in default (8). int effective_threads(); +// The thread count the process-global backend is using right now (0 before it +// is created). A per-call `n_threads` in run_graph applies to that call only, so +// this does not change across a call. For tests. +int backend_thread_count(); + // Gallocr buffer size (bytes) reserved for the most recent single-backend (CPU) // run_graph compute. Used by tests to assert that banded attention memory scales // O(T*window), not O(T^2). Reflects the high-water mark of the persistent diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index d015ced..1f13adc 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -7,6 +7,7 @@ function(pk_add_test name) endfunction() pk_add_test(test_smoke) +pk_add_test(test_run_graph_threads) pk_add_test(test_nbest_json) pk_add_test(test_tdt_beam_core) pk_add_test(test_backend_device) diff --git a/tests/test_run_graph_threads.cpp b/tests/test_run_graph_threads.cpp new file mode 100644 index 0000000..1e06914 --- /dev/null +++ b/tests/test_run_graph_threads.cpp @@ -0,0 +1,61 @@ +// A per-call n_threads in run_graph must not change the thread count of the +// process-global backend afterwards. Silero VAD asks for one thread per chunk; +// before this was scoped, a VAD pass left the ASR decode that followed on one +// thread. The Silero part runs only when PARAKEET_TEST_SILERO_GGUF is set. +#include +#include +#include +#include + +#include "backend.hpp" +#include "ggml.h" +#include "ggml_graph.hpp" +#include "silero_vad.hpp" + +static int failures = 0; +#define CHECK(c) do { if (!(c)) { std::fprintf(stderr, "FAIL %s:%d %s\n", __FILE__, __LINE__, #c); ++failures; } } while (0) + +static bool tiny_graph(int n_threads) { + std::vector out; + const float x[4] = {1, 2, 3, 4}; + const bool ok = pk::run_graph(0, n_threads, [&](ggml_context* ctx) -> ggml_tensor* { + const int64_t ne[1] = {4}; + ggml_tensor* in = pk::graph_input_tensor(ctx, GGML_TYPE_F32, 1, ne, x, sizeof(x)); + return ggml_scale(ctx, in, 2.0f); + }, out); + return ok && out.size() == 4 && out[3] == 8.0f; +} + +int main() { + pk::set_num_threads(0); // no --threads + CHECK(tiny_graph(0)); + const int before = pk::backend_thread_count(); + CHECK(before == pk::effective_threads()); + + CHECK(tiny_graph(1)); + CHECK(pk::backend_thread_count() == before); + CHECK(tiny_graph(before == 3 ? 2 : 3)); + CHECK(pk::backend_thread_count() == before); + + if (const char* gguf = std::getenv("PARAKEET_TEST_SILERO_GGUF")) { + std::string err; + auto vad = pk::SileroVad::load(gguf, &err); + CHECK(vad != nullptr); + if (vad) { + std::vector pcm(16000, 0.0f); + CHECK(!vad->probabilities(pcm.data(), pcm.size(), 16000).empty()); + CHECK(pk::backend_thread_count() == before); + } + } + + // A global override still wins over the per-call count. + pk::set_num_threads(2); + CHECK(tiny_graph(1)); + CHECK(pk::backend_thread_count() == 2); + pk::set_num_threads(0); + + pk::shutdown_backend(); + if (failures) return 1; + std::printf("ok\n"); + return 0; +}