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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/workflows/validate_wrapper.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
11 changes: 10 additions & 1 deletion CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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()
Expand Down
18 changes: 15 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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 \
Expand All @@ -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.

Expand Down
82 changes: 82 additions & 0 deletions src/llama_dart_tts_eval_internal.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,82 @@
#pragma once

#include "ggml.h"

#include <atomic>
#include <cstdint>

static constexpr double llama_dart_tts_eval_budget = 2.5e9;

static inline bool llama_dart_tts_cancel_observed(std::atomic<bool> *latched,
const int8_t *flag) {
using atomic_flag_byte = std::atomic<int8_t>;
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<const atomic_flag_byte *>(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<double>(ggml_nelements(node)) *
static_cast<double>(node->src[0]->ne[0]);
default:
return static_cast<double>(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;
}
90 changes: 81 additions & 9 deletions src/llama_dart_wrapper.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -77,6 +78,7 @@ struct llama_dart_tts {
llama_sampler *sampler = nullptr;
mtmd_bitmap *speaker = nullptr;
std::atomic<bool> 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;
Expand All @@ -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) {
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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");
}
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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");
Expand Down Expand Up @@ -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) {
Expand Down
26 changes: 24 additions & 2 deletions src/llama_dart_wrapper.h
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
8 changes: 8 additions & 0 deletions tests/tts_api_test.c
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}
Loading
Loading