From 57f69e627a6a9fc8fc32d1506dcaaec4570cc05c Mon Sep 17 00:00:00 2001 From: Jhin Lee Date: Tue, 22 Sep 2026 18:02:31 -0400 Subject: [PATCH 1/4] fix: let a cancel stop the Qwen3-TTS audio decode inside a step llama_dart_tts_step checked for a cancel only on entry, so a cancel that landed in a step running the code2wav decode waited out the whole decode: 1.3-1.4 s on CPU and 0.36 s on Metal on an M4 Max. Export llama_dart_tts_eval_callback for mtmd_context_params.cb_eval. Inside a step it splits mtmd graph computes into chunks of about 2.5e9 multiply-adds, ends chunks only after a MUL_MAT, and stops at the next chunk end once the task is cancelled. The step re-checks the cancel after frame generation and audio output and returns CANCELLED. Export llama_dart_tts_set_cancel_flag so a caller can cancel through a byte it owns without sharing the task pointer with the cancelling thread. Refs https://github.com/leehack/llamadart-native/issues/86 --- .github/workflows/validate_wrapper.yml | 1 + CMakeLists.txt | 12 +- README.md | 16 +- src/llama_dart_tts_eval_internal.h | 64 ++++++ src/llama_dart_wrapper.cpp | 101 +++++++++- src/llama_dart_wrapper.h | 25 ++- tests/tts_api_test.c | 8 + tests/tts_eval_test.cpp | 263 +++++++++++++++++++++++++ tools/tts_smoke.cpp | 212 +++++++++++++++++++- tools/validate_exports.py | 2 + 10 files changed, 683 insertions(+), 21 deletions(-) create mode 100644 src/llama_dart_tts_eval_internal.h create mode 100644 tests/tts_eval_test.cpp diff --git a/.github/workflows/validate_wrapper.yml b/.github/workflows/validate_wrapper.yml index 20d4633..f63edae 100644 --- a/.github/workflows/validate_wrapper.yml +++ b/.github/workflows/validate_wrapper.yml @@ -261,6 +261,7 @@ jobs: cmake --build build/wrapper-contract --target llamadart_speculative_api_test llamadart_tts_api_test + llamadart_tts_eval_test llamadart_mtmd_compat_test llamadart_tts_smoke - name: Run wrapper contract tests diff --git a/CMakeLists.txt b/CMakeLists.txt index 6e4a644..c65f412 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -224,6 +224,15 @@ if (LLAMADART_BUILD_TESTS) target_include_directories(llamadart_tts_api_test PRIVATE src) add_test(NAME llamadart_tts_api_test COMMAND llamadart_tts_api_test) + add_executable(llamadart_tts_eval_test tests/tts_eval_test.cpp) + target_compile_features(llamadart_tts_eval_test PRIVATE cxx_std_17) + target_include_directories(llamadart_tts_eval_test PRIVATE src) + target_link_libraries(llamadart_tts_eval_test PRIVATE ggml) + # Beside the ggml backend modules that ggml_backend_load_all() finds. + set_target_properties(llamadart_tts_eval_test PROPERTIES + RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/bin") + add_test(NAME llamadart_tts_eval_test COMMAND llamadart_tts_eval_test) + add_executable(llamadart_mtmd_compat_test tests/mtmd_compat_test.cpp) target_compile_features(llamadart_mtmd_compat_test PRIVATE cxx_std_17) target_include_directories(llamadart_mtmd_compat_test PRIVATE src) @@ -233,7 +242,8 @@ if (LLAMADART_BUILD_TESTS) add_test(NAME llamadart_mtmd_compat_test COMMAND llamadart_mtmd_compat_test) # Contract assertions must execute against optimized Release libraries too. - foreach(test_target llamadart_speculative_api_test llamadart_tts_api_test) + foreach(test_target llamadart_speculative_api_test llamadart_tts_api_test + llamadart_tts_eval_test) if (MSVC) target_compile_options(${test_target} PRIVATE /UNDEBUG) else() diff --git a/README.md b/README.md index df97438..a45dc7f 100644 --- a/README.md +++ b/README.md @@ -249,8 +249,13 @@ give the TTS task exclusive access to both contexts until the task reaches a terminal state or is reset. The upstream API does not currently expose completed PCM incrementally, so the -wrapper's step API is cancellable between prompt batches and generation frames -but is not a real-time audio stream. Public Dart support should remain +wrapper's step API is not a real-time audio stream. A step sees a cancel when it +starts and when its frame generation or audio output returns. Qwen3-TTS decodes +audio in 72-frame windows, inside the step that fills a window and the step that +ends speech. Install `llama_dart_tts_eval_callback` as the mtmd context's +`cb_eval` to stop that decode at its next chunk boundary, and attach a +caller-owned cancel byte with `llama_dart_tts_set_cancel_flag` to cancel from a +thread that must not touch the task. Public Dart support should remain experimental and capability-gated until artifact and platform validation is complete. @@ -264,7 +269,7 @@ cmake -S . -B build/tts-smoke -G Ninja \ -DLLAMADART_BUILD_TTS_SMOKE=ON cmake --build build/tts-smoke --target \ llamadart_speculative_api_test llamadart_tts_api_test \ - llamadart_mtmd_compat_test llamadart_tts_smoke + llamadart_tts_eval_test llamadart_mtmd_compat_test llamadart_tts_smoke ctest --test-dir build/tts-smoke --output-on-failure build/tts-smoke/llamadart_tts_smoke \ /path/to/Qwen3-TTS-model.gguf \ @@ -276,6 +281,11 @@ build/tts-smoke/llamadart_tts_smoke \ The smoke checks capability metadata, cancellation/reset, two consecutive syntheses, 24 kHz mono PCM metadata, finite/non-silent output, and WAV writing. +It also checks the eval callback: uncancelled PCM stays byte-identical, frame +steps stay whole, no chunk of the final decode exceeds a quarter of it, and a +cancel issued a quarter of the way into a decode step, through +`llama_dart_tts_cancel` or a cancel flag, makes that step return `CANCELLED` +within a third of the decode's duration. Omit `--gpu` for a CPU-only run. An optional speaker-reference audio path may appear before the final `--gpu` flag. diff --git a/src/llama_dart_tts_eval_internal.h b/src/llama_dart_tts_eval_internal.h new file mode 100644 index 0000000..335f2ba --- /dev/null +++ b/src/llama_dart_tts_eval_internal.h @@ -0,0 +1,64 @@ +#pragma once + +#include "ggml.h" + +// Scheduler work between chunk boundaries inside one TTS step: multiply-adds +// for MUL_MAT, elements for other computed nodes. +static constexpr double llama_dart_tts_eval_budget = 2.5e9; + +struct llama_dart_tts_eval_chunker { + double budget = llama_dart_tts_eval_budget; + double work = 0.0; + bool boundary_due = false; + bool past_first_boundary = false; +}; + +static inline double llama_dart_tts_eval_node_work(const ggml_tensor *node) { + switch (node->op) { + case GGML_OP_NONE: + case GGML_OP_VIEW: + case GGML_OP_RESHAPE: + case GGML_OP_PERMUTE: + case GGML_OP_TRANSPOSE: + return 0.0; + case GGML_OP_MUL_MAT: + return static_cast(ggml_nelements(node)) * + static_cast(node->src[0]->ne[0]); + default: + return static_cast(ggml_nelements(node)); + } +} + +// Answers one ggml_backend_sched_eval_callback query for a step in progress. +// +// Budget boundaries fall only after a MUL_MAT node. A boundary splits any +// fusion that spans it, which can change the output; the CPU and Metal +// backends fuse nothing that continues past a MUL_MAT. +// +// A break stops only the current scheduler split; later splits still run. +// The code predictor runs first in a step and feeds its sampled indices to +// get_rows, so no break happens before the step's first budget boundary. +static inline bool llama_dart_tts_eval_answer( + llama_dart_tts_eval_chunker *chunker, const ggml_tensor *node, bool ask, + bool cancelled) { + if (!ask) { + return !cancelled; + } + if ((node->flags & GGML_TENSOR_FLAG_COMPUTE) == 0) { + return false; + } + if (cancelled && chunker->past_first_boundary) { + return true; + } + chunker->work += llama_dart_tts_eval_node_work(node); + if (chunker->work >= chunker->budget) { + chunker->boundary_due = true; + } + if (!chunker->boundary_due || node->op != GGML_OP_MUL_MAT) { + return false; + } + chunker->work = 0.0; + chunker->boundary_due = false; + chunker->past_first_boundary = true; + return true; +} diff --git a/src/llama_dart_wrapper.cpp b/src/llama_dart_wrapper.cpp index 8bfbb23..8a13cff 100644 --- a/src/llama_dart_wrapper.cpp +++ b/src/llama_dart_wrapper.cpp @@ -2,6 +2,7 @@ #include "llama_dart_mtp_internal.h" #include "llama_dart_mtmd_compat.h" #include "llama_dart_speculative_compat.h" +#include "llama_dart_tts_eval_internal.h" #include "common.h" #include "llama-ext.h" @@ -77,6 +78,7 @@ struct llama_dart_tts { llama_sampler *sampler = nullptr; mtmd_bitmap *speaker = nullptr; std::atomic cancel_requested{false}; + const int8_t *cancel_flag = nullptr; llama_dart_tts_state state = LLAMA_DART_TTS_STATE_IDLE; llama_seq_id sequence_id = 0; bool owns_sequence = false; @@ -94,6 +96,42 @@ struct llama_dart_tts { static void llama_dart_tts_release_task_resources(llama_dart_tts *tts); +static bool llama_dart_tts_flag_raised(const int8_t *flag) { + using atomic_flag_byte = std::atomic; + static_assert(sizeof(atomic_flag_byte) == sizeof(int8_t) && + alignof(atomic_flag_byte) == alignof(int8_t) && + atomic_flag_byte::is_always_lock_free, + "a caller-owned cancel byte must be readable atomically"); + return reinterpret_cast(flag)->load( + std::memory_order_relaxed) != 0; +} + +static bool llama_dart_tts_cancelled(const llama_dart_tts *tts) { + return tts->cancel_requested.load(std::memory_order_acquire) || + (tts->cancel_flag != nullptr && + llama_dart_tts_flag_raised(tts->cancel_flag)); +} + +struct llama_dart_tts_eval_scope; + +static thread_local llama_dart_tts_eval_scope *llama_dart_tts_active_eval = + nullptr; + +struct llama_dart_tts_eval_scope { + const llama_dart_tts *tts; + llama_dart_tts_eval_chunker chunker; + llama_dart_tts_eval_scope *previous; + + explicit llama_dart_tts_eval_scope(const llama_dart_tts *task) + : tts(task), previous(llama_dart_tts_active_eval) { + llama_dart_tts_active_eval = this; + } + ~llama_dart_tts_eval_scope() { llama_dart_tts_active_eval = previous; } + llama_dart_tts_eval_scope(const llama_dart_tts_eval_scope &) = delete; + llama_dart_tts_eval_scope & + operator=(const llama_dart_tts_eval_scope &) = delete; +}; + static llama_dart_tts_status llama_dart_tts_fail( llama_dart_tts *tts, llama_dart_tts_status status, const char *message) { if (tts != nullptr) { @@ -131,6 +169,15 @@ static void llama_dart_tts_release_task_resources(llama_dart_tts *tts) { llama_memory_seq_rm(llama_get_memory(tts->llama), tts->sequence_id, 0, -1); tts->owns_sequence = false; } + tts->cancel_flag = nullptr; +} + +static llama_dart_tts_status llama_dart_tts_mark_cancelled( + llama_dart_tts *tts) { + tts->state = LLAMA_DART_TTS_STATE_CANCELLED; + tts->error = "TTS task cancelled"; + llama_dart_tts_release_task_resources(tts); + return LLAMA_DART_TTS_STATUS_CANCELLED; } static llama_dart_tts_model_type llama_dart_tts_model_type_from_upstream( @@ -199,8 +246,12 @@ static llama_dart_tts_status llama_dart_tts_finish_output( const char *data = nullptr; size_t data_len = 0; int64_t sample_count = 0; - if (mtmd_helper_gen_audio_get_output(tts->generator, &sample_rate, &data, - &data_len, &sample_count) != 0) { + const int32_t output_status = mtmd_helper_gen_audio_get_output( + tts->generator, &sample_rate, &data, &data_len, &sample_count); + if (llama_dart_tts_cancelled(tts)) { + return llama_dart_tts_mark_cancelled(tts); + } + if (output_status != 0) { return llama_dart_tts_fail(tts, LLAMA_DART_TTS_STATUS_UPSTREAM_ERROR, "audio output conversion failed"); } @@ -883,13 +934,12 @@ LLAMADART_API enum llama_dart_tts_status llama_dart_tts_step( } const bool active = tts->state == LLAMA_DART_TTS_STATE_PROCESSING_PROMPT || tts->state == LLAMA_DART_TTS_STATE_GENERATING; - if (active && tts->cancel_requested.load(std::memory_order_acquire)) { - tts->state = LLAMA_DART_TTS_STATE_CANCELLED; - tts->error = "TTS task cancelled"; - llama_dart_tts_release_task_resources(tts); + if (active && llama_dart_tts_cancelled(tts)) { + const auto status = llama_dart_tts_mark_cancelled(tts); llama_dart_tts_write_progress(tts, out_progress); - return LLAMA_DART_TTS_STATUS_CANCELLED; + return status; } + llama_dart_tts_eval_scope eval_scope(tts); if (tts->state == LLAMA_DART_TTS_STATE_PROCESSING_PROMPT) { const int32_t remaining = mtmd_helper_gen_audio_step_prompt( tts->generator, tts->prompt_batch_size); @@ -940,10 +990,17 @@ LLAMADART_API enum llama_dart_tts_status llama_dart_tts_step( const float *state = llama_get_embeddings_ith(tts->llama, -1); const float *next_state = nullptr; bool stop = false; - if (state == nullptr || + const bool generated = + state != nullptr && llama_dart_tts_step_gen(&mtmd_helper_gen_audio_step_gen, tts->generator, sampled, state, &next_state, - &stop) != 0) { + &stop) == 0; + if (llama_dart_tts_cancelled(tts)) { + const auto status = llama_dart_tts_mark_cancelled(tts); + llama_dart_tts_write_progress(tts, out_progress); + return status; + } + if (!generated) { const auto status = llama_dart_tts_fail( tts, LLAMA_DART_TTS_STATUS_UPSTREAM_ERROR, "TTS generation step failed"); @@ -973,6 +1030,32 @@ LLAMADART_API void llama_dart_tts_cancel(struct llama_dart_tts *tts) { } } +LLAMADART_API enum llama_dart_tts_status +llama_dart_tts_set_cancel_flag(struct llama_dart_tts *tts, + const int8_t *flag) { + if (tts == nullptr || flag == nullptr) { + return LLAMA_DART_TTS_STATUS_INVALID_ARGUMENT; + } + if (tts->state != LLAMA_DART_TTS_STATE_PROCESSING_PROMPT && + tts->state != LLAMA_DART_TTS_STATE_GENERATING) { + return llama_dart_tts_error(tts, LLAMA_DART_TTS_STATUS_INVALID_STATE, + "no TTS task is active"); + } + tts->cancel_flag = flag; + return LLAMA_DART_TTS_STATUS_OK; +} + +LLAMADART_API bool llama_dart_tts_eval_callback(struct ggml_tensor *tensor, + bool ask, void *user_data) { + (void)user_data; + llama_dart_tts_eval_scope *scope = llama_dart_tts_active_eval; + if (scope == nullptr) { + return !ask; + } + return llama_dart_tts_eval_answer(&scope->chunker, tensor, ask, + llama_dart_tts_cancelled(scope->tts)); +} + LLAMADART_API enum llama_dart_tts_status llama_dart_tts_reset(struct llama_dart_tts *tts) { if (tts == nullptr) { diff --git a/src/llama_dart_wrapper.h b/src/llama_dart_wrapper.h index ce3c072..fbf5d0d 100644 --- a/src/llama_dart_wrapper.h +++ b/src/llama_dart_wrapper.h @@ -181,10 +181,31 @@ LLAMADART_API enum llama_dart_tts_status llama_dart_tts_step( struct llama_dart_tts * tts, struct llama_dart_tts_progress * out_progress); -// May be called from another thread. Cancellation is observed between native -// prompt and generation steps. +// May be called from another thread. A step returns CANCELLED when the cancel +// arrives before the step starts or before its frame generation or audio +// output returns. With llama_dart_tts_eval_callback installed, the audio +// decode stops at its next chunk boundary instead of running to the end. LLAMADART_API void llama_dart_tts_cancel(struct llama_dart_tts * tts); +// Attaches a caller-owned cancel byte to the active task. Any thread may set +// *flag to nonzero; steps then behave as after llama_dart_tts_cancel. The task +// reads *flag only inside llama_dart_tts_step and drops it when the task ends, +// restarts, is reset, or is freed; *flag must stay readable until then. +// Returns INVALID_STATE when no task is active. +LLAMADART_API enum llama_dart_tts_status llama_dart_tts_set_cancel_flag( + struct llama_dart_tts * tts, + const int8_t * flag); + +// ggml_backend_sched_eval_callback for mtmd_context_params.cb_eval; +// cb_eval_user_data is unused. While llama_dart_tts_step runs on the calling +// thread, it splits mtmd graph computes into chunks and, once the task is +// cancelled, ends each split's compute at its next chunk boundary. Outside a +// step it requests no tensors. +LLAMADART_API bool llama_dart_tts_eval_callback( + struct ggml_tensor * tensor, + bool ask, + void * user_data); + // Clears task state and the task sequence from the caller-owned llama context. LLAMADART_API enum llama_dart_tts_status llama_dart_tts_reset( struct llama_dart_tts * tts); diff --git a/tests/tts_api_test.c b/tests/tts_api_test.c index 23d906c..a82ba88 100644 --- a/tests/tts_api_test.c +++ b/tests/tts_api_test.c @@ -51,6 +51,14 @@ int main(void) { assert(strcmp(llama_dart_tts_last_error(NULL), "invalid TTS handle") == 0); llama_dart_tts_cancel(NULL); + const int8_t flag = 1; + assert(llama_dart_tts_set_cancel_flag(NULL, &flag) == LLAMA_DART_TTS_STATUS_INVALID_ARGUMENT); + + struct ggml_tensor node; + memset(&node, 0, sizeof(node)); + assert(!llama_dart_tts_eval_callback(&node, true, NULL)); + assert(llama_dart_tts_eval_callback(&node, false, NULL)); + llama_dart_tts_free(NULL); return 0; } diff --git a/tests/tts_eval_test.cpp b/tests/tts_eval_test.cpp new file mode 100644 index 0000000..446f275 --- /dev/null +++ b/tests/tts_eval_test.cpp @@ -0,0 +1,263 @@ +#ifdef NDEBUG +#error "Wrapper contract tests require active assertions in every configuration" +#endif + +#include "llama_dart_tts_eval_internal.h" + +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml.h" + +#include +#include +#include + +namespace { + +constexpr int64_t kDim = 64; +constexpr int64_t kCols = 8; +constexpr int kLayers = 6; +constexpr double kMatMulWork = static_cast(kDim * kCols * kDim); +constexpr double kRowWork = static_cast(kDim * kCols); + +ggml_tensor *computed(ggml_tensor *tensor) { + tensor->flags |= GGML_TENSOR_FLAG_COMPUTE; + return tensor; +} + +void test_node_work(ggml_context *ctx) { + ggml_tensor *a = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kDim, 32); + ggml_tensor *b = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kDim, kCols); + ggml_tensor *product = ggml_mul_mat(ctx, a, b); + assert(llama_dart_tts_eval_node_work(product) == 32.0 * kCols * kDim); + assert(llama_dart_tts_eval_node_work(ggml_add(ctx, product, product)) == + 32.0 * kCols); + assert(llama_dart_tts_eval_node_work(ggml_view_1d(ctx, product, 8, 0)) == + 0.0); + assert(llama_dart_tts_eval_node_work(ggml_reshape_1d(ctx, product, 256)) == + 0.0); + assert(llama_dart_tts_eval_node_work(ggml_transpose(ctx, product)) == 0.0); + assert(llama_dart_tts_eval_node_work( + ggml_permute(ctx, product, 1, 0, 2, 3)) == 0.0); + assert(llama_dart_tts_eval_node_work(a) == 0.0); +} + +void test_boundaries_follow_budget_and_mul_mat(ggml_context *ctx) { + ggml_tensor *w = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kDim, kDim); + ggml_tensor *x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kDim, kCols); + ggml_tensor *product = computed(ggml_mul_mat(ctx, w, x)); + ggml_tensor *sum = computed(ggml_add(ctx, product, product)); + ggml_tensor *idle_product = ggml_mul_mat(ctx, w, x); + + llama_dart_tts_eval_chunker chunker; + chunker.budget = kMatMulWork + kRowWork; + assert(!llama_dart_tts_eval_answer(&chunker, product, true, false)); + assert(!llama_dart_tts_eval_answer(&chunker, idle_product, true, false)); + assert(chunker.work == kMatMulWork); + assert(!llama_dart_tts_eval_answer(&chunker, sum, true, false)); + assert(chunker.boundary_due); + assert(!llama_dart_tts_eval_answer(&chunker, sum, true, false)); + assert(!llama_dart_tts_eval_answer(&chunker, idle_product, true, false)); + assert(llama_dart_tts_eval_answer(&chunker, product, true, false)); + assert(chunker.work == 0.0 && !chunker.boundary_due); + assert(chunker.past_first_boundary); + assert(!llama_dart_tts_eval_answer(&chunker, product, true, false)); + assert(llama_dart_tts_eval_answer(&chunker, product, false, false)); + assert(!llama_dart_tts_eval_answer(&chunker, product, false, true)); +} + +void test_cancel_waits_for_first_boundary(ggml_context *ctx) { + ggml_tensor *w = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kDim, kDim); + ggml_tensor *x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kDim, kCols); + ggml_tensor *product = computed(ggml_mul_mat(ctx, w, x)); + ggml_tensor *sum = computed(ggml_add(ctx, product, product)); + ggml_tensor *idle_sum = ggml_add(ctx, product, product); + + llama_dart_tts_eval_chunker chunker; + chunker.budget = 2 * kMatMulWork; + assert(!llama_dart_tts_eval_answer(&chunker, product, true, true)); + assert(!llama_dart_tts_eval_answer(&chunker, sum, true, true)); + assert(llama_dart_tts_eval_answer(&chunker, product, true, true)); + assert(llama_dart_tts_eval_answer(&chunker, sum, true, true)); + assert(!llama_dart_tts_eval_answer(&chunker, idle_sum, true, true)); + assert(!llama_dart_tts_eval_answer(&chunker, sum, true, false)); +} + +struct probe { + llama_dart_tts_eval_chunker chunker; + bool cancelled = false; + bool cancel_at_first_boundary = false; + bool stopped = false; + int asks = 0; + int asks_after_stop = 0; + int boundaries = 0; + int non_mul_mat_boundaries = 0; +}; + +bool probe_callback(ggml_tensor *node, bool ask, void *user_data) { + auto *p = static_cast(user_data); + if (ask && p->stopped) { + ++p->asks_after_stop; + } + const bool answer = + llama_dart_tts_eval_answer(&p->chunker, node, ask, p->cancelled); + if (!ask) { + p->stopped = p->stopped || !answer; + return answer; + } + ++p->asks; + if (answer) { + ++p->boundaries; + p->non_mul_mat_boundaries += node->op == GGML_OP_MUL_MAT ? 0 : 1; + p->cancelled = p->cancelled || p->cancel_at_first_boundary; + } + return answer; +} + +struct scheduled_graph { + ggml_backend_t backend = nullptr; + ggml_context *weights_ctx = nullptr; + ggml_backend_buffer_t weights = nullptr; + ggml_context *graph_ctx = nullptr; + ggml_cgraph *graph = nullptr; + ggml_tensor *output = nullptr; + ggml_backend_sched_t sched = nullptr; + + scheduled_graph() { + ggml_backend_load_all(); + backend = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + assert(backend != nullptr); + + ggml_init_params weight_params = { + ggml_tensor_overhead() * (2 * kLayers + 1), nullptr, true}; + weights_ctx = ggml_init(weight_params); + ggml_tensor *input = + ggml_new_tensor_2d(weights_ctx, GGML_TYPE_F32, kDim, kCols); + std::vector matrices; + std::vector gains; + for (int layer = 0; layer < kLayers; ++layer) { + matrices.push_back( + ggml_new_tensor_2d(weights_ctx, GGML_TYPE_F32, kDim, kDim)); + gains.push_back(ggml_new_tensor_1d(weights_ctx, GGML_TYPE_F32, kDim)); + } + weights = ggml_backend_alloc_ctx_tensors(weights_ctx, backend); + assert(weights != nullptr); + fill(input, 1); + for (int layer = 0; layer < kLayers; ++layer) { + fill(matrices[layer], 7 + layer); + fill(gains[layer], 101 + layer); + } + + ggml_init_params graph_params = { + ggml_tensor_overhead() * GGML_DEFAULT_GRAPH_SIZE + + ggml_graph_overhead(), + nullptr, true}; + graph_ctx = ggml_init(graph_params); + ggml_tensor *hidden = input; + for (int layer = 0; layer < kLayers; ++layer) { + hidden = ggml_mul_mat(graph_ctx, matrices[layer], hidden); + hidden = ggml_rms_norm(graph_ctx, hidden, 1e-6f); + hidden = ggml_mul(graph_ctx, hidden, gains[layer]); + } + output = hidden; + ggml_set_output(output); + graph = ggml_new_graph(graph_ctx); + ggml_build_forward_expand(graph, output); + assert(ggml_graph_n_nodes(graph) == 3 * kLayers); + + sched = ggml_backend_sched_new(&backend, nullptr, 1, + GGML_DEFAULT_GRAPH_SIZE, false, false); + assert(sched != nullptr); + } + + ~scheduled_graph() { + ggml_backend_sched_free(sched); + ggml_free(graph_ctx); + ggml_backend_buffer_free(weights); + ggml_free(weights_ctx); + ggml_backend_free(backend); + } + + static void fill(ggml_tensor *tensor, int seed) { + std::vector values(static_cast(ggml_nelements(tensor))); + uint32_t state = static_cast(seed) * 2654435761u; + for (float &value : values) { + state = state * 1664525u + 1013904223u; + value = static_cast(state >> 8) / 16777216.0f - 0.5f; + } + ggml_backend_tensor_set(tensor, values.data(), 0, ggml_nbytes(tensor)); + } + + ggml_status compute(probe *p) { + ggml_backend_sched_reset(sched); + ggml_backend_sched_set_eval_callback(sched, p ? probe_callback : nullptr, + p); + return ggml_backend_sched_graph_compute(sched, graph); + } + + std::vector read_output() const { + std::vector values(static_cast(ggml_nelements(output))); + ggml_backend_tensor_get(output, values.data(), 0, ggml_nbytes(output)); + return values; + } +}; + +// The budget first overflows on the first RMS_NORM, so the first boundary +// moves to the second MUL_MAT; every later MUL_MAT then ends a chunk. +constexpr double kSchedulerBudget = kMatMulWork + 0.5 * kRowWork; + +void test_chunked_compute_matches_unchunked(scheduled_graph &g) { + assert(g.compute(nullptr) == GGML_STATUS_SUCCESS); + const std::vector expected = g.read_output(); + + probe p; + p.chunker.budget = kSchedulerBudget; + assert(g.compute(&p) == GGML_STATUS_SUCCESS); + assert(p.boundaries == kLayers - 1); + assert(p.non_mul_mat_boundaries == 0); + assert(!p.stopped); + const std::vector chunked = g.read_output(); + assert(std::memcmp(expected.data(), chunked.data(), + expected.size() * sizeof(float)) == 0); +} + +void test_cancel_stops_at_chunk_boundary(scheduled_graph &g) { + probe p; + p.chunker.budget = kSchedulerBudget; + p.cancel_at_first_boundary = true; + assert(g.compute(&p) == GGML_STATUS_SUCCESS); + assert(p.stopped); + assert(p.boundaries == 1); + assert(p.asks == 4); + assert(p.asks_after_stop == 0); +} + +void test_early_cancel_runs_to_first_boundary(scheduled_graph &g) { + probe p; + p.chunker.budget = kSchedulerBudget; + p.cancelled = true; + assert(g.compute(&p) == GGML_STATUS_SUCCESS); + assert(p.stopped); + assert(p.boundaries == 1); + assert(p.asks == 4); + assert(p.asks_after_stop == 0); +} + +} // namespace + +int main() { + ggml_log_set([](ggml_log_level, const char *, void *) {}, nullptr); + + ggml_init_params params = {ggml_tensor_overhead() * 64, nullptr, true}; + ggml_context *ctx = ggml_init(params); + test_node_work(ctx); + test_boundaries_follow_budget_and_mul_mat(ctx); + test_cancel_waits_for_first_boundary(ctx); + ggml_free(ctx); + + scheduled_graph graph; + test_chunked_compute_matches_unchunked(graph); + test_cancel_stops_at_chunk_boundary(graph); + test_early_cancel_runs_to_first_boundary(graph); + return 0; +} diff --git a/tools/tts_smoke.cpp b/tools/tts_smoke.cpp index 12c1476..3ff8842 100644 --- a/tools/tts_smoke.cpp +++ b/tools/tts_smoke.cpp @@ -4,6 +4,8 @@ #include "mtmd.h" #include +#include +#include #include #include #include @@ -13,8 +15,36 @@ #include #include #include +#include #include +using steady_clock = std::chrono::steady_clock; + +static const char *const long_text = + "Hello from Llama Dart. This sentence is here only to make the speech long " + "enough that the decoder buffers a full window of codec frames before the " + "end, so the flush runs in the middle of synthesis."; + +static std::vector chunk_ends; + +static bool recording_eval_callback(struct ggml_tensor *tensor, bool ask, void *user_data) { + const bool answer = llama_dart_tts_eval_callback(tensor, ask, user_data); + if (!ask) { + chunk_ends.push_back(steady_clock::now()); + } + return answer; +} + +struct step_trace { + double ms; + int chunk_boundaries; + double longest_chunk_ms; +}; + +static double elapsed_ms(steady_clock::time_point from, steady_clock::time_point to) { + return std::chrono::duration(to - from).count(); +} + static void write_u16(std::ofstream &out, uint16_t value) { const char bytes[] = {static_cast(value), static_cast(value >> 8)}; out.write(bytes, sizeof(bytes)); @@ -80,14 +110,29 @@ static bool synthesize(llama_dart_tts *tts, const llama_dart_tts_info &info, std::vector *pcm, llama_dart_tts_progress *progress, - double *rms) { + double *rms, + std::vector *trace = nullptr) { if (llama_dart_tts_start(tts, request) != LLAMA_DART_TTS_STATUS_OK) { std::fprintf(stderr, "failed to start TTS: %s\n", llama_dart_tts_last_error(tts)); return false; } progress->struct_size = sizeof(*progress); do { + chunk_ends.clear(); + const auto started = steady_clock::now(); const llama_dart_tts_status status = llama_dart_tts_step(tts, progress); + const auto returned = steady_clock::now(); + if (trace != nullptr) { + double longest = 0.0; + auto previous = started; + chunk_ends.push_back(returned); + for (const auto &end : chunk_ends) { + longest = std::max(longest, elapsed_ms(previous, end)); + previous = end; + } + trace->push_back({elapsed_ms(started, returned), + static_cast(chunk_ends.size()) - 1, longest}); + } if (status != LLAMA_DART_TTS_STATUS_OK) { std::fprintf(stderr, "TTS failed: %s\n", llama_dart_tts_last_error(tts)); return false; @@ -137,6 +182,99 @@ static bool synthesize(llama_dart_tts *tts, return true; } +static bool cancel_during_long_step(llama_dart_tts *tts, + llama_dart_tts_request *request, + std::atomic *flag, + double cancel_after_ms, + double return_within_ms, + double *cancel_to_return_ms, + int *frames_before_cancel) { + static_assert(sizeof(std::atomic) == sizeof(int8_t), + "the cancel flag must be a plain byte"); + if (llama_dart_tts_start(tts, request) != LLAMA_DART_TTS_STATUS_OK) { + std::fprintf(stderr, "failed to start in-step cancel probe: %s\n", + llama_dart_tts_last_error(tts)); + return false; + } + if (flag != nullptr && + llama_dart_tts_set_cancel_flag(tts, reinterpret_cast(flag)) != + LLAMA_DART_TTS_STATUS_OK) { + std::fprintf(stderr, "failed to attach cancel flag\n"); + return false; + } + std::atomic running_step{-1}; + std::atomic running_since{0}; + std::atomic cancelled_step{-1}; + std::atomic cancelled_at{0}; + std::atomic done{false}; + const auto origin = steady_clock::now(); + auto since_origin = [&origin] { + return std::chrono::duration_cast( + steady_clock::now() - origin) + .count(); + }; + std::thread canceller([&] { + while (!done.load()) { + const int step = running_step.load(); + if (step >= 0 && since_origin() - running_since.load() >= cancel_after_ms * 1000.0) { + cancelled_at.store(since_origin()); + cancelled_step.store(step); + if (flag != nullptr) { + flag->store(1); + } else { + llama_dart_tts_cancel(tts); + } + return; + } + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + }); + llama_dart_tts_progress progress{}; + llama_dart_tts_status status = LLAMA_DART_TTS_STATUS_OK; + int step = 0; + int64_t returned_at = 0; + for (;; ++step) { + progress.struct_size = sizeof(progress); + running_since.store(since_origin()); + running_step.store(step); + status = llama_dart_tts_step(tts, &progress); + running_step.store(-1); + returned_at = since_origin(); + if (status != LLAMA_DART_TTS_STATUS_OK || + progress.state == LLAMA_DART_TTS_STATE_COMPLETED) { + break; + } + } + done.store(true); + canceller.join(); + const bool same_step = cancelled_step.load() == step; + *cancel_to_return_ms = (returned_at - cancelled_at.load()) / 1000.0; + *frames_before_cancel = progress.frames_generated; + if (llama_dart_tts_reset(tts) != LLAMA_DART_TTS_STATUS_OK) { + std::fprintf(stderr, "reset after in-step cancel probe failed\n"); + return false; + } + if (cancelled_step.load() < 0) { + std::fprintf(stderr, "no step ran %.1f ms, so the in-step cancel never fired\n", + cancel_after_ms); + return false; + } + if (!same_step || status != LLAMA_DART_TTS_STATUS_CANCELLED || + progress.state != LLAMA_DART_TTS_STATE_CANCELLED) { + std::fprintf(stderr, + "in-step cancel at step %d ended at step %d with status %d state %d\n", + cancelled_step.load(), step, static_cast(status), + static_cast(progress.state)); + return false; + } + if (*cancel_to_return_ms > return_within_ms) { + std::fprintf(stderr, "in-step cancel took %.1f ms to return, bound %.1f ms\n", + *cancel_to_return_ms, return_within_ms); + return false; + } + return true; +} + int main(int argc, char **argv) { const bool use_gpu = argc > 1 && std::strcmp(argv[argc - 1], "--gpu") == 0; const int value_argc = argc - (use_gpu ? 1 : 0); @@ -180,9 +318,12 @@ int main(int argc, char **argv) { } mtmd_context_params mtmd_params = mtmd_context_params_default(); mtmd_params.use_gpu = use_gpu; + std::unique_ptr plain_mtmd( + mtmd_init_from_file(argv[2], model.get(), mtmd_params), mtmd_free); + mtmd_params.cb_eval = recording_eval_callback; std::unique_ptr mtmd( mtmd_init_from_file(argv[2], model.get(), mtmd_params), mtmd_free); - if (mtmd == nullptr) { + if (plain_mtmd == nullptr || mtmd == nullptr) { std::fprintf(stderr, "failed to load mmproj\n"); return 1; } @@ -206,7 +347,9 @@ int main(int argc, char **argv) { llama_dart_tts_status status = LLAMA_DART_TTS_STATUS_OK; std::unique_ptr tts( llama_dart_tts_init(context.get(), mtmd.get(), &status), llama_dart_tts_free); - if (tts == nullptr) { + std::unique_ptr plain_tts( + llama_dart_tts_init(context.get(), plain_mtmd.get(), &status), llama_dart_tts_free); + if (tts == nullptr || plain_tts == nullptr) { std::fprintf(stderr, "failed to initialize TTS wrapper: %d\n", static_cast(status)); return 1; } @@ -224,6 +367,15 @@ int main(int argc, char **argv) { std::fprintf(stderr, "non-finite sampling validation failed\n"); return 1; } + std::atomic cancel_flag{0}; + const int8_t *cancel_flag_address = reinterpret_cast(&cancel_flag); + if (llama_dart_tts_set_cancel_flag(tts.get(), cancel_flag_address) != + LLAMA_DART_TTS_STATUS_INVALID_STATE || + llama_dart_tts_set_cancel_flag(tts.get(), nullptr) != + LLAMA_DART_TTS_STATUS_INVALID_ARGUMENT) { + std::fprintf(stderr, "cancel flag validation failed\n"); + return 1; + } if (llama_dart_tts_start(tts.get(), &request) != LLAMA_DART_TTS_STATUS_OK) { std::fprintf(stderr, "failed to start cancellation probe: %s\n", llama_dart_tts_last_error(tts.get())); @@ -239,10 +391,52 @@ int main(int argc, char **argv) { return 1; } + std::vector plain_pcm; + llama_dart_tts_progress plain_progress{}; + double plain_rms = 0.0; + if (!synthesize(plain_tts.get(), &request, info, &plain_pcm, &plain_progress, &plain_rms)) { + return 1; + } + std::vector first_pcm; llama_dart_tts_progress first_progress{}; double first_rms = 0.0; - if (!synthesize(tts.get(), &request, info, &first_pcm, &first_progress, &first_rms)) { + std::vector first_trace; + if (!synthesize(tts.get(), &request, info, &first_pcm, &first_progress, &first_rms, + &first_trace)) { + return 1; + } + if (first_pcm.size() != plain_pcm.size() || + std::memcmp(first_pcm.data(), plain_pcm.data(), first_pcm.size() * sizeof(float)) != 0) { + std::fprintf(stderr, "the eval callback changed uncancelled PCM\n"); + return 1; + } + const double decode_ms = first_trace.back().ms; + for (const step_trace &trace : first_trace) { + if (trace.ms < decode_ms / 4 && trace.chunk_boundaries != 0) { + std::fprintf(stderr, "a %.1f ms frame step was split into chunks\n", trace.ms); + return 1; + } + } + if (first_trace.back().longest_chunk_ms > decode_ms / 4) { + std::fprintf(stderr, "a %.1f ms chunk of the %.1f ms final audio decode exceeds a quarter\n", + first_trace.back().longest_chunk_ms, decode_ms); + return 1; + } + + double final_cancel_ms = 0.0; + int final_cancel_frames = 0; + if (!cancel_during_long_step(tts.get(), &request, nullptr, decode_ms / 4, decode_ms / 3, + &final_cancel_ms, &final_cancel_frames)) { + return 1; + } + llama_dart_tts_request long_request = request; + long_request.text = long_text; + long_request.text_length = std::char_traits::length(long_text); + double window_cancel_ms = 0.0; + int window_cancel_frames = 0; + if (!cancel_during_long_step(tts.get(), &long_request, &cancel_flag, decode_ms / 4, + decode_ms / 3, &window_cancel_ms, &window_cancel_frames)) { return 1; } @@ -263,10 +457,16 @@ int main(int argc, char **argv) { } std::printf("PASS backend=%s model_type=%d sample_rate=%d " "first_samples=%zu second_samples=%zu " - "first_frames=%d second_frames=%d first_rms=%.6f second_rms=%.6f\n", + "first_frames=%d second_frames=%d first_rms=%.6f second_rms=%.6f " + "decode_ms=%.1f decode_chunks=%d longest_chunk_ms=%.1f " + "final_cancel_ms=%.1f final_cancel_frames=%d " + "window_cancel_ms=%.1f window_cancel_frames=%d\n", use_gpu ? "gpu" : "cpu", static_cast(info.model_type), info.sample_rate, first_pcm.size(), second_pcm.size(), first_progress.frames_generated, - second_progress.frames_generated, first_rms, second_rms); + second_progress.frames_generated, first_rms, second_rms, decode_ms, + first_trace.back().chunk_boundaries + 1, first_trace.back().longest_chunk_ms, + final_cancel_ms, final_cancel_frames, + window_cancel_ms, window_cancel_frames); return 0; } diff --git a/tools/validate_exports.py b/tools/validate_exports.py index eb31e68..5201a58 100644 --- a/tools/validate_exports.py +++ b/tools/validate_exports.py @@ -19,6 +19,8 @@ "llama_dart_tts_start", "llama_dart_tts_step", "llama_dart_tts_cancel", + "llama_dart_tts_set_cancel_flag", + "llama_dart_tts_eval_callback", "llama_dart_tts_reset", "llama_dart_tts_get_output_info", "llama_dart_tts_read_pcm", From 29ec218614a3cb0cfded054d18bf1bdcefdb6596 Mon Sep 17 00:00:00 2001 From: Jhin Lee Date: Tue, 22 Sep 2026 18:15:01 -0400 Subject: [PATCH 2/4] test: fire the smoke's in-step cancel only after the first frame The prompt step of the long text runs about 216 ms on CPU and the first frame step about 82 ms cold on Metal, close to a quarter of the decode. Waiting for a frame keeps the cancel on the decode step. --- tools/tts_smoke.cpp | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tools/tts_smoke.cpp b/tools/tts_smoke.cpp index 3ff8842..d2920db 100644 --- a/tools/tts_smoke.cpp +++ b/tools/tts_smoke.cpp @@ -203,6 +203,7 @@ static bool cancel_during_long_step(llama_dart_tts *tts, return false; } std::atomic running_step{-1}; + std::atomic running_after_frames{0}; std::atomic running_since{0}; std::atomic cancelled_step{-1}; std::atomic cancelled_at{0}; @@ -216,7 +217,8 @@ static bool cancel_during_long_step(llama_dart_tts *tts, std::thread canceller([&] { while (!done.load()) { const int step = running_step.load(); - if (step >= 0 && since_origin() - running_since.load() >= cancel_after_ms * 1000.0) { + if (step >= 0 && running_after_frames.load() > 0 && + since_origin() - running_since.load() >= cancel_after_ms * 1000.0) { cancelled_at.store(since_origin()); cancelled_step.store(step); if (flag != nullptr) { @@ -234,8 +236,10 @@ static bool cancel_during_long_step(llama_dart_tts *tts, int step = 0; int64_t returned_at = 0; for (;; ++step) { + const int frames = progress.frames_generated; progress.struct_size = sizeof(progress); running_since.store(since_origin()); + running_after_frames.store(frames); running_step.store(step); status = llama_dart_tts_step(tts, &progress); running_step.store(-1); @@ -255,7 +259,7 @@ static bool cancel_during_long_step(llama_dart_tts *tts, return false; } if (cancelled_step.load() < 0) { - std::fprintf(stderr, "no step ran %.1f ms, so the in-step cancel never fired\n", + std::fprintf(stderr, "no step after a frame ran %.1f ms, so the in-step cancel never fired\n", cancel_after_ms); return false; } From b855b550479a10eb1cb464194349785e99bdeb90 Mon Sep 17 00:00:00 2001 From: Jhin Lee Date: Tue, 22 Sep 2026 21:19:36 -0400 Subject: [PATCH 3/4] fix: keep a TTS task cancelled once a step sees the cancel Lowering the cancel flag byte after the eval callback broke a decode let the step return OK with the truncated decode's PCM. The first read that sees a cancel now latches it for the rest of the task. The smoke lowers the byte right after the break at the end-of-speech and frame-72 decodes, and fails any step that returns OK after a break. --- README.md | 4 +- src/llama_dart_tts_eval_internal.h | 22 +++++++++ src/llama_dart_wrapper.cpp | 21 ++------ src/llama_dart_wrapper.h | 9 ++-- tests/tts_eval_test.cpp | 21 ++++++++ tools/tts_smoke.cpp | 79 ++++++++++++++++++++++++++++-- 6 files changed, 130 insertions(+), 26 deletions(-) diff --git a/README.md b/README.md index a45dc7f..c5fd015 100644 --- a/README.md +++ b/README.md @@ -285,7 +285,9 @@ It also checks the eval callback: uncancelled PCM stays byte-identical, frame steps stay whole, no chunk of the final decode exceeds a quarter of it, and a cancel issued a quarter of the way into a decode step, through `llama_dart_tts_cancel` or a cancel flag, makes that step return `CANCELLED` -within a third of the decode's duration. +within a third of the decode's duration. That holds at the end-of-speech and +frame-72 decodes even when the flag returns to zero right after the decode +breaks, and no step returns `OK` after a break. Omit `--gpu` for a CPU-only run. An optional speaker-reference audio path may appear before the final `--gpu` flag. diff --git a/src/llama_dart_tts_eval_internal.h b/src/llama_dart_tts_eval_internal.h index 335f2ba..a8abcd6 100644 --- a/src/llama_dart_tts_eval_internal.h +++ b/src/llama_dart_tts_eval_internal.h @@ -2,10 +2,32 @@ #include "ggml.h" +#include +#include + // Scheduler work between chunk boundaries inside one TTS step: multiply-adds // for MUL_MAT, elements for other computed nodes. static constexpr double llama_dart_tts_eval_budget = 2.5e9; +static inline bool llama_dart_tts_cancel_observed(std::atomic *latched, + const int8_t *flag) { + using atomic_flag_byte = std::atomic; + static_assert(sizeof(atomic_flag_byte) == sizeof(int8_t) && + alignof(atomic_flag_byte) == alignof(int8_t) && + atomic_flag_byte::is_always_lock_free, + "a caller-owned cancel byte must be readable atomically"); + if (latched->load(std::memory_order_acquire)) { + return true; + } + if (flag == nullptr || + reinterpret_cast(flag)->load( + std::memory_order_relaxed) == 0) { + return false; + } + latched->store(true, std::memory_order_release); + return true; +} + struct llama_dart_tts_eval_chunker { double budget = llama_dart_tts_eval_budget; double work = 0.0; diff --git a/src/llama_dart_wrapper.cpp b/src/llama_dart_wrapper.cpp index 8a13cff..cf5cfc6 100644 --- a/src/llama_dart_wrapper.cpp +++ b/src/llama_dart_wrapper.cpp @@ -96,20 +96,9 @@ struct llama_dart_tts { static void llama_dart_tts_release_task_resources(llama_dart_tts *tts); -static bool llama_dart_tts_flag_raised(const int8_t *flag) { - using atomic_flag_byte = std::atomic; - static_assert(sizeof(atomic_flag_byte) == sizeof(int8_t) && - alignof(atomic_flag_byte) == alignof(int8_t) && - atomic_flag_byte::is_always_lock_free, - "a caller-owned cancel byte must be readable atomically"); - return reinterpret_cast(flag)->load( - std::memory_order_relaxed) != 0; -} - -static bool llama_dart_tts_cancelled(const llama_dart_tts *tts) { - return tts->cancel_requested.load(std::memory_order_acquire) || - (tts->cancel_flag != nullptr && - llama_dart_tts_flag_raised(tts->cancel_flag)); +static bool llama_dart_tts_cancelled(llama_dart_tts *tts) { + return llama_dart_tts_cancel_observed(&tts->cancel_requested, + tts->cancel_flag); } struct llama_dart_tts_eval_scope; @@ -118,11 +107,11 @@ static thread_local llama_dart_tts_eval_scope *llama_dart_tts_active_eval = nullptr; struct llama_dart_tts_eval_scope { - const llama_dart_tts *tts; + llama_dart_tts *tts; llama_dart_tts_eval_chunker chunker; llama_dart_tts_eval_scope *previous; - explicit llama_dart_tts_eval_scope(const llama_dart_tts *task) + explicit llama_dart_tts_eval_scope(llama_dart_tts *task) : tts(task), previous(llama_dart_tts_active_eval) { llama_dart_tts_active_eval = this; } diff --git a/src/llama_dart_wrapper.h b/src/llama_dart_wrapper.h index fbf5d0d..7ea39a1 100644 --- a/src/llama_dart_wrapper.h +++ b/src/llama_dart_wrapper.h @@ -188,10 +188,11 @@ LLAMADART_API enum llama_dart_tts_status llama_dart_tts_step( LLAMADART_API void llama_dart_tts_cancel(struct llama_dart_tts * tts); // Attaches a caller-owned cancel byte to the active task. Any thread may set -// *flag to nonzero; steps then behave as after llama_dart_tts_cancel. The task -// reads *flag only inside llama_dart_tts_step and drops it when the task ends, -// restarts, is reset, or is freed; *flag must stay readable until then. -// Returns INVALID_STATE when no task is active. +// *flag to nonzero. The task reads *flag only inside llama_dart_tts_step. Once +// a read sees nonzero, the task behaves as after llama_dart_tts_cancel, even if +// *flag returns to zero. The task drops the flag when it ends, is reset, or is +// freed; *flag must stay readable until then. Returns INVALID_STATE when no +// task is active. LLAMADART_API enum llama_dart_tts_status llama_dart_tts_set_cancel_flag( struct llama_dart_tts * tts, const int8_t * flag); diff --git a/tests/tts_eval_test.cpp b/tests/tts_eval_test.cpp index 446f275..5523db3 100644 --- a/tests/tts_eval_test.cpp +++ b/tests/tts_eval_test.cpp @@ -8,7 +8,9 @@ #include "ggml-backend.h" #include "ggml.h" +#include #include +#include #include #include @@ -83,6 +85,24 @@ void test_cancel_waits_for_first_boundary(ggml_context *ctx) { assert(!llama_dart_tts_eval_answer(&chunker, sum, true, false)); } +void test_cancel_flag_latches() { + std::atomic latched{false}; + std::atomic flag{0}; + const int8_t *byte = reinterpret_cast(&flag); + assert(!llama_dart_tts_cancel_observed(&latched, nullptr)); + assert(!llama_dart_tts_cancel_observed(&latched, byte)); + assert(!latched.load()); + flag.store(1); + assert(llama_dart_tts_cancel_observed(&latched, byte)); + flag.store(0); + assert(llama_dart_tts_cancel_observed(&latched, byte)); + assert(llama_dart_tts_cancel_observed(&latched, nullptr)); + + std::atomic requested{true}; + assert(llama_dart_tts_cancel_observed(&requested, nullptr)); + assert(llama_dart_tts_cancel_observed(&requested, byte)); +} + struct probe { llama_dart_tts_eval_chunker chunker; bool cancelled = false; @@ -253,6 +273,7 @@ int main() { test_node_work(ctx); test_boundaries_follow_budget_and_mul_mat(ctx); test_cancel_waits_for_first_boundary(ctx); + test_cancel_flag_latches(); ggml_free(ctx); scheduled_graph graph; diff --git a/tools/tts_smoke.cpp b/tools/tts_smoke.cpp index d2920db..905c10f 100644 --- a/tools/tts_smoke.cpp +++ b/tools/tts_smoke.cpp @@ -25,16 +25,34 @@ static const char *const long_text = "enough that the decoder buffers a full window of codec frames before the " "end, so the flush runs in the middle of synthesis."; +static constexpr int qwen3_window_frames = 72; + static std::vector chunk_ends; +static int task_breaks = 0; +static std::atomic *lower_after_break = nullptr; static bool recording_eval_callback(struct ggml_tensor *tensor, bool ask, void *user_data) { const bool answer = llama_dart_tts_eval_callback(tensor, ask, user_data); if (!ask) { chunk_ends.push_back(steady_clock::now()); + if (!answer) { + ++task_breaks; + if (lower_after_break != nullptr) { + lower_after_break->store(0); + } + } } return answer; } +static bool succeeded_after_break(int step, llama_dart_tts_status status) { + if (task_breaks == 0 || status != LLAMA_DART_TTS_STATUS_OK) { + return false; + } + std::fprintf(stderr, "step %d returned OK after %d decode break(s)\n", step, task_breaks); + return true; +} + struct step_trace { double ms; int chunk_boundaries; @@ -116,7 +134,9 @@ static bool synthesize(llama_dart_tts *tts, std::fprintf(stderr, "failed to start TTS: %s\n", llama_dart_tts_last_error(tts)); return false; } + task_breaks = 0; progress->struct_size = sizeof(*progress); + int step = 0; do { chunk_ends.clear(); const auto started = steady_clock::now(); @@ -133,6 +153,9 @@ static bool synthesize(llama_dart_tts *tts, trace->push_back({elapsed_ms(started, returned), static_cast(chunk_ends.size()) - 1, longest}); } + if (succeeded_after_break(step++, status)) { + return false; + } if (status != LLAMA_DART_TTS_STATUS_OK) { std::fprintf(stderr, "TTS failed: %s\n", llama_dart_tts_last_error(tts)); return false; @@ -185,6 +208,7 @@ static bool synthesize(llama_dart_tts *tts, static bool cancel_during_long_step(llama_dart_tts *tts, llama_dart_tts_request *request, std::atomic *flag, + bool lower_flag_after_break, double cancel_after_ms, double return_within_ms, double *cancel_to_return_ms, @@ -202,6 +226,8 @@ static bool cancel_during_long_step(llama_dart_tts *tts, std::fprintf(stderr, "failed to attach cancel flag\n"); return false; } + task_breaks = 0; + lower_after_break = lower_flag_after_break ? flag : nullptr; std::atomic running_step{-1}; std::atomic running_after_frames{0}; std::atomic running_since{0}; @@ -235,6 +261,7 @@ static bool cancel_during_long_step(llama_dart_tts *tts, llama_dart_tts_status status = LLAMA_DART_TTS_STATUS_OK; int step = 0; int64_t returned_at = 0; + bool succeeded_after_a_break = false; for (;; ++step) { const int frames = progress.frames_generated; progress.struct_size = sizeof(progress); @@ -244,6 +271,7 @@ static bool cancel_during_long_step(llama_dart_tts *tts, status = llama_dart_tts_step(tts, &progress); running_step.store(-1); returned_at = since_origin(); + succeeded_after_a_break = succeeded_after_a_break || succeeded_after_break(step, status); if (status != LLAMA_DART_TTS_STATUS_OK || progress.state == LLAMA_DART_TTS_STATE_COMPLETED) { break; @@ -251,6 +279,7 @@ static bool cancel_during_long_step(llama_dart_tts *tts, } done.store(true); canceller.join(); + lower_after_break = nullptr; const bool same_step = cancelled_step.load() == step; *cancel_to_return_ms = (returned_at - cancelled_at.load()) / 1000.0; *frames_before_cancel = progress.frames_generated; @@ -263,6 +292,13 @@ static bool cancel_during_long_step(llama_dart_tts *tts, cancel_after_ms); return false; } + if (succeeded_after_a_break) { + return false; + } + if (lower_flag_after_break && task_breaks == 0) { + std::fprintf(stderr, "the decode never broke, so the flag was never lowered\n"); + return false; + } if (!same_step || status != LLAMA_DART_TTS_STATUS_CANCELLED || progress.state != LLAMA_DART_TTS_STATE_CANCELLED) { std::fprintf(stderr, @@ -430,16 +466,45 @@ int main(int argc, char **argv) { double final_cancel_ms = 0.0; int final_cancel_frames = 0; - if (!cancel_during_long_step(tts.get(), &request, nullptr, decode_ms / 4, decode_ms / 3, - &final_cancel_ms, &final_cancel_frames)) { + if (!cancel_during_long_step(tts.get(), &request, nullptr, false, decode_ms / 4, + decode_ms / 3, &final_cancel_ms, &final_cancel_frames)) { return 1; } llama_dart_tts_request long_request = request; long_request.text = long_text; long_request.text_length = std::char_traits::length(long_text); + + double lowered_final_ms = 0.0; + int lowered_final_frames = 0; + std::atomic lowered_final_flag{0}; + const bool lowered_final_ok = cancel_during_long_step( + tts.get(), &request, &lowered_final_flag, true, decode_ms / 4, decode_ms / 3, + &lowered_final_ms, &lowered_final_frames); + if (!lowered_final_ok) { + std::fprintf(stderr, "end-of-speech probe with the flag lowered after the break failed\n"); + } + double lowered_window_ms = 0.0; + int lowered_window_frames = 0; + std::atomic lowered_window_flag{0}; + bool lowered_window_ok = cancel_during_long_step( + tts.get(), &long_request, &lowered_window_flag, true, decode_ms / 4, decode_ms / 3, + &lowered_window_ms, &lowered_window_frames); + if (lowered_window_ok && lowered_window_frames != qwen3_window_frames - 1) { + std::fprintf(stderr, + "the lowered-flag cancel landed after %d frames, not in the frame-%d window\n", + lowered_window_frames, qwen3_window_frames); + lowered_window_ok = false; + } + if (!lowered_window_ok) { + std::fprintf(stderr, "window probe with the flag lowered after the break failed\n"); + } + if (!lowered_final_ok || !lowered_window_ok) { + return 1; + } + double window_cancel_ms = 0.0; int window_cancel_frames = 0; - if (!cancel_during_long_step(tts.get(), &long_request, &cancel_flag, decode_ms / 4, + if (!cancel_during_long_step(tts.get(), &long_request, &cancel_flag, false, decode_ms / 4, decode_ms / 3, &window_cancel_ms, &window_cancel_frames)) { return 1; } @@ -464,13 +529,17 @@ int main(int argc, char **argv) { "first_frames=%d second_frames=%d first_rms=%.6f second_rms=%.6f " "decode_ms=%.1f decode_chunks=%d longest_chunk_ms=%.1f " "final_cancel_ms=%.1f final_cancel_frames=%d " - "window_cancel_ms=%.1f window_cancel_frames=%d\n", + "window_cancel_ms=%.1f window_cancel_frames=%d " + "lowered_final_ms=%.1f lowered_final_frames=%d " + "lowered_window_ms=%.1f lowered_window_frames=%d\n", use_gpu ? "gpu" : "cpu", static_cast(info.model_type), info.sample_rate, first_pcm.size(), second_pcm.size(), first_progress.frames_generated, second_progress.frames_generated, first_rms, second_rms, decode_ms, first_trace.back().chunk_boundaries + 1, first_trace.back().longest_chunk_ms, final_cancel_ms, final_cancel_frames, - window_cancel_ms, window_cancel_frames); + window_cancel_ms, window_cancel_frames, + lowered_final_ms, lowered_final_frames, + lowered_window_ms, lowered_window_frames); return 0; } From d1a168f52aa567f041161be34de33f9ede7ae6ef Mon Sep 17 00:00:00 2001 From: Jhin Lee Date: Tue, 22 Sep 2026 21:19:36 -0400 Subject: [PATCH 4/4] chore: drop TTS eval comments that document no hazard --- CMakeLists.txt | 1 - src/llama_dart_tts_eval_internal.h | 4 ---- tests/tts_eval_test.cpp | 2 -- 3 files changed, 7 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index c65f412..2d9a505 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -228,7 +228,6 @@ if (LLAMADART_BUILD_TESTS) target_compile_features(llamadart_tts_eval_test PRIVATE cxx_std_17) target_include_directories(llamadart_tts_eval_test PRIVATE src) target_link_libraries(llamadart_tts_eval_test PRIVATE ggml) - # Beside the ggml backend modules that ggml_backend_load_all() finds. set_target_properties(llamadart_tts_eval_test PROPERTIES RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/bin") add_test(NAME llamadart_tts_eval_test COMMAND llamadart_tts_eval_test) diff --git a/src/llama_dart_tts_eval_internal.h b/src/llama_dart_tts_eval_internal.h index a8abcd6..d36a45f 100644 --- a/src/llama_dart_tts_eval_internal.h +++ b/src/llama_dart_tts_eval_internal.h @@ -5,8 +5,6 @@ #include #include -// Scheduler work between chunk boundaries inside one TTS step: multiply-adds -// for MUL_MAT, elements for other computed nodes. static constexpr double llama_dart_tts_eval_budget = 2.5e9; static inline bool llama_dart_tts_cancel_observed(std::atomic *latched, @@ -51,8 +49,6 @@ static inline double llama_dart_tts_eval_node_work(const ggml_tensor *node) { } } -// Answers one ggml_backend_sched_eval_callback query for a step in progress. -// // Budget boundaries fall only after a MUL_MAT node. A boundary splits any // fusion that spans it, which can change the output; the CPU and Metal // backends fuse nothing that continues past a MUL_MAT. diff --git a/tests/tts_eval_test.cpp b/tests/tts_eval_test.cpp index 5523db3..da83333 100644 --- a/tests/tts_eval_test.cpp +++ b/tests/tts_eval_test.cpp @@ -222,8 +222,6 @@ struct scheduled_graph { } }; -// The budget first overflows on the first RMS_NORM, so the first boundary -// moves to the second MUL_MAT; every later MUL_MAT then ends a chunk. constexpr double kSchedulerBudget = kMatMulWork + 0.5 * kRowWork; void test_chunked_compute_matches_unchunked(scheduled_graph &g) {