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..2d9a505 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -224,6 +224,14 @@ 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) + 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 +241,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..c5fd015 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,13 @@ 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. 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 new file mode 100644 index 0000000..d36a45f --- /dev/null +++ b/src/llama_dart_tts_eval_internal.h @@ -0,0 +1,82 @@ +#pragma once + +#include "ggml.h" + +#include +#include + +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; + 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)); + } +} + +// 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..cf5cfc6 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,31 @@ struct llama_dart_tts { static void llama_dart_tts_release_task_resources(llama_dart_tts *tts); +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; + +static thread_local llama_dart_tts_eval_scope *llama_dart_tts_active_eval = + nullptr; + +struct llama_dart_tts_eval_scope { + llama_dart_tts *tts; + llama_dart_tts_eval_chunker chunker; + llama_dart_tts_eval_scope *previous; + + explicit llama_dart_tts_eval_scope(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 +158,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 +235,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 +923,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 +979,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 +1019,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..7ea39a1 100644 --- a/src/llama_dart_wrapper.h +++ b/src/llama_dart_wrapper.h @@ -181,10 +181,32 @@ 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. 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); + +// 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..da83333 --- /dev/null +++ b/tests/tts_eval_test.cpp @@ -0,0 +1,282 @@ +#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 +#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)); +} + +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; + 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; + } +}; + +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); + test_cancel_flag_latches(); + 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..905c10f 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,54 @@ #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 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; + 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 +128,34 @@ 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; } + task_breaks = 0; progress->struct_size = sizeof(*progress); + int step = 0; 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 (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; @@ -137,6 +205,116 @@ 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, + bool lower_flag_after_break, + 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; + } + 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}; + 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 && 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) { + 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; + bool succeeded_after_a_break = false; + 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); + 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; + } + } + 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; + 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 after a frame ran %.1f ms, so the in-step cancel never fired\n", + 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, + "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 +358,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 +387,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 +407,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 +431,81 @@ 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, 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, false, decode_ms / 4, + decode_ms / 3, &window_cancel_ms, &window_cancel_frames)) { return 1; } @@ -263,10 +526,20 @@ 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 " + "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); + 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, + lowered_final_ms, lowered_final_frames, + lowered_window_ms, lowered_window_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",