diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp index 88ca7e8e8e..66e1e51a76 100644 --- a/tools/mtmd/mtmd-helper-gen.cpp +++ b/tools/mtmd/mtmd-helper-gen.cpp @@ -6,6 +6,7 @@ #include #include +#include #include #include @@ -16,9 +17,6 @@ // // Audio generation helpers // -// model-specific pipeline logic is dispatched on mtmd_gen_audio_get_info(mctx).type; -// the public surface (init/free/reset/set_input/step/get_output) stays model-agnostic -// static llama_token find_special_token(const llama_vocab * vocab, const std::string & piece) { const int32_t n = llama_vocab_n_tokens(vocab); @@ -51,15 +49,303 @@ static void write_wav16(std::vector & buf, const std::vector & pcm, } } -struct mtmd_helper_gen_audio { +class mtmd_gen_audio_pipeline { +public: + mtmd_gen_audio_pipeline(llama_context * lctx, mtmd_context * mctx) + : lctx(lctx), mctx(mctx), model(llama_get_model(lctx)), vocab(llama_model_get_vocab(model)), + n_embd(llama_model_n_embd(model)), info(mtmd_gen_audio_get_info(mctx)) {} + virtual ~mtmd_gen_audio_pipeline() = default; + + virtual void reset() = 0; + virtual int32_t set_input(const mtmd_helper_gen_audio_inp * inp) = 0; + // sampled may be LLAMA_TOKEN_NULL for pipelines with no discrete backbone token + // (e.g. continuous/diffusion models); such pipelines read whatever they need + // directly off h_state_in instead + virtual int32_t step(llama_token sampled, const float * h_state_in, const float ** h_state_out) = 0; + virtual int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len) = 0; + +protected: llama_context * lctx; mtmd_context * mctx; const llama_model * model; const llama_vocab * vocab; - int n_embd = 0; + int n_embd; mtmd_gen_audio_info info; +}; - // qwen3tts: vocab specials fixed across the whole session, looked up once +// Qwen3-TTS: dual-track discrete AR (backbone codec_0 + MTP code-predictor for the +// remaining 15 codebooks) into a windowed causal conv/transformer decode (code2wav) +class qwen3tts_gen_audio_pipeline : public mtmd_gen_audio_pipeline { +public: + using mtmd_gen_audio_pipeline::mtmd_gen_audio_pipeline; + + void reset() override { + pos = 0; + codes_buf.clear(); + c2w_state.clear(); + audio_pcm.clear(); + overlay.clear(); + overlay_idx = 0; + h_state_buf.clear(); + out_buf.clear(); + } + + int32_t set_input(const mtmd_helper_gen_audio_inp * inp) override { + reset(); + + if (!ensure_cache()) { + return 1; + } + + const std::string lang = inp->lang ? inp->lang : "english"; + const llama_token c_lang = find_special_token(vocab, ("<|codec_language_" + lang + "|>").c_str()); + if (c_lang == LLAMA_TOKEN_NULL) { + LOG_ERR("mtmd_helper_gen_audio: unknown language '%s'\n", lang.c_str()); + return 1; + } + + std::vector speaker_embd; + if (inp->speaker_ref) { + if (!encode_speaker(inp->speaker_ref, speaker_embd)) { + return 1; + } + } + + const int n_e = n_embd; + auto row = [&](llama_token t) { + return std::vector(tok_embd.begin() + (size_t) t * n_e, + tok_embd.begin() + (size_t) (t + 1) * n_e); + }; + auto sum_row = [&](llama_token a, llama_token b) { + std::vector va = row(a), vb = row(b); + for (int i = 0; i < n_e; i++) va[(size_t) i] += vb[(size_t) i]; + return va; + }; + auto sum_vec = [&](llama_token a, const std::vector & vb) { + std::vector va = row(a); + for (int i = 0; i < n_e; i++) va[(size_t) i] += vb[(size_t) i]; + return va; + }; + + // upstream chat wrap, then slices: [0:3] role, [3:-5] utterance body + const std::string full = "<|im_start|>assistant\n" + std::string(inp->prompt, inp->prompt_len) + + "<|im_end|>\n<|im_start|>assistant\n"; + std::vector ids(full.size() + 16); + int n_ids = llama_tokenize(vocab, full.c_str(), (int32_t) full.size(), ids.data(), (int32_t) ids.size(), + false, true); + if (n_ids < 8) { + LOG_ERR("mtmd_helper_gen_audio: tokenization failed\n"); + return 1; + } + ids.resize((size_t) n_ids); + + std::vector> prompt; + for (int i = 0; i < 3; i++) prompt.push_back(row(ids[(size_t) i])); + prompt.push_back(sum_row(tts_pad, c_think)); + prompt.push_back(sum_row(tts_pad, c_think_b)); + prompt.push_back(sum_row(tts_pad, c_lang)); + prompt.push_back(sum_row(tts_pad, c_think_e)); + if (!speaker_embd.empty()) prompt.push_back(sum_vec(tts_pad, speaker_embd)); + prompt.push_back(sum_row(tts_bos, codec_pad)); + for (int i = 3; i < n_ids - 5; i++) prompt.push_back(sum_row(ids[(size_t) i], codec_pad)); + prompt.push_back(sum_row(tts_eos, codec_pad)); + prompt.push_back(sum_row(tts_pad, codec_bos)); + + const int n_prompt = (int) prompt.size(); + + // the talker rides the qwen3vl interleaved mrope: positions carry + // n_pos_per_embd sections laid out [section * n_tokens + i], all + // equal for a pure text/codec stream + mrope = llama_model_rope_type(model) == LLAMA_ROPE_TYPE_MROPE || + llama_model_rope_type(model) == LLAMA_ROPE_TYPE_IMROPE; + const int n_pos_per_embd = mrope ? 4 : 1; + + std::vector embd_buf((size_t) n_prompt * (size_t) n_e); + for (int i = 0; i < n_prompt; i++) { + memcpy(embd_buf.data() + (size_t) i * n_e, prompt[(size_t) i].data(), (size_t) n_e * sizeof(float)); + } + + decode_embd_batch batch_embd(embd_buf.data(), n_prompt, n_pos_per_embd, n_e); + if (mrope) batch_embd.set_position_mrope_1d(0, 0); + else batch_embd.set_position_normal(0, 0); + batch_embd.batch.logits[n_prompt - 1] = 1; + + if (llama_decode(lctx, batch_embd.batch) != 0) { + LOG_ERR("mtmd_helper_gen_audio: prefill decode failed\n"); + return 1; + } + + pos = n_prompt; + top_k = inp->top_k > 0 ? inp->top_k : 50; + top_p = inp->top_p > 0 ? inp->top_p : 1.0f; + out_type = inp->out_type; + + // the text stream keeps flowing during generation: the input after + // frame k adds trailing text row k on top of the codes embedding, + // then tts_eos, then tts_pad once the utterance is spent + for (int i = 3; i < n_ids - 5; i++) overlay.push_back(row(ids[(size_t) i])); + overlay.push_back(row(tts_eos)); + overlay.push_back(row(tts_pad)); + + return 0; + } + + int32_t step(llama_token sampled, const float * h_state_in, const float ** h_state_out) override { + mtmd_gen_inp inp{}; + inp.type = MTMD_GEN_PROCESS_TYPE_GEN_CODE; + inp.code0 = sampled - codec_0; + inp.embd = const_cast(h_state_in); + inp.top_k = top_k; + inp.top_p = top_p; + mtmd_gen_out out{}; + if (mtmd_gen_audio_process(mctx, &inp, &out) != 0) { + LOG_ERR("mtmd_helper_gen_audio: gen_code process failed\n"); + return 1; + } + + codes_buf.insert(codes_buf.end(), out.codes, out.codes + out.n_codes); + if (out.n_codes > 0 && codes_buf.size() / out.n_codes >= window_frames) { + if (!flush_c2w()) { + return 1; + } + } + + std::vector fb(out.embd, out.embd + n_embd); + const auto & ov = overlay[std::min(overlay_idx, overlay.size() - 1)]; + for (int i = 0; i < n_embd; i++) fb[(size_t) i] += ov[(size_t) i]; + overlay_idx++; + + const int n_pos_per_embd = mrope ? 4 : 1; + decode_embd_batch batch_embd(fb.data(), 1, n_pos_per_embd, n_embd); + if (mrope) batch_embd.set_position_mrope_1d(pos, 0); + else batch_embd.set_position_normal(pos, 0); + batch_embd.batch.logits[0] = 1; + pos++; + + if (llama_decode(lctx, batch_embd.batch) != 0) { + LOG_ERR("mtmd_helper_gen_audio: decode failed\n"); + return 1; + } + + const float * he = llama_get_embeddings_ith(lctx, -1); + h_state_buf.assign(he, he + n_embd); + *h_state_out = h_state_buf.data(); + + return 0; + } + + int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len) override { + if (!flush_c2w()) { + return 1; + } + + *out_sample_rate = info.sample_rate; + + if (out_type == MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM) { + *out_data = (const char *) audio_pcm.data(); + *out_data_len = audio_pcm.size() * sizeof(float); + return 0; + } + + out_buf.clear(); + write_wav16(out_buf, audio_pcm, info.sample_rate); + *out_data = out_buf.data(); + *out_data_len = out_buf.size(); + return 0; + } + +private: + bool ensure_cache() { + if (specials_ok) { + return true; + } + codec_0 = find_special_token(vocab, "<|codec_0|>"); + codec_bos = find_special_token(vocab, "<|codec_bos|>"); + codec_eos = find_special_token(vocab, "<|codec_eos_token|>"); + codec_pad = find_special_token(vocab, "<|codec_pad|>"); + c_think = find_special_token(vocab, "<|codec_think|>"); + c_think_b = find_special_token(vocab, "<|codec_think_bos|>"); + c_think_e = find_special_token(vocab, "<|codec_think_eos|>"); + tts_pad = find_special_token(vocab, ""); + tts_bos = find_special_token(vocab, ""); + tts_eos = find_special_token(vocab, ""); + for (llama_token t : { codec_0, codec_bos, codec_eos, codec_pad, + c_think, c_think_b, c_think_e, + tts_pad, tts_bos, tts_eos }) { + if (t == LLAMA_TOKEN_NULL) { + LOG_ERR("mtmd_helper_gen_audio: missing a required special token in vocab\n"); + return false; + } + } + const uint32_t n_tok_embd = llama_model_get_tok_embd(model, nullptr); + if (n_tok_embd == 0) { + LOG_ERR("mtmd_helper_gen_audio: model has no token embeddings\n"); + return false; + } + tok_embd.resize(n_tok_embd); + llama_model_get_tok_embd(model, tok_embd.data()); + specials_ok = true; + return true; + } + + // encodes a reference wav (already loaded as a bitmap) through the mmproj's + // speaker encoder, returning the single x-vector embedding row it produces + bool encode_speaker(mtmd_bitmap * bitmap, std::vector & out) { + if (!mtmd_support_audio(mctx)) { + LOG_ERR("mtmd_helper_gen_audio: mmproj has no speaker/audio encoder\n"); + return false; + } + const std::string marker = mtmd_default_marker(); + mtmd_input_text text{ marker.c_str(), marker.size(), false, true }; + mtmd_input_chunks * chunks = mtmd_input_chunks_init(); + const mtmd_bitmap * bptr = bitmap; + bool ok = mtmd_tokenize(mctx, chunks, &text, &bptr, 1) == 0; + if (ok) { + ok = false; + for (size_t i = 0; i < mtmd_input_chunks_size(chunks); i++) { + const mtmd_input_chunk * chunk = mtmd_input_chunks_get(chunks, i); + if (mtmd_input_chunk_get_type(chunk) != MTMD_INPUT_CHUNK_TYPE_AUDIO) { + continue; + } + if (mtmd_encode_chunk(mctx, chunk) != 0) { + LOG_ERR("mtmd_helper_gen_audio: speaker encode failed\n"); + break; + } + const float * embd = mtmd_get_output_embd(mctx); + const size_t n = (size_t) llama_model_n_embd_inp(model) * mtmd_input_chunk_get_n_tokens(chunk); + out.assign(embd, embd + n); + ok = true; + break; + } + } + mtmd_input_chunks_free(chunks); + return ok; + } + + // runs one CODE2WAV process() call on whatever is currently buffered, carrying + // the persisted state (KV cache + conv left-context) across batches + bool flush_c2w() { + if (codes_buf.empty()) { + return true; + } + mtmd_gen_inp inp{}; + inp.type = MTMD_GEN_PROCESS_TYPE_CODE2WAV; + inp.codes = codes_buf.data(); + inp.n_codes = codes_buf.size(); + inp.state_data = c2w_state.empty() ? nullptr : (const char *) c2w_state.data(); + inp.state_size = c2w_state.size(); + mtmd_gen_out out{}; + if (mtmd_gen_audio_process(mctx, &inp, &out) != 0) { + LOG_ERR("mtmd_helper_gen_audio: code2wav process failed\n"); + return false; + } + audio_pcm.insert(audio_pcm.end(), out.audio, out.audio + out.n_samples); + c2w_state.assign(out.state_data, out.state_data + out.state_size); + codes_buf.clear(); + return true; + } + + // vocab specials fixed across the whole session, looked up once bool specials_ok = false; llama_token codec_0 = LLAMA_TOKEN_NULL; llama_token codec_bos = LLAMA_TOKEN_NULL; @@ -92,104 +378,22 @@ struct mtmd_helper_gen_audio { std::vector out_buf; }; -static bool qwen3tts_ensure_cache(mtmd_helper_gen_audio * ctx) { - if (ctx->specials_ok) { - return true; +static std::unique_ptr make_pipeline(llama_context * lctx, mtmd_context * mctx) { + switch (mtmd_gen_audio_get_info(mctx).type) { + case MTMD_GEN_AUDIO_TYPE_QWEN3TTS: + return std::unique_ptr(new qwen3tts_gen_audio_pipeline(lctx, mctx)); + default: + return nullptr; } - ctx->codec_0 = find_special_token(ctx->vocab, "<|codec_0|>"); - ctx->codec_bos = find_special_token(ctx->vocab, "<|codec_bos|>"); - ctx->codec_eos = find_special_token(ctx->vocab, "<|codec_eos_token|>"); - ctx->codec_pad = find_special_token(ctx->vocab, "<|codec_pad|>"); - ctx->c_think = find_special_token(ctx->vocab, "<|codec_think|>"); - ctx->c_think_b = find_special_token(ctx->vocab, "<|codec_think_bos|>"); - ctx->c_think_e = find_special_token(ctx->vocab, "<|codec_think_eos|>"); - ctx->tts_pad = find_special_token(ctx->vocab, ""); - ctx->tts_bos = find_special_token(ctx->vocab, ""); - ctx->tts_eos = find_special_token(ctx->vocab, ""); - for (llama_token t : { ctx->codec_0, ctx->codec_bos, ctx->codec_eos, ctx->codec_pad, - ctx->c_think, ctx->c_think_b, ctx->c_think_e, - ctx->tts_pad, ctx->tts_bos, ctx->tts_eos }) { - if (t == LLAMA_TOKEN_NULL) { - LOG_ERR("mtmd_helper_gen_audio: missing a required special token in vocab\n"); - return false; - } - } - const uint32_t n_tok_embd = llama_model_get_tok_embd(ctx->model, nullptr); - if (n_tok_embd == 0) { - LOG_ERR("mtmd_helper_gen_audio: model has no token embeddings\n"); - return false; - } - ctx->tok_embd.resize(n_tok_embd); - llama_model_get_tok_embd(ctx->model, ctx->tok_embd.data()); - ctx->specials_ok = true; - return true; } -// encodes a reference wav (already loaded as a bitmap) through the mmproj's -// speaker encoder, returning the single x-vector embedding row it produces -static bool qwen3tts_encode_speaker(mtmd_helper_gen_audio * ctx, mtmd_bitmap * bitmap, std::vector & out) { - if (!mtmd_support_audio(ctx->mctx)) { - LOG_ERR("mtmd_helper_gen_audio: mmproj has no speaker/audio encoder\n"); - return false; - } - const std::string marker = mtmd_default_marker(); - mtmd_input_text text{ marker.c_str(), marker.size(), false, true }; - mtmd_input_chunks * chunks = mtmd_input_chunks_init(); - const mtmd_bitmap * bptr = bitmap; - bool ok = mtmd_tokenize(ctx->mctx, chunks, &text, &bptr, 1) == 0; - if (ok) { - ok = false; - for (size_t i = 0; i < mtmd_input_chunks_size(chunks); i++) { - const mtmd_input_chunk * chunk = mtmd_input_chunks_get(chunks, i); - if (mtmd_input_chunk_get_type(chunk) != MTMD_INPUT_CHUNK_TYPE_AUDIO) { - continue; - } - if (mtmd_encode_chunk(ctx->mctx, chunk) != 0) { - LOG_ERR("mtmd_helper_gen_audio: speaker encode failed\n"); - break; - } - const float * embd = mtmd_get_output_embd(ctx->mctx); - const size_t n = (size_t) llama_model_n_embd_inp(ctx->model) * mtmd_input_chunk_get_n_tokens(chunk); - out.assign(embd, embd + n); - ok = true; - break; - } - } - mtmd_input_chunks_free(chunks); - return ok; -} - -// runs one CODE2WAV process() call on whatever is currently buffered, carrying -// the persisted state (KV cache + conv left-context) across batches -static bool qwen3tts_flush_c2w(mtmd_helper_gen_audio * ctx) { - if (ctx->codes_buf.empty()) { - return true; - } - mtmd_gen_inp inp{}; - inp.type = MTMD_GEN_PROCESS_TYPE_CODE2WAV; - inp.codes = ctx->codes_buf.data(); - inp.n_codes = ctx->codes_buf.size(); - inp.state_data = ctx->c2w_state.empty() ? nullptr : (const char *) ctx->c2w_state.data(); - inp.state_size = ctx->c2w_state.size(); - mtmd_gen_out out{}; - if (mtmd_gen_audio_process(ctx->mctx, &inp, &out) != 0) { - LOG_ERR("mtmd_helper_gen_audio: code2wav process failed\n"); - return false; - } - ctx->audio_pcm.insert(ctx->audio_pcm.end(), out.audio, out.audio + out.n_samples); - ctx->c2w_state.assign(out.state_data, out.state_data + out.state_size); - ctx->codes_buf.clear(); - return true; -} +struct mtmd_helper_gen_audio { + std::unique_ptr pipeline; +}; mtmd_helper_gen_audio * mtmd_helper_gen_audio_init(struct llama_context * lctx, struct mtmd_context * mctx) { - auto * ctx = new mtmd_helper_gen_audio(); - ctx->lctx = lctx; - ctx->mctx = mctx; - ctx->model = llama_get_model(lctx); - ctx->vocab = llama_model_get_vocab(ctx->model); - ctx->n_embd = llama_model_n_embd(ctx->model); - ctx->info = mtmd_gen_audio_get_info(mctx); + auto * ctx = new mtmd_helper_gen_audio(); + ctx->pipeline = make_pipeline(lctx, mctx); return ctx; } @@ -198,186 +402,31 @@ void mtmd_helper_gen_audio_free(mtmd_helper_gen_audio * ctx) { } void mtmd_helper_gen_audio_reset(mtmd_helper_gen_audio * ctx) { - ctx->pos = 0; - ctx->codes_buf.clear(); - ctx->c2w_state.clear(); - ctx->audio_pcm.clear(); - ctx->overlay.clear(); - ctx->overlay_idx = 0; - ctx->h_state_buf.clear(); - ctx->out_buf.clear(); + if (ctx->pipeline) { + ctx->pipeline->reset(); + } } int32_t mtmd_helper_gen_audio_set_input(mtmd_helper_gen_audio * ctx, const mtmd_helper_gen_audio_inp * inp) { - mtmd_helper_gen_audio_reset(ctx); - - if (ctx->info.type != MTMD_GEN_AUDIO_TYPE_QWEN3TTS) { + if (!ctx->pipeline) { LOG_ERR("mtmd_helper_gen_audio: unsupported or missing gen-audio pipeline\n"); return 1; } - if (!qwen3tts_ensure_cache(ctx)) { - return 1; - } - - const std::string lang = inp->lang ? inp->lang : "english"; - const llama_token c_lang = find_special_token(ctx->vocab, ("<|codec_language_" + lang + "|>").c_str()); - if (c_lang == LLAMA_TOKEN_NULL) { - LOG_ERR("mtmd_helper_gen_audio: unknown language '%s'\n", lang.c_str()); - return 1; - } - - std::vector speaker_embd; - if (inp->speaker_ref) { - if (!qwen3tts_encode_speaker(ctx, inp->speaker_ref, speaker_embd)) { - return 1; - } - } - - const int n_embd = ctx->n_embd; - auto row = [&](llama_token t) { - return std::vector(ctx->tok_embd.begin() + (size_t) t * n_embd, - ctx->tok_embd.begin() + (size_t) (t + 1) * n_embd); - }; - auto sum_row = [&](llama_token a, llama_token b) { - std::vector va = row(a), vb = row(b); - for (int i = 0; i < n_embd; i++) va[(size_t) i] += vb[(size_t) i]; - return va; - }; - auto sum_vec = [&](llama_token a, const std::vector & vb) { - std::vector va = row(a); - for (int i = 0; i < n_embd; i++) va[(size_t) i] += vb[(size_t) i]; - return va; - }; - - // upstream chat wrap, then slices: [0:3] role, [3:-5] utterance body - const std::string full = "<|im_start|>assistant\n" + std::string(inp->prompt, inp->prompt_len) + - "<|im_end|>\n<|im_start|>assistant\n"; - std::vector ids(full.size() + 16); - int n_ids = llama_tokenize(ctx->vocab, full.c_str(), (int32_t) full.size(), ids.data(), (int32_t) ids.size(), - false, true); - if (n_ids < 8) { - LOG_ERR("mtmd_helper_gen_audio: tokenization failed\n"); - return 1; - } - ids.resize((size_t) n_ids); - - std::vector> prompt; - for (int i = 0; i < 3; i++) prompt.push_back(row(ids[(size_t) i])); - prompt.push_back(sum_row(ctx->tts_pad, ctx->c_think)); - prompt.push_back(sum_row(ctx->tts_pad, ctx->c_think_b)); - prompt.push_back(sum_row(ctx->tts_pad, c_lang)); - prompt.push_back(sum_row(ctx->tts_pad, ctx->c_think_e)); - if (!speaker_embd.empty()) prompt.push_back(sum_vec(ctx->tts_pad, speaker_embd)); - prompt.push_back(sum_row(ctx->tts_bos, ctx->codec_pad)); - for (int i = 3; i < n_ids - 5; i++) prompt.push_back(sum_row(ids[(size_t) i], ctx->codec_pad)); - prompt.push_back(sum_row(ctx->tts_eos, ctx->codec_pad)); - prompt.push_back(sum_row(ctx->tts_pad, ctx->codec_bos)); - - const int n_prompt = (int) prompt.size(); - - // the talker rides the qwen3vl interleaved mrope: positions carry - // n_pos_per_embd sections laid out [section * n_tokens + i], all - // equal for a pure text/codec stream - ctx->mrope = llama_model_rope_type(ctx->model) == LLAMA_ROPE_TYPE_MROPE || - llama_model_rope_type(ctx->model) == LLAMA_ROPE_TYPE_IMROPE; - const int n_pos_per_embd = ctx->mrope ? 4 : 1; - - std::vector embd_buf((size_t) n_prompt * (size_t) n_embd); - for (int i = 0; i < n_prompt; i++) { - memcpy(embd_buf.data() + (size_t) i * n_embd, prompt[(size_t) i].data(), (size_t) n_embd * sizeof(float)); - } - - decode_embd_batch batch_embd(embd_buf.data(), n_prompt, n_pos_per_embd, n_embd); - if (ctx->mrope) batch_embd.set_position_mrope_1d(0, 0); - else batch_embd.set_position_normal(0, 0); - batch_embd.batch.logits[n_prompt - 1] = 1; - - if (llama_decode(ctx->lctx, batch_embd.batch) != 0) { - LOG_ERR("mtmd_helper_gen_audio: prefill decode failed\n"); - return 1; - } - - ctx->pos = n_prompt; - ctx->top_k = inp->top_k > 0 ? inp->top_k : 50; - ctx->top_p = inp->top_p > 0 ? inp->top_p : 1.0f; - ctx->out_type = inp->out_type; - - // the text stream keeps flowing during generation: the input after - // frame k adds trailing text row k on top of the codes embedding, - // then tts_eos, then tts_pad once the utterance is spent - for (int i = 3; i < n_ids - 5; i++) ctx->overlay.push_back(row(ids[(size_t) i])); - ctx->overlay.push_back(row(ctx->tts_eos)); - ctx->overlay.push_back(row(ctx->tts_pad)); - - return 0; + return ctx->pipeline->set_input(inp); } int32_t mtmd_helper_gen_audio_step(mtmd_helper_gen_audio * ctx, llama_token sampled, const float * h_state_in, const float ** h_state_out) { - if (ctx->info.type != MTMD_GEN_AUDIO_TYPE_QWEN3TTS) { + if (!ctx->pipeline) { return 1; } - - mtmd_gen_inp inp{}; - inp.type = MTMD_GEN_PROCESS_TYPE_GEN_CODE; - inp.code0 = sampled - ctx->codec_0; - inp.embd = const_cast(h_state_in); - inp.top_k = ctx->top_k; - inp.top_p = ctx->top_p; - mtmd_gen_out out{}; - if (mtmd_gen_audio_process(ctx->mctx, &inp, &out) != 0) { - LOG_ERR("mtmd_helper_gen_audio: gen_code process failed\n"); - return 1; - } - - ctx->codes_buf.insert(ctx->codes_buf.end(), out.codes, out.codes + out.n_codes); - if (out.n_codes > 0 && ctx->codes_buf.size() / out.n_codes >= ctx->window_frames) { - if (!qwen3tts_flush_c2w(ctx)) { - return 1; - } - } - - std::vector fb(out.embd, out.embd + ctx->n_embd); - const auto & ov = ctx->overlay[std::min(ctx->overlay_idx, ctx->overlay.size() - 1)]; - for (int i = 0; i < ctx->n_embd; i++) fb[(size_t) i] += ov[(size_t) i]; - ctx->overlay_idx++; - - const int n_pos_per_embd = ctx->mrope ? 4 : 1; - decode_embd_batch batch_embd(fb.data(), 1, n_pos_per_embd, ctx->n_embd); - if (ctx->mrope) batch_embd.set_position_mrope_1d(ctx->pos, 0); - else batch_embd.set_position_normal(ctx->pos, 0); - batch_embd.batch.logits[0] = 1; - ctx->pos++; - - if (llama_decode(ctx->lctx, batch_embd.batch) != 0) { - LOG_ERR("mtmd_helper_gen_audio: decode failed\n"); - return 1; - } - - const float * he = llama_get_embeddings_ith(ctx->lctx, -1); - ctx->h_state_buf.assign(he, he + ctx->n_embd); - *h_state_out = ctx->h_state_buf.data(); - - return 0; + return ctx->pipeline->step(sampled, h_state_in, h_state_out); } int32_t mtmd_helper_gen_audio_get_output(mtmd_helper_gen_audio * ctx, int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len) { - if (!qwen3tts_flush_c2w(ctx)) { + if (!ctx->pipeline) { return 1; } - - *out_sample_rate = ctx->info.sample_rate; - - if (ctx->out_type == MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM) { - *out_data = (const char *) ctx->audio_pcm.data(); - *out_data_len = ctx->audio_pcm.size() * sizeof(float); - return 0; - } - - ctx->out_buf.clear(); - write_wav16(ctx->out_buf, ctx->audio_pcm, ctx->info.sample_rate); - *out_data = ctx->out_buf.data(); - *out_data_len = ctx->out_buf.size(); - return 0; + return ctx->pipeline->get_output(out_sample_rate, out_data, out_data_len); }