Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 18 additions & 3 deletions src/ggml_graph.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,11 @@ Backend& global_backend() {
return global_backend_locked();
}

int backend_thread_count() {
std::lock_guard<std::recursive_mutex> 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/
Expand All @@ -107,16 +112,26 @@ bool run_graph(size_t /*mem_bytes*/, int n_threads,
std::lock_guard<std::recursive_mutex> 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);
}

Expand Down
5 changes: 5 additions & 0 deletions src/ggml_graph.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
61 changes: 61 additions & 0 deletions tests/test_run_graph_threads.cpp
Original file line number Diff line number Diff line change
@@ -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 <cstdio>
#include <cstdlib>
#include <string>
#include <vector>

#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<float> 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<float> 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;
}
Loading