From a9df03d08ed1dc554d2a00c08b95cbb23ad18ee3 Mon Sep 17 00:00:00 2001 From: Xuan Son Nguyen Date: Fri, 31 Jul 2026 19:59:41 +0200 Subject: [PATCH] code2wav preserve kv between calls --- tools/mtmd/clip-model.h | 2 +- tools/mtmd/clip.cpp | 67 +++++-- tools/mtmd/clip.h | 6 +- tools/mtmd/models/models.h | 52 +++++- tools/mtmd/models/qwen3tts-gen.cpp | 272 ++++++++++++++++++++++------- tools/mtmd/mtmd.cpp | 13 +- tools/mtmd/mtmd.h | 8 +- tools/tts/qwen3-tts.cpp | 55 +++--- 8 files changed, 365 insertions(+), 110 deletions(-) diff --git a/tools/mtmd/clip-model.h b/tools/mtmd/clip-model.h index 13095c7074..e3e3c6601a 100644 --- a/tools/mtmd/clip-model.h +++ b/tools/mtmd/clip-model.h @@ -145,7 +145,7 @@ struct clip_hparams { int32_t wav_upsample_n_block = 0; int32_t wav_dac_n_block = 0; int32_t wav_dac_n_res = 0; - int32_t wav_c2w_window_frames = 0; // number of frames code2wav decodes per call + int32_t wav_tfm_swa = 0; // pre_transformer's KV cache size, in frames // mimo-v2.5: LLM-side connector (input_local_transformer) int32_t audio_local_n_layer = 0; diff --git a/tools/mtmd/clip.cpp b/tools/mtmd/clip.cpp index 2d20740dd5..9709ed2075 100644 --- a/tools/mtmd/clip.cpp +++ b/tools/mtmd/clip.cpp @@ -1723,8 +1723,8 @@ struct clip_model_loader { hparams.wav_upsample_n_block = 2; hparams.wav_dac_n_block = 4; hparams.wav_dac_n_res = 3; - // code2wav keeps no state between calls, so it decodes this many frames each time to get enough left context - hparams.wav_c2w_window_frames = 24; + // matches the reference decoder's sliding_window (speech_tokenizer/config.json) + hparams.wav_tfm_swa = 72; } break; case PROJECTOR_TYPE_PADDLEOCR: { @@ -4728,21 +4728,40 @@ bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) { if (params->gen_process == CLIP_GEN_PROCESS_CODE2WAV) { GGML_ASSERT(params->codes != nullptr); - // the caller sends codes frame-major (frame 0's codes, then frame 1's, ...). the graph wants them group-major (all frames of codebook 0, then all frames of codebook 1, ...). pad the front with code 0 if fewer frames than the window. - const int64_t n_codes = model.gen_code_head_w->ne[2] + 1; - const int64_t n_frames_w = hparams.wav_c2w_window_frames; + // the caller sends codes frame-major (frame 0's codes, then frame + // 1's, ...), for however many frames it has (up to the window + // size). the graph wants them group-major (all frames of codebook + // 0, then all frames of codebook 1, ...), padded to exactly one + // window. real frames go first (so their RoPE positions/state + // stay in sequence), code 0 pads the rear if there are fewer -- + // the corresponding tail of out_audio is trimmed off below. + const int64_t n_codes = model.gen_code_head_w->ne[2] + 1; + const int64_t n_frames_w = hparams.wav_tfm_swa; const int64_t n_frames = (int64_t) params->codes->size() / n_codes; - const int64_t n_use = std::min(n_frames, n_frames_w); - const int64_t dst0 = n_frames_w - n_use; // front padding - const int64_t src0 = n_frames - n_use; // newest frames + GGML_ASSERT(n_frames > 0 && n_frames <= n_frames_w); std::vector codes(n_frames_w * n_codes, 0); - for (int64_t f = 0; f < n_use; f++) { + for (int64_t f = 0; f < n_frames; f++) { for (int64_t g = 0; g < n_codes; g++) { - codes[g * n_frames_w + dst0 + f] = (*params->codes)[(src0 + f) * n_codes + g]; + codes[g * n_frames_w + f] = (*params->codes)[f * n_codes + g]; } } set_input_i32("inp_codes", codes); + + // upload the state carried over from the previous call, or + // zero-fill on a cold start (no previous state, or wrong size) + size_t offset = 0; + for (const auto & slot : list_c2w_state_slots(hparams, model)) { + ggml_tensor * t = get_inp_tensor(("state_in_" + slot.name).c_str()); + const size_t nb = ggml_nbytes(t); + if (params->state_in && params->state_in->size() >= offset + nb) { + ggml_backend_tensor_set(t, params->state_in->data() + offset, 0, nb); + } else { + std::vector zeros(nb, 0); + ggml_backend_tensor_set(t, zeros.data(), 0, nb); + } + offset += nb; + } } else { std::vector code0 = { params->code0 }; set_input_i32("inp_code0", code0); @@ -5223,6 +5242,34 @@ bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) { auto & out_audio = *params->out_audio; out_audio.resize(ggml_nelements(audio)); ggml_backend_tensor_get(audio, out_audio.data(), 0, ggml_nbytes(audio)); + + // a short (rear-padded) batch only has real audio for its real frames; + // the tail generated from the code-0 padding is discarded + const int64_t n_codes = model.gen_code_head_w->ne[2] + 1; + const int64_t n_frames_w = hparams.wav_tfm_swa; + const int64_t n_frames = (int64_t) params->codes->size() / n_codes; + if (n_frames < n_frames_w) { + const size_t hop = out_audio.size() / n_frames_w; + out_audio.resize((size_t) n_frames * hop); + } + } + if (params->state_out != nullptr) { + auto & state_out = *params->state_out; + size_t total = 0; + for (const auto & slot : list_c2w_state_slots(hparams, model)) { + total += (size_t) (slot.ne0 * slot.ne1) * sizeof(float); + } + state_out.resize(total); + size_t offset = 0; + for (const auto & slot : list_c2w_state_slots(hparams, model)) { + ggml_tensor * t = ggml_graph_get_tensor(gf, ("state_out_" + slot.name).c_str()); + if (t == nullptr) { + GGML_ABORT("state_out requested but graph has no \"state_out_%s\" tensor", slot.name.c_str()); + } + const size_t nb = ggml_nbytes(t); + ggml_backend_tensor_get(t, state_out.data() + offset, 0, nb); + offset += nb; + } } // diff --git a/tools/mtmd/clip.h b/tools/mtmd/clip.h index 8c1b96d666..e969b2b9b1 100644 --- a/tools/mtmd/clip.h +++ b/tools/mtmd/clip.h @@ -108,9 +108,13 @@ struct clip_encode_params { std::vector * out_codes = nullptr; // CODE2WAV: codes holds this frame's 16 RVQ codes, out_audio receives the - // decoded PCM samples (F32) + // decoded PCM samples (F32). state_in is the state from the previous + // call (null or wrong size means cold start, state is zero-filled). + // state_out receives the state to pass into the next call. const std::vector * codes = nullptr; std::vector * out_audio = nullptr; + const std::vector * state_in = nullptr; + std::vector * state_out = nullptr; }; bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params); diff --git a/tools/mtmd/models/models.h b/tools/mtmd/models/models.h index 90b4b2f5f2..cf422a4d84 100644 --- a/tools/mtmd/models/models.h +++ b/tools/mtmd/models/models.h @@ -2,6 +2,11 @@ #include "../clip-graph.h" +#include +#include +#include +#include + /* * IMPORTANT: The mtmd module does NOT accept pull requests that are fully or predominantly AI-generated. * We encourage human contributors to ensure the quality and reliability of the codebase. @@ -286,28 +291,57 @@ struct clip_graph_qwen3tts_gen : clip_graph { // // code2wav: RVQ codes -> raw PCM (quantizer + pre_conv + pre_transformer + upsample + DAC). - // Each call decodes a window of T frames from scratch, RoPE positions 0..T-1. No state is kept between calls. A window of about 24 frames gives enough left context without a persistent KV cache. + // Processes one frame per call (T=1). Every causal conv/transpose-conv and + // the pre_transformer's attention carry real state across calls (state_in + // / state_out), so there is no left-context zero-padding at call boundaries. // struct code2wav : clip_graph { code2wav(const clip_graph & parent) : clip_graph(parent) {} ggml_cgraph * build() override { GGML_ABORT("call decode() instead"); } - ggml_tensor * causal_conv1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int dilation) const; - ggml_tensor * causal_conv1d_dw(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b) const; - ggml_tensor * causal_conv_transpose1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int stride) const; + // state carried in from the previous call (by slot name, see + // list_c2w_state_slots()), filled in by build() before calling decode() + std::map state_in; + // state to persist for the next call, filled in by decode(); each + // entry's tensor must be added to the graph outputs by build() + mutable std::vector> state_out; + + // stateful conv ops: read their left-context (or overlap-add tail, for + // the transpose conv) from state_in[state_name], append the updated + // state to state_out + ggml_tensor * causal_conv1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int dilation, const std::string & state_name) const; + ggml_tensor * causal_conv1d_dw(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, const std::string & state_name) const; + ggml_tensor * causal_conv_transpose1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int stride, const std::string & state_name) const; ggml_tensor * snake(ggml_tensor * x, ggml_tensor * alpha, ggml_tensor * beta) const; ggml_tensor * quant_decode(ggml_tensor * inp_codes) const; - ggml_tensor * tfm_layer_forward(ggml_tensor * cur, const clip_layer & layer, ggml_tensor * pos, ggml_tensor * mask) const; - ggml_tensor * convnext_block(ggml_tensor * x, const clip_code2wav::upsample_block & blk) const; - ggml_tensor * dac_res_unit(ggml_tensor * x, const clip_code2wav::dac_res & res, int dilation) const; + // il: layer index, used to look up this layer's K/V slice of the sliding-window KV cache + ggml_tensor * tfm_layer_forward(ggml_tensor * cur, const clip_layer & layer, int il) const; + ggml_tensor * convnext_block(ggml_tensor * x, const clip_code2wav::upsample_block & blk, const std::string & state_prefix) const; + ggml_tensor * dac_res_unit(ggml_tensor * x, const clip_code2wav::dac_res & res, int dilation, const std::string & state_name) const; - // inp_codes: [T, n_codes] I32, T frames of RVQ codes (group-major: all T frames of codebook 0, then all T frames of codebook 1, etc.). - // returns audio samples for all T frames, [n_samples] F32, clamped to [-1, 1]. + // inp_codes: [1, n_codes] I32, one frame of RVQ codes. + // returns this frame's audio samples, [n_samples] F32, clamped to [-1, 1]. ggml_tensor * decode(ggml_tensor * inp_codes) const; }; }; +// one persisted state buffer used by code2wav (conv left-context, transpose-conv +// overlap tail, one layer's K/V slice of the pre_transformer's sliding-window +// cache, or its running position counter), named so build() and clip.cpp's +// (de)serialization agree on layout. ne0/ne1 is the tensor's own shape (conv +// states are time-first [T, C] like their input; KV cache is channel-first +// [C, T] like q/k/v). +struct c2w_state_slot { + std::string name; + int64_t ne0; + int64_t ne1; +}; +// computed purely from hparams/model tensor shapes, no graph needed -- used by +// both clip_graph_qwen3tts_gen::code2wav::decode() (to create/collect state +// tensors) and clip.cpp (to (de)serialize the flat state_data byte buffer) +std::vector list_c2w_state_slots(const clip_hparams & hparams, const clip_model & model); + struct clip_graph_kimik25 : clip_graph { clip_graph_kimik25(clip_ctx * ctx, const clip_image_f32 & img) : clip_graph(ctx, img) {} ggml_cgraph * build() override; diff --git a/tools/mtmd/models/qwen3tts-gen.cpp b/tools/mtmd/models/qwen3tts-gen.cpp index f63f39fc8c..6753be372d 100644 --- a/tools/mtmd/models/qwen3tts-gen.cpp +++ b/tools/mtmd/models/qwen3tts-gen.cpp @@ -292,53 +292,87 @@ ggml_tensor * clip_graph_qwen3tts_gen::code_gen::step( return cache_set(out_code_cache, pos, sampled); } -// causal conv1d, stride 1: left-pad (K-1)*dilation zeros, then a plain conv. +// causal conv1d, stride 1: prepend the persisted left-context (from the +// previous call) instead of zero-padding, then a plain (unpadded) conv. // x: [T, IC] (T-first, matches ggml_conv_1d's native layout). w: [K, IC, OC]. -// returns [T, OC]. -ggml_tensor * clip_graph_qwen3tts_gen::code2wav::causal_conv1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int dilation) const { +// state_name empty means K == 1, no left-context needed. returns [T, OC]. +ggml_tensor * clip_graph_qwen3tts_gen::code2wav::causal_conv1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int dilation, const std::string & state_name) const { const int K = (int) w->ne[0]; const int pad = (K - 1) * dilation; - ggml_tensor * x_pad = pad > 0 ? ggml_pad_ext(ctx0, x, pad, 0, 0, 0, 0, 0, 0, 0) : x; - ggml_tensor * y = ggml_conv_1d(ctx0, w, x_pad, 1, 0, dilation); // [T, OC, 1] + ggml_tensor * x_full = x; + if (pad > 0) { + ggml_tensor * left = state_in.at(state_name); // [pad, IC] + x_full = ggml_concat(ctx0, left, x, 0); + } + ggml_tensor * y = ggml_conv_1d(ctx0, w, x_full, 1, 0, dilation); // [T, OC, 1] y = ggml_reshape_2d(ctx0, y, y->ne[0], y->ne[1]); if (b) { y = ggml_add(ctx0, y, ggml_reshape_2d(ctx0, b, 1, b->ne[0])); } + if (pad > 0) { + ggml_tensor * new_left = ggml_cont(ctx0, ggml_view_2d(ctx0, x_full, pad, x_full->ne[1], x_full->nb[1], + (size_t) (x_full->ne[0] - pad) * x_full->nb[0])); + state_out.push_back({state_name, new_left}); + } return y; } // causal depthwise conv1d, stride 1, dilation 1, kernel from w's shape. -// x: [T, C]. w: [K, 1, C]. returns [T, C]. -ggml_tensor * clip_graph_qwen3tts_gen::code2wav::causal_conv1d_dw(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b) const { +// x: [T, C]. w: [K, 1, C]. returns [T, C]. see causal_conv1d for the state contract. +ggml_tensor * clip_graph_qwen3tts_gen::code2wav::causal_conv1d_dw(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, const std::string & state_name) const { const int K = (int) w->ne[0]; const int pad = K - 1; - ggml_tensor * x_pad = pad > 0 ? ggml_pad_ext(ctx0, x, pad, 0, 0, 0, 0, 0, 0, 0) : x; - ggml_tensor * y = ggml_conv_1d_dw(ctx0, w, x_pad, 1, 0, 1); // [T, C, 1] + ggml_tensor * x_full = x; + if (pad > 0) { + ggml_tensor * left = state_in.at(state_name); // [pad, C] + x_full = ggml_concat(ctx0, left, x, 0); + } + ggml_tensor * y = ggml_conv_1d_dw(ctx0, w, x_full, 1, 0, 1); // [T, C, 1] y = ggml_reshape_2d(ctx0, y, y->ne[0], y->ne[1]); if (b) { y = ggml_add(ctx0, y, ggml_reshape_2d(ctx0, b, 1, b->ne[0])); } + if (pad > 0) { + ggml_tensor * new_left = ggml_cont(ctx0, ggml_view_2d(ctx0, x_full, pad, x_full->ne[1], x_full->nb[1], + (size_t) (x_full->ne[0] - pad) * x_full->nb[0])); + state_out.push_back({state_name, new_left}); + } return y; } -// causal ConvTranspose1d: full (non-causal) transpose conv, then trim the -// right (kernel - stride) frames that would otherwise leak future context. -// x: [T, IC] (plain matrix). w: [K, OC, IC]. returns [T * stride, OC]. -ggml_tensor * clip_graph_qwen3tts_gen::code2wav::causal_conv_transpose1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int stride) const { - const int K = (int) w->ne[0]; - const int trim = K - stride; +// causal ConvTranspose1d with a persisted overlap-add tail: transposed conv +// output windows from adjacent input frames overlap by (kernel - stride) +// samples. The overlap that would otherwise leak into the next call's output +// is carried forward as state and added in, instead of being discarded. +// x: [T, IC] (plain matrix). w: [K, OC, IC]. state_name empty means +// K == stride (no overlap, e.g. the upsample blocks here). returns [T * stride, OC]. +ggml_tensor * clip_graph_qwen3tts_gen::code2wav::causal_conv_transpose1d(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int stride, const std::string & state_name) const { + const int K = (int) w->ne[0]; + const int trim = K - stride; + const int64_t emit_len = x->ne[0] * stride; - ggml_tensor * y = ggml_conv_transpose_1d(ctx0, w, x, stride, 0, 1); // [T*stride + trim, OC, 1, 1] + ggml_tensor * y = ggml_conv_transpose_1d(ctx0, w, x, stride, 0, 1); // [emit_len + trim, OC, 1, 1] y = ggml_reshape_2d(ctx0, y, y->ne[0], y->ne[1]); + + ggml_tensor * out = y; if (trim > 0) { - y = ggml_cont(ctx0, ggml_view_2d(ctx0, y, y->ne[0] - trim, y->ne[1], y->nb[1], 0)); + ggml_tensor * tail = state_in.at(state_name); // [trim, OC] + ggml_tensor * head = ggml_add(ctx0, ggml_view_2d(ctx0, y, trim, y->ne[1], y->nb[1], 0), tail); + if (emit_len > trim) { + ggml_tensor * middle = ggml_view_2d(ctx0, y, emit_len - trim, y->ne[1], y->nb[1], (size_t) trim * y->nb[0]); + out = ggml_concat(ctx0, head, middle, 0); + } else { + out = head; + } + ggml_tensor * new_tail = ggml_cont(ctx0, ggml_view_2d(ctx0, y, trim, y->ne[1], y->nb[1], (size_t) emit_len * y->nb[0])); + state_out.push_back({state_name, new_tail}); } if (b) { - y = ggml_add(ctx0, y, ggml_reshape_2d(ctx0, b, 1, b->ne[0])); + out = ggml_add(ctx0, out, ggml_reshape_2d(ctx0, b, 1, b->ne[0])); } - return y; + return out; } // SnakeBeta activation: y = x + sin(alpha*x)^2 * inv_beta (alpha/inv_beta @@ -382,35 +416,85 @@ ggml_tensor * clip_graph_qwen3tts_gen::code2wav::quant_decode(ggml_tensor * inp_ return hidden; } -// one pre_transformer layer, causal over the T-frame window (RoPE positions 0..T-1) -ggml_tensor * clip_graph_qwen3tts_gen::code2wav::tfm_layer_forward(ggml_tensor * cur, const clip_layer & layer, ggml_tensor * pos, ggml_tensor * mask) const { +// one pre_transformer layer, over a batch of N = sliding_window new frames. +// Attention runs over [persisted (W-1)-frame prefix from the previous batch] +// + [this batch's N new frames]: the prefix supplies left-context for the +// batch's own early frames, a banded causal mask keeps each query within its +// W-frame window, and RoPE uses a real, ever-increasing position counter +// (persisted) so the prefix's already-baked-in rotation phase lines up with +// the new frames'. The persisted state for the next batch is just the last +// (W-1) frames of this batch (the old prefix falls out of every query's +// window once a full new batch has passed). +ggml_tensor * clip_graph_qwen3tts_gen::code2wav::tfm_layer_forward(ggml_tensor * cur, const clip_layer & layer, int il) const { const int n_head = hparams.wav_tfm_n_head; const int n_head_kv = hparams.wav_tfm_n_head_kv; const int64_t d_head = layer.q_w->ne[1] / n_head; const float kq_scale = 1.0f / sqrtf((float) d_head); - const int64_t T = cur->ne[1]; + const int64_t W = hparams.wav_tfm_swa; // == N, frames per batch + const int64_t N = cur->ne[1]; + const int64_t prefix = W - 1; + const int64_t total_kv = prefix + N; ggml_tensor * residual = cur; ggml_tensor * h = ggml_rms_norm(ctx0, cur, hparams.wav_tfm_eps); h = ggml_mul(ctx0, h, layer.ln_1_w); - ggml_tensor * q = ggml_mul_mat(ctx0, layer.q_w, h); - ggml_tensor * k = ggml_mul_mat(ctx0, layer.k_w, h); - ggml_tensor * v = ggml_mul_mat(ctx0, layer.v_w, h); + ggml_tensor * q = ggml_mul_mat(ctx0, layer.q_w, h); // [n_head*d_head, N] + ggml_tensor * k = ggml_mul_mat(ctx0, layer.k_w, h); // [n_head_kv*d_head, N] + ggml_tensor * v = ggml_mul_mat(ctx0, layer.v_w, h); // [n_head_kv*d_head, N] - q = ggml_reshape_3d(ctx0, q, d_head, n_head, T); - k = ggml_reshape_3d(ctx0, k, d_head, n_head_kv, T); + q = ggml_reshape_3d(ctx0, q, d_head, n_head, N); + k = ggml_reshape_3d(ctx0, k, d_head, n_head_kv, N); + + // real, ever-increasing positions: base (persisted) .. base+N-1 + ggml_tensor * base = ggml_reshape_1d(ctx0, state_in.at("tfm_pos"), 1); + ggml_tensor * offset = ggml_arange(ctx0, 0.0f, (float) N, 1.0f); + ggml_tensor * pos = ggml_cast(ctx0, ggml_add(ctx0, offset, base), GGML_TYPE_I32); q = ggml_rope_ext(ctx0, q, pos, nullptr, (int) d_head, GGML_ROPE_TYPE_NEOX, 0, hparams.wav_tfm_rope_theta, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); k = ggml_rope_ext(ctx0, k, pos, nullptr, (int) d_head, GGML_ROPE_TYPE_NEOX, 0, hparams.wav_tfm_rope_theta, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); - ggml_tensor * q_cur = ggml_reshape_4d(ctx0, q, d_head, n_head, T, 1); - ggml_tensor * k_cur = ggml_reshape_4d(ctx0, k, d_head, n_head_kv, T, 1); - ggml_tensor * v_cur = ggml_reshape_4d(ctx0, v, d_head, n_head_kv, T, 1); + // the position counter is layer-independent -- only push it once, from layer 0 + if (il == 0) { + state_out.push_back({"tfm_pos", ggml_scale_bias(ctx0, state_in.at("tfm_pos"), 1.0f, (float) N)}); + } - ggml_tensor * attn_out = build_attn(layer.o_w, layer.o_b, q_cur, k_cur, v_cur, mask, kq_scale, 0); + ggml_tensor * k_new = ggml_reshape_2d(ctx0, k, d_head * n_head_kv, N); + ggml_tensor * v_new = ggml_reshape_2d(ctx0, v, d_head * n_head_kv, N); + + ggml_tensor * old_k = state_in.at("tfm_k_" + std::to_string(il)); // [d_head*n_head_kv, W-1] + ggml_tensor * old_v = state_in.at("tfm_v_" + std::to_string(il)); + + ggml_tensor * k_full = ggml_concat(ctx0, old_k, k_new, 1); // [.., prefix+N] + ggml_tensor * v_full = ggml_concat(ctx0, old_v, v_new, 1); + + // next batch's prefix: the last (W-1) frames of this batch + state_out.push_back({"tfm_k_" + std::to_string(il), + ggml_cont(ctx0, ggml_view_2d(ctx0, k_full, k_full->ne[0], prefix, k_full->nb[1], (size_t) N * k_full->nb[1]))}); + state_out.push_back({"tfm_v_" + std::to_string(il), + ggml_cont(ctx0, ggml_view_2d(ctx0, v_full, v_full->ne[0], prefix, v_full->nb[1], (size_t) N * v_full->nb[1]))}); + + // banded causal mask over [total_kv keys, N queries]: key j (local index + // into the concatenated prefix+batch) is visible to query i (local index + // into the batch, offset by `prefix` so it lines up with its own slot in + // the concatenated sequence) iff 0 <= (prefix+i) - j < W + ggml_tensor * pos_k = ggml_reshape_2d(ctx0, ggml_arange(ctx0, 0.0f, (float) total_kv, 1.0f), total_kv, 1); + ggml_tensor * pos_q = ggml_reshape_2d(ctx0, ggml_arange(ctx0, (float) prefix, (float) (prefix + N), 1.0f), 1, N); + ggml_tensor * pos_q_grid = ggml_repeat_4d(ctx0, pos_q, total_kv, N, 1, 1); + ggml_tensor * diff = ggml_sub(ctx0, pos_q_grid, pos_k); // [total_kv, N] + + ggml_tensor * causal_keep = ggml_step(ctx0, ggml_scale_bias(ctx0, diff, 1.0f, 0.5f)); // diff >= 0 + ggml_tensor * in_window = ggml_step(ctx0, ggml_scale_bias(ctx0, diff, -1.0f, (float) W - 0.5f)); // diff < W + ggml_tensor * keep = ggml_mul(ctx0, causal_keep, in_window); + ggml_tensor * mask = ggml_reshape_4d(ctx0, ggml_log(ctx0, keep), total_kv, N, 1, 1); // 0 = keep, -inf = masked + + ggml_tensor * q_cur = ggml_reshape_4d(ctx0, q, d_head, n_head, N, 1); + ggml_tensor * k_cur = ggml_reshape_4d(ctx0, k_full, d_head, n_head_kv, total_kv, 1); + ggml_tensor * v_cur = ggml_reshape_4d(ctx0, v_full, d_head, n_head_kv, total_kv, 1); + + ggml_tensor * attn_out = build_attn(layer.o_w, layer.o_b, q_cur, k_cur, v_cur, mask, kq_scale, il); if (layer.ls_1_w) { attn_out = ggml_mul(ctx0, attn_out, layer.ls_1_w); } @@ -433,10 +517,10 @@ ggml_tensor * clip_graph_qwen3tts_gen::code2wav::tfm_layer_forward(ggml_tensor * // dwconv -> LayerNorm -> pwconv1 -> GELU -> pwconv2 -> layer scale -> residual. // x: [T, C] T-first; LayerNorm/pwconv need C on ne0, so this transposes in // and back out around them. -ggml_tensor * clip_graph_qwen3tts_gen::code2wav::convnext_block(ggml_tensor * x, const clip_code2wav::upsample_block & blk) const { +ggml_tensor * clip_graph_qwen3tts_gen::code2wav::convnext_block(ggml_tensor * x, const clip_code2wav::upsample_block & blk, const std::string & state_prefix) const { ggml_tensor * residual = x; - ggml_tensor * h = causal_conv1d_dw(x, blk.dwconv_w, blk.dwconv_b); // [T, C] + ggml_tensor * h = causal_conv1d_dw(x, blk.dwconv_w, blk.dwconv_b, state_prefix + "_dwconv"); // [T, C] ggml_tensor * hc = ggml_cont(ctx0, ggml_transpose(ctx0, h)); // [C, T] hc = ggml_norm(ctx0, hc, 1e-6f); @@ -456,75 +540,73 @@ ggml_tensor * clip_graph_qwen3tts_gen::code2wav::convnext_block(ggml_tensor * x, // SnakeBeta -> dilated causal conv (k=7) -> SnakeBeta -> pointwise causal conv (k=1) -> residual. // x: [T, C]. returns [T, C]. -ggml_tensor * clip_graph_qwen3tts_gen::code2wav::dac_res_unit(ggml_tensor * x, const clip_code2wav::dac_res & res, int dilation) const { +ggml_tensor * clip_graph_qwen3tts_gen::code2wav::dac_res_unit(ggml_tensor * x, const clip_code2wav::dac_res & res, int dilation, const std::string & state_name) const { ggml_tensor * residual = x; ggml_tensor * h = snake(x, res.act1_alpha, res.act1_beta); - h = causal_conv1d(h, res.conv1_w, res.conv1_b, dilation); + h = causal_conv1d(h, res.conv1_w, res.conv1_b, dilation, state_name); h = snake(h, res.act2_alpha, res.act2_beta); - h = causal_conv1d(h, res.conv2_w, res.conv2_b, 1); + h = causal_conv1d(h, res.conv2_w, res.conv2_b, 1, ""); // k=1, no left-context needed return ggml_add(ctx0, residual, h); } -// RVQ codes -> raw PCM for the whole T-frame window +// RVQ codes -> raw PCM for a batch of N = sliding_window frames ggml_tensor * clip_graph_qwen3tts_gen::code2wav::decode(ggml_tensor * inp_codes) const { const auto & c2w = model.c2w; - const int64_t T = inp_codes->ne[0]; - // 1. quantizer decode: T frames of 16 codes -> [512, T] (C-first) + // 1. quantizer decode: N frames of 16 codes -> [512, N] (C-first) ggml_tensor * hidden = quant_decode(inp_codes); - // 2. pre_conv: [512, T] -> T-first [T, 512] -> causal conv k=3 -> [T, 1024] - ggml_tensor * x = ggml_cont(ctx0, ggml_transpose(ctx0, hidden)); // [T, 512] - x = causal_conv1d(x, c2w.pre_conv_w, c2w.pre_conv_b, 1); // [T, 1024] + // 2. pre_conv: [512, N] -> T-first [N, 512] -> causal conv k=3 -> [N, 1024] + ggml_tensor * x = ggml_cont(ctx0, ggml_transpose(ctx0, hidden)); // [N, 512] + x = causal_conv1d(x, c2w.pre_conv_w, c2w.pre_conv_b, 1, "pre_conv"); // [N, 1024] cb(x, "wav_pre_conv_out", -1); - // 3. pre_transformer: back to C-first [1024, T], project down to hidden_size, run 8 layers causal over the T-frame window, project back up - ggml_tensor * cur = ggml_cont(ctx0, ggml_transpose(ctx0, x)); // [1024, T] + // 3. pre_transformer: back to C-first [1024, N], project down to hidden_size, + // run 8 layers with a persisted sliding-window KV cache, project back up + ggml_tensor * cur = ggml_cont(ctx0, ggml_transpose(ctx0, x)); // [1024, N] cur = ggml_mul_mat(ctx0, c2w.tfm_in_proj_w, cur); - cur = ggml_add(ctx0, cur, c2w.tfm_in_proj_b); // [512 (tfm hidden), T] - - ggml_tensor * pos = ggml_cast(ctx0, ggml_arange(ctx0, 0.0f, (float) T, 1.0f), GGML_TYPE_I32); - ggml_tensor * tri = ggml_tri(ctx0, ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, T, T), 1.0f), - GGML_TRI_TYPE_LOWER_DIAG); - ggml_tensor * mask = ggml_reshape_4d(ctx0, ggml_log(ctx0, tri), T, T, 1, 1); // 0 = keep, -inf = masked + cur = ggml_add(ctx0, cur, c2w.tfm_in_proj_b); // [512 (tfm hidden), N] for (int il = 0; il < hparams.wav_tfm_n_layer; il++) { - cur = tfm_layer_forward(cur, c2w.tfm_layers[il], pos, mask); + cur = tfm_layer_forward(cur, c2w.tfm_layers[il], il); } + cur = ggml_rms_norm(ctx0, cur, hparams.wav_tfm_eps); cur = ggml_mul(ctx0, cur, c2w.tfm_output_norm_w); cur = ggml_mul_mat(ctx0, c2w.tfm_out_proj_w, cur); - cur = ggml_add(ctx0, cur, c2w.tfm_out_proj_b); // [1024, T] + cur = ggml_add(ctx0, cur, c2w.tfm_out_proj_b); // [1024, N] cb(cur, "wav_tfm_out", -1); - // 4. upsample: 2x (causal ConvTranspose1d, stride 2 + ConvNeXt block), back to T-first - x = ggml_cont(ctx0, ggml_transpose(ctx0, cur)); // [T, 1024] + // 4. upsample: 2x (causal ConvTranspose1d, stride 2 + ConvNeXt block), back to T-first. + // kernel == stride here, so there's no transpose-conv overlap, no state needed. + x = ggml_cont(ctx0, ggml_transpose(ctx0, cur)); // [N, 1024] for (size_t il = 0; il < c2w.upsample.size(); il++) { const auto & up = c2w.upsample[il]; - x = causal_conv_transpose1d(x, up.conv_w, up.conv_b, 2); - x = convnext_block(x, up); + x = causal_conv_transpose1d(x, up.conv_w, up.conv_b, 2, ""); + x = convnext_block(x, up, "up" + std::to_string(il)); cb(x, "wav_upsample_out", (int) il); } // 5. DAC decoder: conv_pre -> n blocks (SnakeBeta -> ConvTranspose1d -> 3 res units) -> conv_post - static constexpr int DAC_STRIDES[4] = { 8, 5, 4, 3 }; - static constexpr int DAC_DILATIONS[3] = { 1, 3, 9 }; + static constexpr int DAC_DILATIONS[3] = { 1, 3, 9 }; - x = causal_conv1d(x, c2w.dac_entry_w, c2w.dac_entry_b, 1); + x = causal_conv1d(x, c2w.dac_entry_w, c2w.dac_entry_b, 1, "dac_entry"); cb(x, "wav_dac_entry_out", -1); for (size_t il = 0; il < c2w.dac.size(); il++) { const auto & blk = c2w.dac[il]; + const int stride = (int) (blk.conv_w->ne[0] / 2); // kernel == 2*stride for all 4 blocks + const std::string blk_name = "dac" + std::to_string(il); x = snake(x, blk.snake_alpha, blk.snake_beta); - x = causal_conv_transpose1d(x, blk.conv_w, blk.conv_b, DAC_STRIDES[il]); + x = causal_conv_transpose1d(x, blk.conv_w, blk.conv_b, stride, blk_name + "_tail"); for (size_t ir = 0; ir < blk.res.size(); ir++) { - x = dac_res_unit(x, blk.res[ir], DAC_DILATIONS[ir]); + x = dac_res_unit(x, blk.res[ir], DAC_DILATIONS[ir], blk_name + "_res" + std::to_string(ir)); } cb(x, "wav_dac_block_out", (int) il); } x = snake(x, c2w.dac_post_snake_alpha, c2w.dac_post_snake_beta); - x = causal_conv1d(x, c2w.dac_post_conv_w, c2w.dac_post_conv_b, 1); // [n_samples, 1] + x = causal_conv1d(x, c2w.dac_post_conv_w, c2w.dac_post_conv_b, 1, "dac_post_conv"); // [n_samples, 1] x = ggml_clamp(ctx0, x, -1.0f, 1.0f); x = ggml_reshape_1d(ctx0, x, x->ne[0]); @@ -532,6 +614,54 @@ ggml_tensor * clip_graph_qwen3tts_gen::code2wav::decode(ggml_tensor * inp_codes) return x; } +// enumerates code2wav's persisted state buffers: the running RoPE position +// counter, one K/V slot per pre_transformer layer, and one left-context (or +// overlap-add tail) slot per stateful conv. Pure hparams/tensor-shape lookup, +// no graph needed -- shared by build() (to create/collect the state tensors) +// and clip.cpp (to (de)serialize the flat state_data byte buffer). +std::vector list_c2w_state_slots(const clip_hparams & hparams, const clip_model & model) { + const auto & c2w = model.c2w; + std::vector slots; + + slots.push_back({"tfm_pos", 1, 1}); + + // persisted prefix is (W-1) frames: the code2wav batch itself supplies the + // other N=W frames of context for its own later frames (see tfm_layer_forward) + const int64_t d_head = c2w.tfm_layers[0].q_w->ne[1] / hparams.wav_tfm_n_head; + const int64_t kv_ch = d_head * hparams.wav_tfm_n_head_kv; + const int64_t prefix = hparams.wav_tfm_swa - 1; + for (int il = 0; il < hparams.wav_tfm_n_layer; il++) { + slots.push_back({"tfm_k_" + std::to_string(il), kv_ch, prefix}); + slots.push_back({"tfm_v_" + std::to_string(il), kv_ch, prefix}); + } + + slots.push_back({"pre_conv", c2w.pre_conv_w->ne[0] - 1, c2w.pre_conv_w->ne[1]}); + + for (size_t il = 0; il < c2w.upsample.size(); il++) { + const auto & up = c2w.upsample[il]; + slots.push_back({"up" + std::to_string(il) + "_dwconv", up.dwconv_w->ne[0] - 1, up.dwconv_w->ne[2]}); + } + + slots.push_back({"dac_entry", c2w.dac_entry_w->ne[0] - 1, c2w.dac_entry_w->ne[1]}); + + static constexpr int DAC_DILATIONS[3] = { 1, 3, 9 }; + for (size_t il = 0; il < c2w.dac.size(); il++) { + const auto & blk = c2w.dac[il]; + const int64_t stride = blk.conv_w->ne[0] / 2; // kernel == 2*stride for all 4 blocks + const std::string blk_name = "dac" + std::to_string(il); + slots.push_back({blk_name + "_tail", stride, blk.conv_w->ne[1]}); + for (size_t ir = 0; ir < blk.res.size(); ir++) { + const auto & res = blk.res[ir]; + slots.push_back({blk_name + "_res" + std::to_string(ir), + (res.conv1_w->ne[0] - 1) * DAC_DILATIONS[ir], res.conv1_w->ne[1]}); + } + } + + slots.push_back({"dac_post_conv", c2w.dac_post_conv_w->ne[0] - 1, c2w.dac_post_conv_w->ne[1]}); + + return slots; +} + // master build(): switches on gen_process to construct either the code_gen sub-graph (backbone hidden state -> 16 RVQ codes + next-step embd) or the code2wav sub-graph (16 RVQ codes -> raw PCM), both hosted in this one clip_ctx. ggml_cgraph * clip_graph_qwen3tts_gen::build() { GGML_ASSERT(n_batch == 1); // this module only ever processes one frame at a time @@ -539,16 +669,30 @@ ggml_cgraph * clip_graph_qwen3tts_gen::build() { if (gen_process == CLIP_GEN_PROCESS_CODE2WAV) { const int64_t n_acoustic = model.gen_code_head_w->ne[2]; // 15 const int n_codes = (int) n_acoustic + 1; // 16 - const int n_frames = hparams.wav_c2w_window_frames; + const int n_frames = hparams.wav_tfm_swa; // frames per batch, == the attention window ggml_tensor * inp_codes = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_frames, n_codes); ggml_set_name(inp_codes, "inp_codes"); ggml_set_input(inp_codes); - ggml_tensor * out_audio = code2wav(*this).decode(inp_codes); + code2wav cg(*this); + for (const auto & slot : list_c2w_state_slots(hparams, model)) { + ggml_tensor * t = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, slot.ne0, slot.ne1); + ggml_set_name(t, ("state_in_" + slot.name).c_str()); + ggml_set_input(t); + cg.state_in[slot.name] = t; + } + + ggml_tensor * out_audio = cg.decode(inp_codes); ggml_set_name(out_audio, "out_audio"); ggml_set_output(out_audio); ggml_build_forward_expand(gf, out_audio); + + for (auto & slot : cg.state_out) { + ggml_set_name(slot.second, ("state_out_" + slot.first).c_str()); + ggml_set_output(slot.second); + ggml_build_forward_expand(gf, slot.second); + } return gf; } diff --git a/tools/mtmd/mtmd.cpp b/tools/mtmd/mtmd.cpp index 889b5e2b3d..2bcd58104e 100644 --- a/tools/mtmd/mtmd.cpp +++ b/tools/mtmd/mtmd.cpp @@ -267,6 +267,7 @@ struct mtmd_context { std::vector gen_out_codes; // this frame's 16 sampled codes (CODE_GEN) std::vector gen_out_embd; // next-step hidden state fed back to backbone (CODE_GEN) std::vector gen_out_audio; // decoded PCM samples for the current frame (CODE2WAV) + std::vector gen_out_state; // state to feed into the next CODE2WAV call bool print_timings; int n_threads; @@ -1632,6 +1633,10 @@ static int32_t mtmd_gen_audio_process_impl(mtmd_context * ctx, const mtmd_gen_in return 1; } std::vector in_codes(inp->codes, inp->codes + inp->n_codes); + std::vector in_state; + if (inp->state_data) { + in_state.assign(inp->state_data, inp->state_data + inp->state_size); + } // code2wav has no hidden-state input, the batch entry is an unused placeholder clip_image_f32 dummy; @@ -1648,14 +1653,18 @@ static int32_t mtmd_gen_audio_process_impl(mtmd_context * ctx, const mtmd_gen_in params.gen_process = CLIP_GEN_PROCESS_CODE2WAV; params.codes = &in_codes; params.out_audio = &ctx->gen_out_audio; + params.state_in = inp->state_data ? &in_state : nullptr; + params.state_out = &ctx->gen_out_state; if (!clip_encode(ctx_clip, ¶ms)) { LOG_ERR("%s: clip_encode failed (code2wav)\n", __func__); return 1; } - out->audio = ctx->gen_out_audio.data(); - out->n_samples = ctx->gen_out_audio.size(); + out->audio = ctx->gen_out_audio.data(); + out->n_samples = ctx->gen_out_audio.size(); + out->state_data = (const char *) ctx->gen_out_state.data(); + out->state_size = ctx->gen_out_state.size(); return 0; } diff --git a/tools/mtmd/mtmd.h b/tools/mtmd/mtmd.h index 35b8170a08..6056ac56f5 100644 --- a/tools/mtmd/mtmd.h +++ b/tools/mtmd/mtmd.h @@ -351,7 +351,9 @@ struct mtmd_gen_inp { // for MTMD_GEN_PROCESS_TYPE_CODE2WAV int32_t * codes; - size_t n_codes; + size_t n_codes; + const char * state_data; + size_t state_size; }; struct mtmd_gen_out { // note: output memory is allocated by the context, valid until next process() call @@ -364,7 +366,9 @@ struct mtmd_gen_out { // for MTMD_GEN_PROCESS_TYPE_CODE2WAV const float * audio; - size_t n_samples; + size_t n_samples; + const char * state_data; + size_t state_size; }; MTMD_API int32_t mtmd_gen_audio_process(mtmd_context * ctx, const struct mtmd_gen_inp * inp, diff --git a/tools/tts/qwen3-tts.cpp b/tools/tts/qwen3-tts.cpp index 407f06cf74..64518e67af 100644 --- a/tools/tts/qwen3-tts.cpp +++ b/tools/tts/qwen3-tts.cpp @@ -14,6 +14,7 @@ #include "gguf.h" #include +#include #include #include #include @@ -105,24 +106,32 @@ static void save_wav16(const char * path, const std::vector & pcm, int ra fclose(f); } -// runs one CODE2WAV process() call on whatever frames are buffered, appends -// the result to audio_out, then clears the buffer -static bool flush_codes(mtmd_context * mctx, std::vector & codes_buf, std::vector & audio_out) { - if (codes_buf.empty()) { - return true; - } +// runs one CODE2WAV process() call on a batch of frames' codes (frame-major; +// the model wants exactly one window's worth, clip.cpp front-pads a shorter +// batch), carrying the persisted state (KV cache + conv left-context) across +// batches; appends the resulting PCM to audio_out and updates state for the +// next batch +static bool code2wav_step(mtmd_context * mctx, const std::vector & codes, std::vector & state, + std::vector & audio_out) { mtmd_gen_inp inp{}; - inp.type = MTMD_GEN_PROCESS_TYPE_CODE2WAV; - inp.codes = codes_buf.data(); - inp.n_codes = codes_buf.size(); + inp.type = MTMD_GEN_PROCESS_TYPE_CODE2WAV; + inp.codes = const_cast(codes.data()); + inp.n_codes = codes.size(); + inp.state_data = state.empty() ? nullptr : (const char *) state.data(); + inp.state_size = state.size(); mtmd_gen_out out{}; - if (mtmd_gen_audio_process(mctx, &inp, &out) != 0) { + const auto t0 = std::chrono::steady_clock::now(); + const int rc = mtmd_gen_audio_process(mctx, &inp, &out); + const auto t1 = std::chrono::steady_clock::now(); + const double ms = std::chrono::duration(t1 - t0).count(); + LOG_INF("code2wav: %zu codes -> %zu samples in %.1f ms\n", codes.size() / 16, out.n_samples, ms); + if (rc != 0) { LOG_ERR("code2wav process failed\n"); return false; } audio_out.insert(audio_out.end(), out.audio, out.audio + out.n_samples); - codes_buf.clear(); + state.assign(out.state_data, out.state_data + out.state_size); return true; } @@ -134,8 +143,6 @@ int main(int argc, char ** argv) { std::string lang = "english"; int max_new = 512; int n_gpu = 999; - const int n_codes_per_frame = 16; - const int window_frames = 24; for (int i = 1; i < argc; i++) { auto next = [&](const char * flag) -> const char * { @@ -278,11 +285,16 @@ int main(int argc, char ** argv) { overlay.push_back(row(tts_eos)); overlay.push_back(row(tts_pad)); + // matches hparams.wav_tfm_sliding_window hardcoded in clip.cpp; code2wav + // batches exactly this many frames per call + const size_t C2W_WINDOW_FRAMES = 72; + // AR loop: sample c0 among the semantic codec rows plus eos, hand the // hidden state to the code predictor (GEN_CODE), buffer the 16 codes it - // returns, feed its embedding back to the talker. Once window_frames - // frames are buffered, run CODE2WAV to turn them into PCM. + // returns. Once a full window is buffered, run CODE2WAV on it, carrying + // its state (KV cache + conv left-context) across batches. std::vector audio; + std::vector c2w_state; std::vector codes_buf; std::vector h((size_t) n_embd), fb((size_t) n_embd); int n_frames = 0; @@ -314,9 +326,9 @@ int main(int argc, char ** argv) { codes_buf.insert(codes_buf.end(), out.codes, out.codes + out.n_codes); memcpy(fb.data(), out.embd, (size_t) n_embd * sizeof(float)); - if ((int) (codes_buf.size() / n_codes_per_frame) >= window_frames) { - if (!flush_codes(mctx, codes_buf, audio)) return 1; - LOG_INF("flushed a %d-frame window, %zu samples so far\n", window_frames, audio.size()); + if (codes_buf.size() / 16 >= C2W_WINDOW_FRAMES) { + if (!code2wav_step(mctx, codes_buf, c2w_state, audio)) return 1; + codes_buf.clear(); } const auto & ov = overlay[std::min((size_t) n_frames, overlay.size() - 1)]; @@ -336,9 +348,10 @@ int main(int argc, char ** argv) { if (llama_decode(lctx, batch) != 0) { LOG_ERR("decode failed at frame %d\n", n_frames); return 1; } } - // flush whatever's left, less than a full window (front-padded with - // code 0 by clip.cpp) - if (!flush_codes(mctx, codes_buf, audio)) return 1; + // flush whatever's left, less than a full window (front-padded with code 0 by clip.cpp) + if (!codes_buf.empty()) { + if (!code2wav_step(mctx, codes_buf, c2w_state, audio)) return 1; + } LOG_INF("generated %d frames, %zu samples (%.2f s)\n", n_frames, audio.size(), (double) audio.size() / 24000.0); save_wav16(out_path, audio, 24000);