]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
mtmd: support Qwen3-TTS (note: breaking change to llama-tts binary) (#26254)
authorXuan-Son Nguyen <redacted>
Tue, 4 Aug 2026 15:26:15 +0000 (17:26 +0200)
committerGitHub <redacted>
Tue, 4 Aug 2026 15:26:15 +0000 (17:26 +0200)
* convert text model

* main model load ok

* convert encoder ok

* speaker encoder loading ok

* speaker enc graph

* adapt vocab for backbone (with some tricks)

* add suppress_tokens

* poc new mtmd gen api

* convert code_predictor to gguf

* load gen_code model ok

* add clip_encode

* wire up

* code gen cgraph init version

Co-authored-by: Pascal <redacted>
* code2wav convert to gguf

* code2wav graph ok

* wire up in/out

* (wip) subgraph

* wire up

* wip, correct code2wav

* demo (to be removed)

* code2wav preserve kv between calls

* demo voice clone

* llama: add llama_model_get_tok_embd

* mtmd_helper_gen_audio API

* fix clamp cold prefix

Co-authored-by: Pascal <redacted>
* fuse snake op

Co-authored-by: Pascal <redacted>
* demo: use proper sampling

* update dev docs

* polymorphism helper

* revamp llama-tts binary

* update docs

* fix compile

* fix lint

* nits

* add guide + docs

* more timings info

* clean up code comments

* security fixes

* update docs

* use ggml_build_forward_select, clean up comments

* fix ci

* use ISO 639-1 language code

* rename CODE2WAV --> GEN_WAV, update docs

* clean up

* clean up tts.cpp

* add seq_id

* add step_prompt()

* mtmd_helper_model_can_chat

* clean up comments

---------

Co-authored-by: Pascal <redacted>
42 files changed:
common/arg.cpp
common/arg.h
common/common.h
conversion/__init__.py
conversion/qwen3tts.py [new file with mode: 0644]
docs/development/HOWTO-add-model.md
gguf-py/gguf/constants.py
gguf-py/gguf/gguf_writer.py
gguf-py/gguf/tensor_mapping.py
skills/code-review/SKILL.md
src/llama-arch.cpp
src/llama-arch.h
src/llama-ext.h
src/llama-model.cpp
src/models/models.h
src/models/qwen3tts.cpp [new file with mode: 0644]
src/models/qwen3vl.cpp
tests/test-llama-archs.cpp
tools/mtmd/CMakeLists.txt
tools/mtmd/README-dev.md
tools/mtmd/clip-graph.h
tools/mtmd/clip-impl.h
tools/mtmd/clip-model.h
tools/mtmd/clip.cpp
tools/mtmd/clip.h
tools/mtmd/models/models.h
tools/mtmd/models/qwen3tts-gen.cpp [new file with mode: 0644]
tools/mtmd/models/qwen3tts-spkenc.cpp [new file with mode: 0644]
tools/mtmd/mtmd-audio.cpp
tools/mtmd/mtmd-audio.h
tools/mtmd/mtmd-cli.cpp
tools/mtmd/mtmd-helper-common.h [new file with mode: 0644]
tools/mtmd/mtmd-helper-gen.cpp [new file with mode: 0644]
tools/mtmd/mtmd-helper.cpp
tools/mtmd/mtmd-helper.h
tools/mtmd/mtmd.cpp
tools/mtmd/mtmd.h
tools/tts/CMakeLists.txt
tools/tts/README.md
tools/tts/convert_pt_to_hf.py [deleted file]
tools/tts/tts-outetts.py [deleted file]
tools/tts/tts.cpp

index b75f4f05f0398412b80f22cc46f4d209b4f95935..86af0ba10a327283f2500f0bb8e48095df547017 100644 (file)
@@ -61,6 +61,7 @@ static std::initializer_list<enum llama_example> mmproj_examples = {
     LLAMA_EXAMPLE_MTMD,
     LLAMA_EXAMPLE_SERVER,
     LLAMA_EXAMPLE_CLI,
+    LLAMA_EXAMPLE_TTS,
 };
 
 static std::string read_file(const std::string & fname) {
@@ -360,7 +361,6 @@ static bool spec_types_is_default(const common_params & params) {
 common_models_handler common_models_handler_init(const common_params & params, llama_example curr_ex) {
     common_download_hf_plan plan;
     common_download_hf_plan plan_spec;
-    common_download_hf_plan plan_voc;
     common_download_opts opts;
 
     const bool spec_type_draft_mtp = std::find(params.speculative.types.begin(),
@@ -413,11 +413,7 @@ common_models_handler common_models_handler_init(const common_params & params, l
         plan_spec = common_download_get_hf_plan(params.speculative.draft.mparams, opts_spec);
     }
 
-    if (!params.vocoder.model.hf_repo.empty()) {
-        plan_voc = common_download_get_hf_plan(params.vocoder.model, opts);
-    }
-
-    return common_models_handler{plan, plan_spec, plan_voc, opts};
+    return common_models_handler{plan, plan_spec, opts};
 }
 
 bool common_models_handler_is_preset_repo(const common_models_handler & handler) {
@@ -467,7 +463,6 @@ void common_models_handler_apply(common_models_handler & handler, common_params
 
     auto & plan      = handler.plan;
     auto & plan_spec = handler.plan_spec;
-    auto & plan_voc  = handler.plan_voc;
 
     auto opts = handler.opts; // copy
     opts.callback = callback;
@@ -482,7 +477,6 @@ void common_models_handler_apply(common_models_handler & handler, common_params
     };
     handle_url(params.model);
     handle_url(params.mmproj);
-    handle_url(params.vocoder.model);
     handle_url(params.speculative.draft.mparams);
 
     // optionally, if docker repo is set, resolve it
@@ -510,14 +504,6 @@ void common_models_handler_apply(common_models_handler & handler, common_params
         task.opts       = opts;
         tasks.push_back(task);
     }
-    if (!params.vocoder.model.url.empty()) {
-        common_download_task task;
-        task.url        = params.vocoder.model.url;
-        task.local_path = params.vocoder.model.path;
-        task.opts       = opts;
-        tasks.push_back(task);
-    }
-
     bool had_spec_url = false;
     if (!params.speculative.draft.mparams.url.empty()) {
         common_download_task task;
@@ -631,11 +617,6 @@ void common_models_handler_apply(common_models_handler & handler, common_params
         had_spec_url = true;
     }
 
-    // handle vocoder plan (e.g. --hf-repo-v)
-    if (!plan_voc.model_files.empty()) {
-        add_tasks(plan_voc.model_files, plan_voc.primary, params.vocoder.model);
-    }
-
     if (!plan.model_files.empty()) {
         add_tasks(plan.model_files, plan.primary, params.model);
     }
@@ -1361,6 +1342,10 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
         params.n_parallel = -1;     // auto by default
     } else if (ex == LLAMA_EXAMPLE_TOKENIZE) {
         params.parse_special = true; // parse special tokens by default, like the old tokenize tool
+    } else if (ex == LLAMA_EXAMPLE_TTS) {
+        params.out_file = "output.wav";
+        params.sampling.penalty_repeat = 1.05f;
+        params.sampling.penalty_last_n = -1;
     }
 
     params.use_color = tty_can_use_colors();
@@ -2983,20 +2968,6 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
             params.model.hf_file = value;
         }
     ).set_examples({LLAMA_EXAMPLE_COMMON, LLAMA_EXAMPLE_DOWNLOAD, LLAMA_EXAMPLE_TOKENIZE}).set_env("LLAMA_ARG_HF_FILE"));
-    add_opt(common_arg(
-        {"-hfv", "-hfrv", "--hf-repo-v"}, "<user>/<model>[:quant]",
-        "Hugging Face model repository for the vocoder model (default: unused)",
-        [](common_params & params, const std::string & value) {
-            params.vocoder.model.hf_repo = value;
-        }
-    ).set_env("LLAMA_ARG_HF_REPO_V"));
-    add_opt(common_arg(
-        {"-hffv", "--hf-file-v"}, "FILE",
-        "Hugging Face model file for the vocoder model (default: unused)",
-        [](common_params & params, const std::string & value) {
-            params.vocoder.model.hf_file = value;
-        }
-    ).set_env("LLAMA_ARG_HF_FILE_V"));
     add_opt(common_arg(
         {"-hft", "--hf-token"}, "TOKEN",
         "Hugging Face access token (default: value from HF_TOKEN environment variable)",
@@ -4272,24 +4243,18 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
     //
 
     add_opt(common_arg(
-        {"-mv", "--model-vocoder"}, "FNAME",
-        "vocoder model for audio generation (default: unused)",
+        {"--tts-lang"}, "FNAME",
+        "language (ISO 639-1) for audio generation\n"
+        "see tts/README.md for per-model usage notes",
         [](common_params & params, const std::string & value) {
-            params.vocoder.model.path = value;
-        }
-    ).set_examples({LLAMA_EXAMPLE_TTS, LLAMA_EXAMPLE_SERVER}));
-     add_opt(common_arg(
-        {"--tts-use-guide-tokens"},
-        "Use guide tokens to improve TTS word recall",
-        [](common_params & params) {
-            params.vocoder.use_guide_tokens = true;
+            params.tts_lang = value;
         }
-    ).set_examples({LLAMA_EXAMPLE_TTS, LLAMA_EXAMPLE_SERVER}));
+    ).set_examples({LLAMA_EXAMPLE_TTS}));
     add_opt(common_arg(
         {"--tts-speaker-file"}, "FNAME",
         "speaker file path for audio generation",
         [](common_params & params, const std::string & value) {
-            params.vocoder.speaker_file = value;
+            params.tts_speaker_file = value;
         }
     ).set_examples({LLAMA_EXAMPLE_TTS}));
 
@@ -4409,16 +4374,6 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
     ).set_examples({LLAMA_EXAMPLE_DEBUG}));
 
     // presets
-    add_opt(common_arg(
-        {"--tts-oute-default"},
-        string_format("use default OuteTTS models (note: can download weights from the internet)"),
-        [](common_params & params) {
-            params.model.hf_repo = "OuteAI/OuteTTS-0.2-500M-GGUF";
-            params.model.hf_file = "OuteTTS-0.2-500M-Q8_0.gguf";
-            params.vocoder.model.hf_repo = "ggml-org/WavTokenizer";
-            params.vocoder.model.hf_file = "WavTokenizer-Large-75-F16.gguf";
-        }
-    ).set_examples({LLAMA_EXAMPLE_TTS}));
 
     add_opt(common_arg(
         {"--embd-gemma-default"},
index 8f609e356fe2f590d384d8a68f30b5c53e8ef4d0..44b9e887cfb907e4f5fb9190c4320e0e49a485a1 100644 (file)
@@ -137,7 +137,6 @@ void common_params_add_preset_options(std::vector<common_arg> & args);
 struct common_models_handler {
     common_download_hf_plan plan;
     common_download_hf_plan plan_spec;
-    common_download_hf_plan plan_voc;
     common_download_opts opts;
 };
 
index 78d0877566b609df3b89455da97f0a7f3542b6dc..3444aa157e9b73727ea2ca6107eb0dc9f9b36a74 100644 (file)
@@ -392,14 +392,6 @@ struct common_params_speculative {
     }
 };
 
-struct common_params_vocoder {
-    struct common_params_model model;
-
-    std::string speaker_file; // speaker file path
-
-    bool use_guide_tokens = false; // enable guide tokens to improve TTS accuracy
-};
-
 struct common_params_diffusion {
     int32_t steps         = 128;
     bool    visual_mode   = false;
@@ -497,7 +489,6 @@ struct common_params {
 
     struct common_params_sampling    sampling;
     struct common_params_speculative speculative;
-    struct common_params_vocoder     vocoder;
     struct common_params_diffusion   diffusion;
 
     struct common_params_model model;
@@ -740,6 +731,10 @@ struct common_params {
     void *                  load_progress_callback_user_data = NULL;
     bool no_alloc = false; // Don't allocate model buffers
 
+    // TTS params
+    std::string tts_lang = "";
+    std::string tts_speaker_file = "";
+
     bool is_gen_docs = false; // whether we are running inside llama-gen-docs
 };
 
index 534f9e309a138407189fc05d674fc46e0e78266b..06c2c50ad24571be56689aa83052b1ee17a8b299 100644 (file)
@@ -210,6 +210,7 @@ TEXT_MODEL_MAP: dict[str, str] = {
     "Qwen3MoeForCausalLM": "qwen",
     "Qwen3NextForCausalLM": "qwen",
     "Qwen3OmniMoeForConditionalGeneration": "qwen3vl",
+    "Qwen3TTSForConditionalGeneration": "qwen3tts",
     "Qwen3VLForConditionalGeneration": "qwen3vl",
     "Qwen3VLMoeForConditionalGeneration": "qwen3vl",
     "Qwen3_5ForCausalLM": "qwen",
@@ -304,6 +305,7 @@ MMPROJ_MODEL_MAP: dict[str, str] = {
     "Qwen2_5_VLForConditionalGeneration": "qwenvl",
     "Qwen3ASRForConditionalGeneration": "qwen3vl",
     "Qwen3OmniMoeForConditionalGeneration": "qwen3vl",
+    "Qwen3TTSForConditionalGeneration": "qwen3tts",
     "Qwen3VLForConditionalGeneration": "qwen3vl",
     "Qwen3VLMoeForConditionalGeneration": "qwen3vl",
     "Qwen3_5ForConditionalGeneration": "qwen3vl",
diff --git a/conversion/qwen3tts.py b/conversion/qwen3tts.py
new file mode 100644 (file)
index 0000000..d21a505
--- /dev/null
@@ -0,0 +1,471 @@
+from __future__ import annotations
+
+import json
+from pathlib import Path
+from typing import Any, Callable, Iterable, TYPE_CHECKING
+
+import torch
+import torch.nn.functional as F
+
+if TYPE_CHECKING:
+    from torch import Tensor
+
+from .base import ModelBase, MmprojModel, TextModel, gguf
+
+# Tricks being used to support this model via existing llama.cpp code paths:
+# - Text projection MLP is folded into the embedding table
+# - codec_embedding is concat to the text embedding table, vocab is extended
+#   example: codec_bos_id(2149) --> "<|codec_bos|>"
+#            codec_eos_token_id(2150) --> "<|codec_eos_token|>"
+#            codec_language_id.chinese(2055) --> "<|codec_language_chinese|>"
+#            other rows --> "<|codec_0|>", "<|codec_1|>", ..., "<|codec_1023|>"
+# - output tensor codec_head is smaller than vocab, so logits will be padded at inference time
+# - suppress_tokens is used to limit the backbone to only sample either semantic or EOS (stop) token
+
+# pipeline stage mapping:
+#   speaker reference encoder --> mapped to normal mtmd audio encoder
+#   backbone --> mapped to normal libllama text model (autoregressive)
+#   code_predictor --> MTMD_GEN_PROCESS_TYPE_GEN_CODE
+#   code2wav --> MTMD_GEN_PROCESS_TYPE_GEN_WAV
+
+# torch activation functions used by Qwen3TTSTalkerResizeMLP (config's hidden_act)
+_ACT2FN = {
+    "silu": F.silu,
+    "gelu": F.gelu,
+    "relu": F.relu,
+}
+
+
+@ModelBase.register("Qwen3TTSForConditionalGeneration")
+class Qwen3TTSTalkerModel(TextModel):
+    model_arch = gguf.MODEL_ARCH.QWEN3TTS
+
+    _TEXT_PROJ_KEYS = (
+        "model.text_embedding.weight",
+        "text_projection.linear_fc1.weight",
+        "text_projection.linear_fc1.bias",
+        "text_projection.linear_fc2.weight",
+        "text_projection.linear_fc2.bias",
+    )
+
+    _text_proj_buffer: dict[str, Tensor]
+    _folded_text_embed: Tensor | None
+    _codec_embed: Tensor | None
+
+    def __init__(self, dir_model: Path, *args, **kwargs):
+        hparams = kwargs.pop("hparams", None)
+        if hparams is None:
+            hparams = ModelBase.load_hparams(dir_model, is_mistral_format=False)
+        raw_talker_config = dict(hparams["talker_config"])
+        self._talker_config = raw_talker_config
+        self.n_codec_vocab = raw_talker_config["vocab_size"]
+        talker_config = dict(raw_talker_config)
+        talker_config["vocab_size"] = talker_config["text_vocab_size"]
+        hparams["text_config"] = talker_config
+        super().__init__(dir_model, *args, hparams=hparams, **kwargs)
+        self._text_proj_buffer = {}
+        self._folded_text_embed = None
+        self._codec_embed = None
+
+    def _codec_token_names(self) -> list[str]:
+        # start every row with a generic name, then override the ones with a
+        # known meaning (bos/eos/language/etc, derived from the *_id fields
+        # of talker_config) with a more descriptive one
+        names = [f"<|codec_{i}|>" for i in range(self.n_codec_vocab)]
+        for key, val in self._talker_config.items():
+            if not key.endswith("_id"):
+                continue
+            prefix = key[:-len("_id")]
+            if isinstance(val, int):
+                names[val] = f"<|{prefix}|>"
+            elif isinstance(val, dict):
+                for subkey, subval in val.items():
+                    names[subval] = f"<|{prefix}_{subkey}|>"
+        return names
+
+    def set_vocab(self):
+        codec_tokens = self._codec_token_names()
+        codec_toktypes = [gguf.TokenType.CONTROL] * len(codec_tokens)
+
+        try:
+            tokens, scores, toktypes = self._create_vocab_sentencepiece()
+            self.gguf_writer.add_tokenizer_model("llama")
+            self.gguf_writer.add_tokenizer_pre("default")
+            tokens += [t.encode("utf-8") for t in codec_tokens]
+            scores += [0.0] * len(codec_tokens)
+            toktypes += codec_toktypes
+            self.gguf_writer.add_token_list(tokens)
+            self.gguf_writer.add_token_scores(scores)
+            self.gguf_writer.add_token_types(toktypes)
+            special_vocab = gguf.SpecialVocab(self.dir_model, n_vocab=len(tokens))
+            special_vocab.add_to_gguf(self.gguf_writer)
+            return
+        except FileNotFoundError:
+            pass
+
+        tokens, toktypes, tokpre = self.get_vocab_base()
+        tokens += codec_tokens
+        toktypes += codec_toktypes
+        self.gguf_writer.add_tokenizer_model("gpt2")
+        self.gguf_writer.add_tokenizer_pre(tokpre)
+        self.gguf_writer.add_token_list(tokens)
+        self.gguf_writer.add_token_types(toktypes)
+
+        special_vocab = gguf.SpecialVocab(self.dir_model, load_merges=True)
+        special_vocab.add_to_gguf(self.gguf_writer)
+
+        # make sure that the model has no chat template, so chat will be disabled
+        self.gguf_writer.add_chat_template(None)
+
+    def set_gguf_parameters(self):
+        super().set_gguf_parameters()
+
+        # note: final vocab layout is [text_vocab | codec_vocab], with text_vocab is actually padded with -inf in cgraph
+        # for codec_vocab, only first 2048 rows can be sampled for semantic code
+        # plus codec_eos_token_id that used for signaling end of generation
+        # ref: https://github.com/QwenLM/Qwen3-TTS/blob/022e286b98fbec7e1e916cb940cdf532cd9f488e/qwen_tts/core/models/modeling_qwen3_tts.py#L2059-L2063
+
+        vocab_size = self.hparams["vocab_size"] + self.n_codec_vocab
+        codec_eos_token_id = self.hparams["vocab_size"] + self._talker_config["codec_eos_token_id"]
+        self.gguf_writer.add_suppress_tokens([
+            i for i in range(vocab_size - 1024, vocab_size)
+            if i != codec_eos_token_id
+        ])
+        self.gguf_writer.add_eos_token_id(codec_eos_token_id)
+        self.gguf_writer.add_add_eos_token(False)
+
+    @classmethod
+    def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
+        name, gen = item
+
+        if not name.startswith("talker.") or name.startswith("talker.code_predictor."):
+            return None
+
+        name = name[len("talker."):]
+        return super().filter_tensors((name, gen))
+
+    def _maybe_emit_token_embd(self) -> Iterable[tuple[str, Tensor]]:
+        if self._folded_text_embed is None or self._codec_embed is None:
+            return
+        combined = torch.cat([self._folded_text_embed, self._codec_embed], dim=0)
+        yield (self.format_tensor_name(gguf.MODEL_TENSOR.TOKEN_EMBD), combined)
+
+    def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
+        # codec_embedding rows are appended after the text vocab, extending the embedding table
+        if name == "model.codec_embedding.weight":
+            self._codec_embed = data_torch
+            yield from self._maybe_emit_token_embd()
+            return
+
+        # codec_head is the output head for the (smaller) codec vocab; logits get padded to
+        # the extended vocab size at inference time
+        if name == "codec_head.weight":
+            yield (self.format_tensor_name(gguf.MODEL_TENSOR.OUTPUT), data_torch)
+            return
+
+        if name in self._TEXT_PROJ_KEYS:
+            self._text_proj_buffer[name] = data_torch
+            if len(self._text_proj_buffer) < len(self._TEXT_PROJ_KEYS):
+                return
+
+            # fold MLP into the embedding table at conversion time, MLP won't be used at inference time anyway
+            act_fn = _ACT2FN[self.hparams["hidden_act"]]
+            embed = self._text_proj_buffer["model.text_embedding.weight"]
+            hidden = act_fn(F.linear(embed,
+                                     self._text_proj_buffer["text_projection.linear_fc1.weight"],
+                                     self._text_proj_buffer["text_projection.linear_fc1.bias"]))
+            folded = F.linear(hidden,
+                              self._text_proj_buffer["text_projection.linear_fc2.weight"],
+                              self._text_proj_buffer["text_projection.linear_fc2.bias"])
+            self._folded_text_embed = folded
+            yield from self._maybe_emit_token_embd()
+            return
+
+        yield from super().modify_tensors(data_torch, name, bid)
+
+
+@ModelBase.register("Qwen3TTSForConditionalGeneration")
+class Qwen3TTSSpeakerEncoderModel(MmprojModel):
+    has_vision_encoder = False
+    has_audio_encoder = True
+
+    # talker.code_predictor.model.layers.{bid}.<key> -> A_GEN_CODE_*
+    # bypass tensor_mapping.py for now to make it simple
+    _CODE_LAYER_TENSOR_MAP = {
+        "input_layernorm": gguf.MODEL_TENSOR.A_GEN_CODE_ATTN_NORM,
+        "self_attn.q_proj": gguf.MODEL_TENSOR.A_GEN_CODE_ATTN_Q,
+        "self_attn.q_norm": gguf.MODEL_TENSOR.A_GEN_CODE_ATTN_Q_NORM,
+        "self_attn.k_proj": gguf.MODEL_TENSOR.A_GEN_CODE_ATTN_K,
+        "self_attn.k_norm": gguf.MODEL_TENSOR.A_GEN_CODE_ATTN_K_NORM,
+        "self_attn.v_proj": gguf.MODEL_TENSOR.A_GEN_CODE_ATTN_V,
+        "self_attn.o_proj": gguf.MODEL_TENSOR.A_GEN_CODE_ATTN_OUT,
+        "post_attention_layernorm": gguf.MODEL_TENSOR.A_GEN_CODE_FFN_NORM,
+        "mlp.gate_proj": gguf.MODEL_TENSOR.A_GEN_CODE_FFN_GATE,
+        "mlp.up_proj": gguf.MODEL_TENSOR.A_GEN_CODE_FFN_UP,
+        "mlp.down_proj": gguf.MODEL_TENSOR.A_GEN_CODE_FFN_DOWN,
+    }
+
+    # note: codebook pages will be stacked to 3D
+    _CODE_GEN_N_CODEBOOKS = 15
+    _code_embed_buffer: dict[int, Tensor] = {}
+    _code_head_buffer: dict[int, Tensor] = {}
+    _wav_config_cache: dict[str, Any] | None = None
+
+    def __init__(self, dir_model: Path, *args, **kwargs):
+        hparams = kwargs.pop("hparams", None)
+        if hparams is None:
+            hparams = ModelBase.load_hparams(dir_model, is_mistral_format=False)
+        hparams["text_config"] = {"hidden_size": hparams["talker_config"]["hidden_size"]}
+        # ECAPA-TDNN has a fixed 4-stage backbone, but MmprojModel.__init__ needs a n_block_keys
+        hparams["speaker_encoder_config"]["n_layers"] = 4
+        super().__init__(dir_model, *args, hparams=hparams, **kwargs)
+        self._wav_config_cache = None
+
+    def get_audio_config(self) -> dict[str, Any] | None:
+        return self.global_config.get("speaker_encoder_config")
+
+    def set_gguf_parameters(self):
+        self.gguf_writer.add_file_type(self.ftype)
+        self.gguf_writer.add_clip_has_audio_encoder(True)
+        self.gguf_writer.add_clip_audio_projector_type(gguf.VisionProjectorType.QWEN3TTS_SPKENC)
+
+        # handle speaker encoder config
+        self.gguf_writer.add_audio_projection_dim(self.n_embd_text)
+        # mel_spectrogram() front-end: sr=24000, n_fft=1024, hop=256, n_mels=128, fmin=0, fmax=12000 (=sr/2, the clip.cpp default)
+        self.gguf_writer.add_audio_num_mel_bins(128)
+        # 3 SE-Res2Net stages; the stem conv, mfa, asp and fc are not counted here
+        self.gguf_writer.add_audio_block_count(3)
+        # ECAPA-TDNN has no attention/FFN, these are dummy to allow clip.cpp to load it
+        self.gguf_writer.add_audio_embedding_length(1536)
+        self.gguf_writer.add_audio_head_count(1)
+        self.gguf_writer.add_audio_feed_forward_length(1536)
+        self.gguf_writer.add_audio_attention_layernorm_eps(1e-5)
+
+        # handle code predictor config
+        self.gguf_writer.add_clip_has_gen_audio_encoder(True)
+        self.gguf_writer.add_clip_gen_audio_projector_type(gguf.VisionProjectorType.QWEN3TTS_GEN)
+        code_predictor_config = self.global_config["talker_config"]["code_predictor_config"]
+        self.gguf_writer.add_gen_audio_projection_dim(self.n_embd_text)
+        self.gguf_writer.add_gen_audio_embedding_length(code_predictor_config["hidden_size"])
+        self.gguf_writer.add_gen_audio_feed_forward_length(code_predictor_config["intermediate_size"])
+        self.gguf_writer.add_gen_audio_block_count(code_predictor_config["num_hidden_layers"])
+        self.gguf_writer.add_gen_audio_head_count(code_predictor_config["num_attention_heads"])
+        self.gguf_writer.add_gen_audio_head_count_kv(code_predictor_config["num_key_value_heads"])
+        self.gguf_writer.add_gen_audio_attention_layernorm_eps(code_predictor_config["rms_norm_eps"])
+        # note: code2wav hparams are hardcoded on the mtmd/clip.cpp side for now, not written here
+
+    def _wav_decoder_config(self) -> dict[str, Any] | None:
+        # code2wav has its own config.json, inside the speech_tokenizer dir
+        if self._wav_config_cache is None:
+            path = self.dir_model / "speech_tokenizer" / "config.json"
+            with open(path, "r", encoding="utf-8") as f:
+                cfg = json.load(f)
+            self._wav_config_cache = cfg["decoder_config"]
+        return self._wav_config_cache
+
+    def tensor_force_quant(self, name, new_name, bid, n_dims):
+        # conv1d/conv1d_dw kernels must be F16, ggml_conv_1d(_dw) has no BF16 path
+        if new_name.endswith(".weight") and (
+            new_name in ("a.gen.wav.pre_conv.weight", "a.gen.wav.dac.entry.weight", "a.gen.wav.dac.post_conv.weight")
+            or (".up.blk." in new_name and new_name.endswith(".dwconv.weight"))
+            or (".dac.blk." in new_name and (new_name.endswith(".conv1.weight") or new_name.endswith(".conv2.weight")))
+        ):
+            return gguf.GGMLQuantizationType.F16
+        # ConvTranspose1d kernels: only F16/F32 are implemented, no BF16
+        if new_name.endswith(".conv.weight") and (".up.blk." in new_name or ".dac.blk." in new_name):
+            return gguf.GGMLQuantizationType.F32
+        return super().tensor_force_quant(name, new_name, bid, n_dims)
+
+    @classmethod
+    def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
+        name, gen = item
+
+        if not (
+            name.startswith("speaker_encoder.")
+            or name.startswith("talker.code_predictor.")
+            or name == "talker.model.codec_embedding.weight"
+        ):
+            return None
+
+        return super().filter_tensors((name, gen))
+
+    def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
+        # code2wav tensors are already named by generate_extra_tensors(), pass them through
+        if name.startswith("a.gen.wav."):
+            yield (name, data_torch)
+            return
+
+        # codebook-0 embedding, fed back to the talker backbone (codebooks 1-15 live in code_predictor)
+        if name == "talker.model.codec_embedding.weight":
+            yield (self.format_tensor_name(gguf.MODEL_TENSOR.A_GEN_CODE_OUT_EMBD), data_torch)
+            return
+
+        if name == "talker.code_predictor.model.norm.weight":
+            yield (self.format_tensor_name(gguf.MODEL_TENSOR.A_GEN_CODE_OUTPUT_NORM), data_torch)
+            return
+
+        if name.startswith("talker.code_predictor.small_to_mtp_projection."):
+            suffix = "." + name.rsplit(".", 1)[1]
+            yield (self.format_tensor_name(gguf.MODEL_TENSOR.A_GEN_CODE_PROJ_IN, suffix=suffix), data_torch)
+            return
+
+        if name.startswith("talker.code_predictor.model.codec_embedding."):
+            idx = int(name.split("codec_embedding.")[1].split(".")[0])
+            self._code_embed_buffer[idx] = data_torch
+            if len(self._code_embed_buffer) < self._CODE_GEN_N_CODEBOOKS:
+                return
+            stacked = torch.stack([self._code_embed_buffer.pop(i) for i in range(self._CODE_GEN_N_CODEBOOKS)], dim=0)
+            yield (self.format_tensor_name(gguf.MODEL_TENSOR.A_GEN_CODE_EMBD), stacked)
+            return
+
+        if name.startswith("talker.code_predictor.lm_head."):
+            idx = int(name.split("lm_head.")[1].split(".")[0])
+            self._code_head_buffer[idx] = data_torch
+            if len(self._code_head_buffer) < self._CODE_GEN_N_CODEBOOKS:
+                return
+            stacked = torch.stack([self._code_head_buffer.pop(i) for i in range(self._CODE_GEN_N_CODEBOOKS)], dim=0)
+            yield (self.format_tensor_name(gguf.MODEL_TENSOR.A_GEN_CODE_HEAD), stacked)
+            return
+
+        if name.startswith("talker.code_predictor.model.layers."):
+            rest = name.split("model.layers.")[1]        # "{bid}.<key>.weight"
+            _, key_with_suffix = rest.split(".", 1)       # "<key>.weight"
+            key = key_with_suffix.rsplit(".", 1)[0]        # "<key>"
+            tensor = self._CODE_LAYER_TENSOR_MAP.get(key)
+            if tensor is not None:
+                yield (self.format_tensor_name(tensor, bid), data_torch)
+                return
+
+        if "res2net_block.blocks." in name:
+            assert bid is not None  # the outer stage index, picked up from the tensor name automatically
+            xid = int(name.split("res2net_block.blocks.")[1].split(".")[0])
+            suffix = "." + name.rsplit(".", 1)[1]
+            new_name = gguf.TENSOR_NAMES[gguf.MODEL_TENSOR.A_ENC_CONV_RES2].format(bid=bid, xid=xid) + suffix
+            yield (new_name, data_torch)
+            return
+
+        yield from super().modify_tensors(data_torch, name, bid)
+
+    def generate_extra_tensors(self) -> Iterable[tuple[str, Tensor]]:
+        yield from self._generate_code2wav_tensors()
+
+    def _generate_code2wav_tensors(self) -> Iterable[tuple[str, Tensor]]:
+        # code2wav weights live in speech_tokenizer/model.safetensors, not the main safetensors
+        from safetensors.torch import load_file
+
+        wav_config = self._wav_decoder_config()
+        state_dict = load_file(self.dir_model / "speech_tokenizer" / "model.safetensors")
+
+        def get(name: str) -> Tensor:
+            return state_dict[name]
+
+        def snake_fold(alpha: Tensor, beta: Tensor) -> tuple[Tensor, Tensor]:
+            # fold SnakeBeta's exp()/reciprocal here, so the graph is only mul/sin/sqr/mul/add
+            return torch.exp(alpha), 1.0 / (torch.exp(beta) + 1e-9)
+
+        def rvq_codebook(prefix: str, n_layers: int) -> Tensor:
+            # checkpoint has EMA accumulators, so codebook[i] = embedding_sum[i] / cluster_usage[i]
+            books = []
+            for i in range(n_layers):
+                embedding_sum = get(f"{prefix}.vq.layers.{i}._codebook.embedding_sum")
+                cluster_usage = get(f"{prefix}.vq.layers.{i}._codebook.cluster_usage")
+                books.append(embedding_sum / cluster_usage.clamp_min(1e-5).unsqueeze(-1))
+            return torch.stack(books, dim=0) if n_layers > 1 else books[0]
+
+        T = gguf.MODEL_TENSOR
+
+        # --- quantizer: RVQ codebook decode ---
+        yield (self.format_tensor_name(T.A_GEN_WAV_QUANT_FIRST_IN), get("decoder.quantizer.rvq_first.input_proj.weight").squeeze(-1))
+        yield (self.format_tensor_name(T.A_GEN_WAV_QUANT_FIRST_OUT), get("decoder.quantizer.rvq_first.output_proj.weight").squeeze(-1))
+        yield (self.format_tensor_name(T.A_GEN_WAV_QUANT_FIRST_CB), rvq_codebook("decoder.quantizer.rvq_first", 1))
+        yield (self.format_tensor_name(T.A_GEN_WAV_QUANT_REST_IN), get("decoder.quantizer.rvq_rest.input_proj.weight").squeeze(-1))
+        yield (self.format_tensor_name(T.A_GEN_WAV_QUANT_REST_OUT), get("decoder.quantizer.rvq_rest.output_proj.weight").squeeze(-1))
+        yield (self.format_tensor_name(T.A_GEN_WAV_QUANT_REST_CB), rvq_codebook("decoder.quantizer.rvq_rest", self._CODE_GEN_N_CODEBOOKS))
+
+        # --- pre_conv ---
+        yield (self.format_tensor_name(T.A_GEN_WAV_PRE_CONV, suffix=".weight"), get("decoder.pre_conv.conv.weight"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_PRE_CONV, suffix=".bias"), get("decoder.pre_conv.conv.bias"))
+
+        # --- pre_transformer ---
+        yield (self.format_tensor_name(T.A_GEN_WAV_TFM_IN_PROJ, suffix=".weight"), get("decoder.pre_transformer.input_proj.weight"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_TFM_IN_PROJ, suffix=".bias"), get("decoder.pre_transformer.input_proj.bias"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_TFM_OUT_PROJ, suffix=".weight"), get("decoder.pre_transformer.output_proj.weight"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_TFM_OUT_PROJ, suffix=".bias"), get("decoder.pre_transformer.output_proj.bias"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_TFM_OUTPUT_NORM), get("decoder.pre_transformer.norm.weight"))
+
+        tfm_layer_map = {
+            "input_layernorm.weight":         T.A_GEN_WAV_TFM_ATTN_NORM,
+            "self_attn.q_proj.weight":        T.A_GEN_WAV_TFM_ATTN_Q,
+            "self_attn.k_proj.weight":        T.A_GEN_WAV_TFM_ATTN_K,
+            "self_attn.v_proj.weight":        T.A_GEN_WAV_TFM_ATTN_V,
+            "self_attn.o_proj.weight":        T.A_GEN_WAV_TFM_ATTN_OUT,
+            "self_attn_layer_scale.scale":    T.A_GEN_WAV_TFM_ATTN_SCALE,
+            "post_attention_layernorm.weight": T.A_GEN_WAV_TFM_FFN_NORM,
+            "mlp.gate_proj.weight":           T.A_GEN_WAV_TFM_FFN_GATE,
+            "mlp.up_proj.weight":             T.A_GEN_WAV_TFM_FFN_UP,
+            "mlp.down_proj.weight":           T.A_GEN_WAV_TFM_FFN_DOWN,
+            "mlp_layer_scale.scale":          T.A_GEN_WAV_TFM_FFN_SCALE,
+        }
+        assert wav_config is not None
+        for bid in range(wav_config["num_hidden_layers"]):
+            for key, tensor_id in tfm_layer_map.items():
+                yield (self.format_tensor_name(tensor_id, bid), get(f"decoder.pre_transformer.layers.{bid}.{key}"))
+
+        # --- upsample: 2x (causal ConvTranspose1d + ConvNeXt block) ---
+        up_map = {
+            "0.conv.weight":     (T.A_GEN_WAV_UP_CONV, ".weight"),
+            "0.conv.bias":       (T.A_GEN_WAV_UP_CONV, ".bias"),
+            "1.dwconv.conv.weight": (T.A_GEN_WAV_UP_DWCONV, ".weight"),
+            "1.dwconv.conv.bias":   (T.A_GEN_WAV_UP_DWCONV, ".bias"),
+            "1.norm.weight":     (T.A_GEN_WAV_UP_NORM, ".weight"),
+            "1.norm.bias":       (T.A_GEN_WAV_UP_NORM, ".bias"),
+            "1.pwconv1.weight":  (T.A_GEN_WAV_UP_PW1, ".weight"),
+            "1.pwconv1.bias":    (T.A_GEN_WAV_UP_PW1, ".bias"),
+            "1.pwconv2.weight":  (T.A_GEN_WAV_UP_PW2, ".weight"),
+            "1.pwconv2.bias":    (T.A_GEN_WAV_UP_PW2, ".bias"),
+            "1.gamma":           (T.A_GEN_WAV_UP_GAMMA, ""),
+        }
+        for bid in range(len(wav_config["upsampling_ratios"])):
+            for key, (tensor_id, suffix) in up_map.items():
+                yield (self.format_tensor_name(tensor_id, bid, suffix=suffix), get(f"decoder.upsample.{bid}.{key}"))
+
+        # --- DAC decoder ---
+        yield (self.format_tensor_name(T.A_GEN_WAV_DAC_ENTRY, suffix=".weight"), get("decoder.decoder.0.conv.weight"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_DAC_ENTRY, suffix=".bias"), get("decoder.decoder.0.conv.bias"))
+
+        n_dac_blocks = len(wav_config["upsample_rates"])
+        for bid in range(n_dac_blocks):
+            py = bid + 1  # decoder.decoder.0 is the entry conv, blocks start at 1
+
+            a, b = snake_fold(get(f"decoder.decoder.{py}.block.0.alpha"), get(f"decoder.decoder.{py}.block.0.beta"))
+            yield (self.format_tensor_name(T.A_GEN_WAV_DAC_UP_SNAKE, bid, suffix=".alpha"), a)
+            yield (self.format_tensor_name(T.A_GEN_WAV_DAC_UP_SNAKE, bid, suffix=".beta"), b)
+            yield (self.format_tensor_name(T.A_GEN_WAV_DAC_UP_CONV, bid, suffix=".weight"), get(f"decoder.decoder.{py}.block.1.conv.weight"))
+            yield (self.format_tensor_name(T.A_GEN_WAV_DAC_UP_CONV, bid, suffix=".bias"), get(f"decoder.decoder.{py}.block.1.conv.bias"))
+
+            for xid in range(3):
+                ridx = xid + 2  # block.2/3/4 are the 3 residual units
+
+                a1, b1 = snake_fold(get(f"decoder.decoder.{py}.block.{ridx}.act1.alpha"), get(f"decoder.decoder.{py}.block.{ridx}.act1.beta"))
+                name1 = gguf.TENSOR_NAMES[T.A_GEN_WAV_DAC_RES_ACT1].format(bid=bid, xid=xid)
+                yield (name1 + ".alpha", a1)
+                yield (name1 + ".beta", b1)
+
+                name_conv1 = gguf.TENSOR_NAMES[T.A_GEN_WAV_DAC_RES_CONV1].format(bid=bid, xid=xid)
+                yield (name_conv1 + ".weight", get(f"decoder.decoder.{py}.block.{ridx}.conv1.conv.weight"))
+                yield (name_conv1 + ".bias", get(f"decoder.decoder.{py}.block.{ridx}.conv1.conv.bias"))
+
+                a2, b2 = snake_fold(get(f"decoder.decoder.{py}.block.{ridx}.act2.alpha"), get(f"decoder.decoder.{py}.block.{ridx}.act2.beta"))
+                name2 = gguf.TENSOR_NAMES[T.A_GEN_WAV_DAC_RES_ACT2].format(bid=bid, xid=xid)
+                yield (name2 + ".alpha", a2)
+                yield (name2 + ".beta", b2)
+
+                name_conv2 = gguf.TENSOR_NAMES[T.A_GEN_WAV_DAC_RES_CONV2].format(bid=bid, xid=xid)
+                yield (name_conv2 + ".weight", get(f"decoder.decoder.{py}.block.{ridx}.conv2.conv.weight"))
+                yield (name_conv2 + ".bias", get(f"decoder.decoder.{py}.block.{ridx}.conv2.conv.bias"))
+
+        a5, b5 = snake_fold(get("decoder.decoder.5.alpha"), get("decoder.decoder.5.beta"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_DAC_POST_SNAKE, suffix=".alpha"), a5)
+        yield (self.format_tensor_name(T.A_GEN_WAV_DAC_POST_SNAKE, suffix=".beta"), b5)
+        yield (self.format_tensor_name(T.A_GEN_WAV_DAC_POST_CONV, suffix=".weight"), get("decoder.decoder.6.conv.weight"))
+        yield (self.format_tensor_name(T.A_GEN_WAV_DAC_POST_CONV, suffix=".bias"), get("decoder.decoder.6.conv.bias"))
index 102f479eb02c2b7d2b2167f2719c9cc6f6dacc9c..270e6b735651eb3fd97f4f34473de12d5aa47fd4 100644 (file)
@@ -133,6 +133,7 @@ Note:
 - To debug the multimodal preprocessor and encoder, you can use [llama-mtmd-debug](tools/mtmd/debug/mtmd-debug.cpp).
 - Adding a model-specific API or CLI is an anti-pattern in `libmtmd`. The goal of `libmtmd` is to provide an easy-to-use, model-agnostic library for multimodal pipeline.
 - In most cases, `llama-mtmd-cli` should not be modified. If a model requires a specific prompt, either let the user provide it or bake it into the Jinja chat template.
+- For audio generation models, see `tools/mtmd/README-dev.md`
 
 ## Tips and tricks
 
index 6b0a26b63d89de68058fd0da676905e5d27dc029..8516222cccbbaa22430d846658ae2a4507017c15 100644 (file)
@@ -323,6 +323,7 @@ class Keys:
         PROJECTOR_TYPE        = "clip.projector_type"
         HAS_VISION_ENCODER    = "clip.has_vision_encoder"
         HAS_AUDIO_ENCODER     = "clip.has_audio_encoder"
+        HAS_GEN_AUDIO_ENCODER = "clip.has_gen_audio_encoder"
         HAS_LLAVA_PROJECTOR   = "clip.has_llava_projector"
 
     class ClipVision:
@@ -397,6 +398,18 @@ class Keys:
             DOWNSAMPLE_RATE = "clip.audio.projector.downsample_rate"
             HEAD_COUNT      = "clip.audio.projector.head_count"
 
+    class ClipGenAudio:
+        PROJECTOR_TYPE      = "clip.gen.audio.projector_type" # for mixed modality models
+        EMBEDDING_LENGTH    = "clip.gen.audio.embedding_length"
+        FEED_FORWARD_LENGTH = "clip.gen.audio.feed_forward_length"
+        BLOCK_COUNT         = "clip.gen.audio.block_count"
+        PROJECTION_DIM      = "clip.gen.audio.projection_dim"
+
+        class Attention:
+            HEAD_COUNT      = "clip.gen.audio.attention.head_count"
+            HEAD_COUNT_KV   = "clip.gen.audio.attention.head_count_kv"
+            LAYERNORM_EPS   = "clip.gen.audio.attention.layer_norm_epsilon"
+
     class Diffusion:
         SHIFT_LOGITS        = "diffusion.shift_logits"
 
@@ -558,6 +571,7 @@ class MODEL_ARCH(IntEnum):
     TALKIE           = auto()
     MELLUM           = auto()
     NANBEIGE         = auto()
+    QWEN3TTS         = auto()
 
 
 class VISION_PROJECTOR_TYPE(IntEnum):
@@ -958,6 +972,65 @@ class MODEL_TENSOR(IntEnum):
     A_ENC_DOWNSAMPLE_CONV = auto() # mimo-audio-tokenizer: post-transformer downsample conv
     A_ENC_DOWNSAMPLE_NORM = auto() # mimo-audio-tokenizer: post-transformer downsample norm
     A_ENC_RVQ_CODEBOOK    = auto() # mimo-audio-tokenizer: residual vector quantizer codebook, per quantizer index
+    A_ENC_CONV_RES2       = auto() # qwen3tts
+    A_ENC_SE_CONV1        = auto() # qwen3tts
+    A_ENC_SE_CONV2        = auto() # qwen3tts
+    A_ENC_ASP_ATTN        = auto() # qwen3tts
+    A_ENC_ASP_TDNN        = auto() # qwen3tts
+    # qwen3tts code_predictor: predicts the remaining RVQ codebooks
+    A_GEN_CODE_PROJ_IN     = auto() # small_to_mtp_projection
+    A_GEN_CODE_EMBD        = auto() # per-codebook embedding table, merged 3D [n_codebooks, vocab, dim]
+    A_GEN_CODE_HEAD        = auto() # per-codebook output head, merged 3D [n_codebooks, vocab, dim]
+    A_GEN_CODE_OUT_EMBD    = auto() # codebook-0 embedding, re-fed into the talker backbone (talker.model.codec_embedding)
+    A_GEN_CODE_ATTN_NORM   = auto()
+    A_GEN_CODE_ATTN_Q      = auto()
+    A_GEN_CODE_ATTN_Q_NORM = auto()
+    A_GEN_CODE_ATTN_K      = auto()
+    A_GEN_CODE_ATTN_K_NORM = auto()
+    A_GEN_CODE_ATTN_V      = auto()
+    A_GEN_CODE_ATTN_OUT    = auto()
+    A_GEN_CODE_FFN_NORM    = auto()
+    A_GEN_CODE_FFN_GATE    = auto()
+    A_GEN_CODE_FFN_UP      = auto()
+    A_GEN_CODE_FFN_DOWN    = auto()
+    A_GEN_CODE_OUTPUT_NORM = auto()
+    # qwen3tts code2wav: RVQ codes -> raw PCM
+    A_GEN_WAV_QUANT_FIRST_IN       = auto() # semantic RVQ, in_proj (1x1 conv, loaded as 2D)
+    A_GEN_WAV_QUANT_FIRST_OUT      = auto() # semantic RVQ, out_proj
+    A_GEN_WAV_QUANT_FIRST_CB = auto() # semantic RVQ codebook (1 layer), folded from embedding_sum/cluster_usage
+    A_GEN_WAV_QUANT_REST_IN        = auto() # acoustic RVQ, in_proj
+    A_GEN_WAV_QUANT_REST_OUT       = auto() # acoustic RVQ, out_proj
+    A_GEN_WAV_QUANT_REST_CB  = auto() # acoustic RVQ codebooks, merged 3D [15, vocab, dim]
+    A_GEN_WAV_PRE_CONV             = auto()
+    A_GEN_WAV_TFM_IN_PROJ          = auto()
+    A_GEN_WAV_TFM_OUT_PROJ         = auto()
+    A_GEN_WAV_TFM_OUTPUT_NORM      = auto()
+    A_GEN_WAV_TFM_ATTN_NORM        = auto()
+    A_GEN_WAV_TFM_ATTN_Q           = auto()
+    A_GEN_WAV_TFM_ATTN_K           = auto()
+    A_GEN_WAV_TFM_ATTN_V           = auto()
+    A_GEN_WAV_TFM_ATTN_OUT         = auto()
+    A_GEN_WAV_TFM_ATTN_SCALE       = auto() # layer scale (gamma) on the attn output
+    A_GEN_WAV_TFM_FFN_NORM         = auto()
+    A_GEN_WAV_TFM_FFN_GATE         = auto()
+    A_GEN_WAV_TFM_FFN_UP           = auto()
+    A_GEN_WAV_TFM_FFN_DOWN         = auto()
+    A_GEN_WAV_TFM_FFN_SCALE        = auto() # layer scale (gamma) on the FFN output
+    A_GEN_WAV_UP_CONV              = auto() # causal ConvTranspose1d, 2x upsample
+    A_GEN_WAV_UP_DWCONV            = auto() # ConvNeXt depthwise conv
+    A_GEN_WAV_UP_NORM              = auto() # ConvNeXt LayerNorm
+    A_GEN_WAV_UP_PW1               = auto() # ConvNeXt pointwise conv 1 (expand)
+    A_GEN_WAV_UP_PW2               = auto() # ConvNeXt pointwise conv 2 (project)
+    A_GEN_WAV_UP_GAMMA             = auto() # ConvNeXt layer scale
+    A_GEN_WAV_DAC_ENTRY            = auto() # DAC conv_pre
+    A_GEN_WAV_DAC_UP_SNAKE         = auto() # DAC per-block SnakeBeta before the upsample conv
+    A_GEN_WAV_DAC_UP_CONV          = auto() # DAC per-block causal ConvTranspose1d
+    A_GEN_WAV_DAC_RES_ACT1         = auto() # DAC residual unit, SnakeBeta before conv1
+    A_GEN_WAV_DAC_RES_CONV1        = auto() # DAC residual unit, dilated causal conv
+    A_GEN_WAV_DAC_RES_ACT2         = auto() # DAC residual unit, SnakeBeta before conv2
+    A_GEN_WAV_DAC_RES_CONV2        = auto() # DAC residual unit, pointwise causal conv
+    A_GEN_WAV_DAC_POST_SNAKE       = auto() # DAC final SnakeBeta
+    A_GEN_WAV_DAC_POST_CONV        = auto() # DAC conv_post -> 1-channel PCM
     A_MMPROJ              = auto()
     A_MMPROJ_FC           = auto()
     A_MM_NORM_PRE         = auto()
@@ -1170,6 +1243,7 @@ MODEL_ARCH_NAMES: dict[MODEL_ARCH, str] = {
     MODEL_ARCH.TALKIE:           "talkie",
     MODEL_ARCH.MELLUM:           "mellum",
     MODEL_ARCH.NANBEIGE:         "nanbeige",
+    MODEL_ARCH.QWEN3TTS:         "qwen3tts",
 }
 
 VISION_PROJECTOR_TYPE_NAMES: dict[VISION_PROJECTOR_TYPE, str] = {
@@ -1567,6 +1641,63 @@ TENSOR_NAMES: dict[MODEL_TENSOR, str] = {
     MODEL_TENSOR.A_ENC_DOWNSAMPLE_CONV:     "a.downsample.conv",
     MODEL_TENSOR.A_ENC_DOWNSAMPLE_NORM:     "a.downsample.norm",
     MODEL_TENSOR.A_ENC_RVQ_CODEBOOK:        "a.rvq.codebook",
+    MODEL_TENSOR.A_ENC_CONV_RES2:           "a.blk.{bid}.res2.{xid}",
+    MODEL_TENSOR.A_ENC_SE_CONV1:            "a.blk.{bid}.se_conv1",
+    MODEL_TENSOR.A_ENC_SE_CONV2:            "a.blk.{bid}.se_conv2",
+    MODEL_TENSOR.A_ENC_ASP_ATTN:            "a.asp_attn",
+    MODEL_TENSOR.A_ENC_ASP_TDNN:            "a.asp_tdnn",
+    MODEL_TENSOR.A_GEN_CODE_PROJ_IN:        "a.gen.code.proj_in",
+    MODEL_TENSOR.A_GEN_CODE_EMBD:           "a.gen.code.embd",
+    MODEL_TENSOR.A_GEN_CODE_HEAD:           "a.gen.code.head",
+    MODEL_TENSOR.A_GEN_CODE_OUT_EMBD:       "a.gen.code.out_embd",
+    MODEL_TENSOR.A_GEN_CODE_ATTN_NORM:      "a.gen.code.blk.{bid}.ln1", # reuses the generic clip.cpp block loader (TN_LN_1)
+    MODEL_TENSOR.A_GEN_CODE_ATTN_Q:         "a.gen.code.blk.{bid}.attn_q",
+    MODEL_TENSOR.A_GEN_CODE_ATTN_Q_NORM:    "a.gen.code.blk.{bid}.attn_q_norm",
+    MODEL_TENSOR.A_GEN_CODE_ATTN_K:         "a.gen.code.blk.{bid}.attn_k",
+    MODEL_TENSOR.A_GEN_CODE_ATTN_K_NORM:    "a.gen.code.blk.{bid}.attn_k_norm",
+    MODEL_TENSOR.A_GEN_CODE_ATTN_V:         "a.gen.code.blk.{bid}.attn_v",
+    MODEL_TENSOR.A_GEN_CODE_ATTN_OUT:       "a.gen.code.blk.{bid}.attn_out",
+    MODEL_TENSOR.A_GEN_CODE_FFN_NORM:       "a.gen.code.blk.{bid}.ln2", # reuses the generic clip.cpp block loader (TN_LN_2)
+    MODEL_TENSOR.A_GEN_CODE_FFN_GATE:       "a.gen.code.blk.{bid}.ffn_gate",
+    MODEL_TENSOR.A_GEN_CODE_FFN_UP:         "a.gen.code.blk.{bid}.ffn_up",
+    MODEL_TENSOR.A_GEN_CODE_FFN_DOWN:       "a.gen.code.blk.{bid}.ffn_down",
+    MODEL_TENSOR.A_GEN_CODE_OUTPUT_NORM:    "a.gen.code.output_norm",
+    MODEL_TENSOR.A_GEN_WAV_QUANT_FIRST_IN:  "a.gen.wav.quant.first.in_proj",
+    MODEL_TENSOR.A_GEN_WAV_QUANT_FIRST_OUT: "a.gen.wav.quant.first.out_proj",
+    MODEL_TENSOR.A_GEN_WAV_QUANT_FIRST_CB:  "a.gen.wav.quant.first.codebook",
+    MODEL_TENSOR.A_GEN_WAV_QUANT_REST_IN:   "a.gen.wav.quant.rest.in_proj",
+    MODEL_TENSOR.A_GEN_WAV_QUANT_REST_OUT:  "a.gen.wav.quant.rest.out_proj",
+    MODEL_TENSOR.A_GEN_WAV_QUANT_REST_CB:   "a.gen.wav.quant.rest.codebook",
+    MODEL_TENSOR.A_GEN_WAV_PRE_CONV:        "a.gen.wav.pre_conv",
+    MODEL_TENSOR.A_GEN_WAV_TFM_IN_PROJ:     "a.gen.wav.tfm.in_proj",
+    MODEL_TENSOR.A_GEN_WAV_TFM_OUT_PROJ:    "a.gen.wav.tfm.out_proj",
+    MODEL_TENSOR.A_GEN_WAV_TFM_OUTPUT_NORM: "a.gen.wav.tfm.output_norm",
+    MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_NORM:   "a.gen.wav.tfm.blk.{bid}.ln1",
+    MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_Q:      "a.gen.wav.tfm.blk.{bid}.attn_q",
+    MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_K:      "a.gen.wav.tfm.blk.{bid}.attn_k",
+    MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_V:      "a.gen.wav.tfm.blk.{bid}.attn_v",
+    MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_OUT:    "a.gen.wav.tfm.blk.{bid}.attn_out",
+    MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_SCALE:  "a.gen.wav.tfm.blk.{bid}.ls1",
+    MODEL_TENSOR.A_GEN_WAV_TFM_FFN_NORM:    "a.gen.wav.tfm.blk.{bid}.ln2",
+    MODEL_TENSOR.A_GEN_WAV_TFM_FFN_GATE:    "a.gen.wav.tfm.blk.{bid}.ffn_gate",
+    MODEL_TENSOR.A_GEN_WAV_TFM_FFN_UP:      "a.gen.wav.tfm.blk.{bid}.ffn_up",
+    MODEL_TENSOR.A_GEN_WAV_TFM_FFN_DOWN:    "a.gen.wav.tfm.blk.{bid}.ffn_down",
+    MODEL_TENSOR.A_GEN_WAV_TFM_FFN_SCALE:   "a.gen.wav.tfm.blk.{bid}.ls2",
+    MODEL_TENSOR.A_GEN_WAV_UP_CONV:         "a.gen.wav.up.blk.{bid}.conv",
+    MODEL_TENSOR.A_GEN_WAV_UP_DWCONV:       "a.gen.wav.up.blk.{bid}.dwconv",
+    MODEL_TENSOR.A_GEN_WAV_UP_NORM:         "a.gen.wav.up.blk.{bid}.norm",
+    MODEL_TENSOR.A_GEN_WAV_UP_PW1:          "a.gen.wav.up.blk.{bid}.pw1",
+    MODEL_TENSOR.A_GEN_WAV_UP_PW2:          "a.gen.wav.up.blk.{bid}.pw2",
+    MODEL_TENSOR.A_GEN_WAV_UP_GAMMA:        "a.gen.wav.up.blk.{bid}.gamma",
+    MODEL_TENSOR.A_GEN_WAV_DAC_ENTRY:       "a.gen.wav.dac.entry",
+    MODEL_TENSOR.A_GEN_WAV_DAC_UP_SNAKE:    "a.gen.wav.dac.blk.{bid}.snake",
+    MODEL_TENSOR.A_GEN_WAV_DAC_UP_CONV:     "a.gen.wav.dac.blk.{bid}.conv",
+    MODEL_TENSOR.A_GEN_WAV_DAC_RES_ACT1:    "a.gen.wav.dac.blk.{bid}.res.{xid}.act1",
+    MODEL_TENSOR.A_GEN_WAV_DAC_RES_CONV1:   "a.gen.wav.dac.blk.{bid}.res.{xid}.conv1",
+    MODEL_TENSOR.A_GEN_WAV_DAC_RES_ACT2:    "a.gen.wav.dac.blk.{bid}.res.{xid}.act2",
+    MODEL_TENSOR.A_GEN_WAV_DAC_RES_CONV2:   "a.gen.wav.dac.blk.{bid}.res.{xid}.conv2",
+    MODEL_TENSOR.A_GEN_WAV_DAC_POST_SNAKE:  "a.gen.wav.dac.post_snake",
+    MODEL_TENSOR.A_GEN_WAV_DAC_POST_CONV:   "a.gen.wav.dac.post_conv",
     MODEL_TENSOR.A_MMPROJ:                  "mm.a.mlp.{bid}",
     MODEL_TENSOR.A_MMPROJ_FC:               "mm.a.fc",
     MODEL_TENSOR.A_MM_NORM_PRE:             "mm.a.norm_pre",
@@ -1821,6 +1952,63 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = {
         MODEL_TENSOR.A_ENC_CONV_NORM,
         MODEL_TENSOR.A_ENC_CONV_PW1,
         MODEL_TENSOR.A_ENC_CONV_PW2,
+        MODEL_TENSOR.A_ENC_CONV_RES2,
+        MODEL_TENSOR.A_ENC_SE_CONV1,
+        MODEL_TENSOR.A_ENC_SE_CONV2,
+        MODEL_TENSOR.A_ENC_ASP_ATTN,
+        MODEL_TENSOR.A_ENC_ASP_TDNN,
+        MODEL_TENSOR.A_GEN_CODE_PROJ_IN,
+        MODEL_TENSOR.A_GEN_CODE_EMBD,
+        MODEL_TENSOR.A_GEN_CODE_HEAD,
+        MODEL_TENSOR.A_GEN_CODE_OUT_EMBD,
+        MODEL_TENSOR.A_GEN_CODE_ATTN_NORM,
+        MODEL_TENSOR.A_GEN_CODE_ATTN_Q,
+        MODEL_TENSOR.A_GEN_CODE_ATTN_Q_NORM,
+        MODEL_TENSOR.A_GEN_CODE_ATTN_K,
+        MODEL_TENSOR.A_GEN_CODE_ATTN_K_NORM,
+        MODEL_TENSOR.A_GEN_CODE_ATTN_V,
+        MODEL_TENSOR.A_GEN_CODE_ATTN_OUT,
+        MODEL_TENSOR.A_GEN_CODE_FFN_NORM,
+        MODEL_TENSOR.A_GEN_CODE_FFN_GATE,
+        MODEL_TENSOR.A_GEN_CODE_FFN_UP,
+        MODEL_TENSOR.A_GEN_CODE_FFN_DOWN,
+        MODEL_TENSOR.A_GEN_CODE_OUTPUT_NORM,
+        MODEL_TENSOR.A_GEN_WAV_QUANT_FIRST_IN,
+        MODEL_TENSOR.A_GEN_WAV_QUANT_FIRST_OUT,
+        MODEL_TENSOR.A_GEN_WAV_QUANT_FIRST_CB,
+        MODEL_TENSOR.A_GEN_WAV_QUANT_REST_IN,
+        MODEL_TENSOR.A_GEN_WAV_QUANT_REST_OUT,
+        MODEL_TENSOR.A_GEN_WAV_QUANT_REST_CB,
+        MODEL_TENSOR.A_GEN_WAV_PRE_CONV,
+        MODEL_TENSOR.A_GEN_WAV_TFM_IN_PROJ,
+        MODEL_TENSOR.A_GEN_WAV_TFM_OUT_PROJ,
+        MODEL_TENSOR.A_GEN_WAV_TFM_OUTPUT_NORM,
+        MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_NORM,
+        MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_Q,
+        MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_K,
+        MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_V,
+        MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_OUT,
+        MODEL_TENSOR.A_GEN_WAV_TFM_ATTN_SCALE,
+        MODEL_TENSOR.A_GEN_WAV_TFM_FFN_NORM,
+        MODEL_TENSOR.A_GEN_WAV_TFM_FFN_GATE,
+        MODEL_TENSOR.A_GEN_WAV_TFM_FFN_UP,
+        MODEL_TENSOR.A_GEN_WAV_TFM_FFN_DOWN,
+        MODEL_TENSOR.A_GEN_WAV_TFM_FFN_SCALE,
+        MODEL_TENSOR.A_GEN_WAV_UP_CONV,
+        MODEL_TENSOR.A_GEN_WAV_UP_DWCONV,
+        MODEL_TENSOR.A_GEN_WAV_UP_NORM,
+        MODEL_TENSOR.A_GEN_WAV_UP_PW1,
+        MODEL_TENSOR.A_GEN_WAV_UP_PW2,
+        MODEL_TENSOR.A_GEN_WAV_UP_GAMMA,
+        MODEL_TENSOR.A_GEN_WAV_DAC_ENTRY,
+        MODEL_TENSOR.A_GEN_WAV_DAC_UP_SNAKE,
+        MODEL_TENSOR.A_GEN_WAV_DAC_UP_CONV,
+        MODEL_TENSOR.A_GEN_WAV_DAC_RES_ACT1,
+        MODEL_TENSOR.A_GEN_WAV_DAC_RES_CONV1,
+        MODEL_TENSOR.A_GEN_WAV_DAC_RES_ACT2,
+        MODEL_TENSOR.A_GEN_WAV_DAC_RES_CONV2,
+        MODEL_TENSOR.A_GEN_WAV_DAC_POST_SNAKE,
+        MODEL_TENSOR.A_GEN_WAV_DAC_POST_CONV,
         MODEL_TENSOR.A_ENC_CONV_NORM_MEAN,
         MODEL_TENSOR.A_ENC_CONV_NORM_VAR,
         MODEL_TENSOR.A_ENC_MEL_FILTERS,
@@ -4648,6 +4836,22 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = {
         MODEL_TENSOR.FFN_DOWN,
         MODEL_TENSOR.FFN_UP,
     ],
+    MODEL_ARCH.QWEN3TTS: [
+        MODEL_TENSOR.TOKEN_EMBD,
+        MODEL_TENSOR.OUTPUT_NORM,
+        MODEL_TENSOR.OUTPUT,
+        MODEL_TENSOR.ATTN_NORM,
+        MODEL_TENSOR.ATTN_Q,
+        MODEL_TENSOR.ATTN_Q_NORM,
+        MODEL_TENSOR.ATTN_K,
+        MODEL_TENSOR.ATTN_K_NORM,
+        MODEL_TENSOR.ATTN_V,
+        MODEL_TENSOR.ATTN_OUT,
+        MODEL_TENSOR.FFN_NORM,
+        MODEL_TENSOR.FFN_GATE,
+        MODEL_TENSOR.FFN_DOWN,
+        MODEL_TENSOR.FFN_UP,
+    ],
 }
 
 # tensors that will not be serialized
@@ -4922,6 +5126,8 @@ class VisionProjectorType:
     GLM4V = "glm4v"
     YOUTUVL = "youtuvl"
     NEMOTRON_V2_VL = "nemotron_v2_vl"
+    QWEN3TTS_SPKENC = "qwen3tts_spkenc" # audio: ECAPA-TDNN speaker encoder
+    QWEN3TTS_GEN = "qwen3tts_gen" # audio generation: code_predictor
     HUNYUANVL      = "hunyuanvl"
     PARAKEET       = "parakeet"  # audio
     MINIMAXM3      = "minimax_m3"
index c5905164c356dfb0310e7478ced926712e33c579..39da9f2c05fb9698d23c7c34dc5ad8fe17da90d3 100644 (file)
@@ -280,6 +280,10 @@ class GGUFWriter:
 
         self.kv_data[0][key] = GGUFValue(value=val, type=vtype, sub_type=sub_type)
 
+    def remove_key(self, key: str) -> None:
+        for kv_data in self.kv_data:
+            kv_data.pop(key, None)
+
     def add_uint8(self, key: str, val: int) -> None:
         self.add_key_value(key,val, GGUFValueType.UINT8)
 
@@ -1144,7 +1148,11 @@ class GGUFWriter:
     def add_precompiled_charsmap(self, charsmap: bytes) -> None:
         self.add_array(Keys.Tokenizer.PRECOMPILED_CHARSMAP, charsmap)
 
-    def add_chat_template(self, value: str | Sequence[Mapping[str, str]]) -> None:
+    def add_chat_template(self, value: str | Sequence[Mapping[str, str]] | None) -> None:
+        if value is None:
+            self.remove_key(Keys.Tokenizer.CHAT_TEMPLATE)
+            return
+
         if not isinstance(value, str):
             template_default = None
             template_names = set()
@@ -1199,6 +1207,9 @@ class GGUFWriter:
     def add_clip_has_audio_encoder(self, value: bool) -> None:
         self.add_bool(Keys.Clip.HAS_AUDIO_ENCODER, value)
 
+    def add_clip_has_gen_audio_encoder(self, value: bool) -> None:
+        self.add_bool(Keys.Clip.HAS_GEN_AUDIO_ENCODER, value)
+
     def add_clip_projector_type(self, value: str) -> None:
         self.add_string(Keys.Clip.PROJECTOR_TYPE, value)
 
@@ -1401,6 +1412,32 @@ class GGUFWriter:
     def add_audio_projector_head_count(self, value: int) -> None:
         self.add_uint32(Keys.ClipAudio.Projector.HEAD_COUNT, value)
 
+    # audio generation (mmproj)
+
+    def add_clip_gen_audio_projector_type(self, value: str) -> None:
+        self.add_string(Keys.ClipGenAudio.PROJECTOR_TYPE, value)
+
+    def add_gen_audio_projection_dim(self, value: int) -> None:
+        self.add_uint32(Keys.ClipGenAudio.PROJECTION_DIM, value)
+
+    def add_gen_audio_embedding_length(self, value: int) -> None:
+        self.add_uint32(Keys.ClipGenAudio.EMBEDDING_LENGTH, value)
+
+    def add_gen_audio_feed_forward_length(self, value: int) -> None:
+        self.add_uint32(Keys.ClipGenAudio.FEED_FORWARD_LENGTH, value)
+
+    def add_gen_audio_block_count(self, value: int) -> None:
+        self.add_uint32(Keys.ClipGenAudio.BLOCK_COUNT, value)
+
+    def add_gen_audio_head_count(self, value: int) -> None:
+        self.add_uint32(Keys.ClipGenAudio.Attention.HEAD_COUNT, value)
+
+    def add_gen_audio_head_count_kv(self, value: int) -> None:
+        self.add_uint32(Keys.ClipGenAudio.Attention.HEAD_COUNT_KV, value)
+
+    def add_gen_audio_attention_layernorm_eps(self, value: float) -> None:
+        self.add_float32(Keys.ClipGenAudio.Attention.LAYERNORM_EPS, value)
+
     def add_xielu_alpha_p(self, values: Sequence[float]):
         self.add_array(Keys.xIELU.ALPHA_P, values)
 
index 1e991b873ceab0b0130f7b70583e3ef2fff94c65..7892342e473dad19aec3d83ce9b6a99d071d9033 100644 (file)
@@ -2109,6 +2109,7 @@ class TensorNameMap:
             "conformer.subsample_conv_projection.layer{bid}.conv", # gemma4
             "sound_encoder.encoder.subsampling.layers.{bid}", # parakeet
             "encoder.conv{bid}", # mimo-audio-tokenizer
+            "speaker_encoder.blocks.{bid}.conv", # qwen3tts speaker encoder (only bid=0, the stem TDNN)
         ),
 
         MODEL_TENSOR.A_ENC_CONV1D_NORM: (
@@ -2126,6 +2127,7 @@ class TensorNameMap:
 
         MODEL_TENSOR.A_ENC_CONV_OUT: (
             "audio_tower.conv_out", # qwen3omni
+            "speaker_encoder.mfa.conv", # qwen3tts speaker encoder: multi-layer feature aggregation
         ),
 
         MODEL_TENSOR.A_PRE_NORM: (),
@@ -2336,7 +2338,8 @@ class TensorNameMap:
         MODEL_TENSOR.A_MMPROJ_FC: (
             "audio.multi_modal_projector.linear", # qwen2audio
             "audio_tower.proj", # qwen2omni
-            "model.audio_tower.output_proj" # gemma4
+            "model.audio_tower.output_proj", # gemma4
+            "speaker_encoder.fc", # qwen3tts speaker encoder: final speaker embedding projection
         ),
 
         MODEL_TENSOR.A_MM_NORM_PRE: (
@@ -2411,6 +2414,7 @@ class TensorNameMap:
             "conformer.layers.{bid}.lconv1d.linear_start", # gemma3n
             "sound_encoder.encoder.layers.{bid}.conv.pointwise_conv1", # parakeet
             "encoder.layers.{bid}.conv.up_conv", # granite_speech
+            "speaker_encoder.blocks.{bid}.tdnn1.conv", # qwen3tts speaker encoder
         ),
 
         MODEL_TENSOR.A_ENC_CONV_PW2: (
@@ -2418,6 +2422,23 @@ class TensorNameMap:
             "conformer.layers.{bid}.lconv1d.linear_end", # gemma3n
             "sound_encoder.encoder.layers.{bid}.conv.pointwise_conv2", # parakeet
             "encoder.layers.{bid}.conv.down_conv", # granite_speech
+            "speaker_encoder.blocks.{bid}.tdnn2.conv", # qwen3tts speaker encoder
+        ),
+
+        MODEL_TENSOR.A_ENC_SE_CONV1: (
+            "speaker_encoder.blocks.{bid}.se_block.conv1", # qwen3tts
+        ),
+
+        MODEL_TENSOR.A_ENC_SE_CONV2: (
+            "speaker_encoder.blocks.{bid}.se_block.conv2", # qwen3tts
+        ),
+
+        MODEL_TENSOR.A_ENC_ASP_ATTN: (
+            "speaker_encoder.asp.conv", # qwen3tts
+        ),
+
+        MODEL_TENSOR.A_ENC_ASP_TDNN: (
+            "speaker_encoder.asp.tdnn.conv", # qwen3tts
         ),
 
         MODEL_TENSOR.A_ENC_NORM_CONV: (
index 726edbb0ccb1864a90fded844a31e1b0141f4b23..b9372ddda870db50b9d8c0024ee000ae74d785d8 100644 (file)
@@ -119,6 +119,7 @@ Public API changes carry a higher bar than internal ones (`CONTRIBUTING.md`). Re
 - In most cases, `build_vit` should be enough to build the transformer graph for vision models. Do not add a loop to build the transformer graph manually, unless you have a very good reason to do so. If you do, please explain why in the PR description.
 - If you need a dedicated preprocessor, there is a high chance that it can be a derived class from one of the existing preprocessors. Check carefully before adding a new preprocessor class.
 - If the model need a new public API in `mtmd.h`, open a discussion first.
+- For audio generation models, see `tools/mtmd/README-dev.md`
 
 ## General (always)
 
index ea0ddd114c0c525d93791026a80e691a8f0edc1d..836cfade226c479c27ee1ac01d0f5750207d0138 100644 (file)
@@ -144,6 +144,7 @@ static const std::map<llm_arch, const char *> LLM_ARCH_NAMES = {
     { LLM_ARCH_TALKIE,           "talkie"           },
     { LLM_ARCH_MELLUM,           "mellum"           },
     { LLM_ARCH_NANBEIGE,         "nanbeige"         },
+    { LLM_ARCH_QWEN3TTS,         "qwen3tts"         },
     { LLM_ARCH_UNKNOWN,          "(unknown)"        },
 };
 
@@ -1026,6 +1027,7 @@ bool llm_arch_supports_sm_tensor(const llm_arch & arch) {
         case LLM_ARCH_MINIMAX_M3:
         case LLM_ARCH_MISTRAL4:
         case LLM_ARCH_KIMI_LINEAR:
+        case LLM_ARCH_QWEN3TTS:
             return false;
         default:
             return true;
index cbc97085ea79e20a2c6b967d70f082f89c4e91bc..49c2a6ac3997c0e111a3d2c896fb5615008eafcc 100644 (file)
@@ -149,6 +149,7 @@ enum llm_arch {
     LLM_ARCH_MINIMAX_M3,
     LLM_ARCH_DFLASH,
     LLM_ARCH_NANBEIGE,
+    LLM_ARCH_QWEN3TTS,
     LLM_ARCH_UNKNOWN,
 };
 
index 348bbae95770f3bcb6b2eb464624a24fa4d8028f..35d6e58adfa815d344dcc884331719810f7f1732 100644 (file)
@@ -124,3 +124,9 @@ LLAMA_API llama_context * llama_get_ctx_other(struct llama_context * ctx);
 LLAMA_API const int32_t * llama_model_target_layer_ids  (const struct llama_model * model);
 // returns the number of extracted layers from target model
 LLAMA_API uint32_t        llama_model_target_layer_ids_n(const struct llama_model * model);
+
+// retrieves the whole token embedding matrix in F32 format (n_embd * n_vocab)
+// returns total number of elements or 0 on error
+// if out is nullptr, returns the number of tokens without writing to out
+// caller must allocate enough memory for out before calling
+LLAMA_API uint32_t llama_model_get_tok_embd(const struct llama_model * model, float * out);
index 333f506de55788524efd8074d2728d2ed33fa72d..dda311c47bbf64c0333c2c48b23f16bf24153e42 100644 (file)
@@ -112,6 +112,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params
             return new llama_model_qwen3vl(params);
         case LLM_ARCH_QWEN3VLMOE:
             return new llama_model_qwen3vlmoe(params);
+        case LLM_ARCH_QWEN3TTS:
+            return new llama_model_qwen3tts(params);
         case LLM_ARCH_PHI2:
             return new llama_model_phi2(params);
         case LLM_ARCH_PHI3:
@@ -2693,6 +2695,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) {
         case LLM_ARCH_QWEN3VLMOE:
         case LLM_ARCH_QWEN35:
         case LLM_ARCH_QWEN35MOE:
+        case LLM_ARCH_QWEN3TTS:
             return LLAMA_ROPE_TYPE_IMROPE;
 
         case LLM_ARCH_GLM4:
@@ -2908,3 +2911,38 @@ const int32_t * llama_model_target_layer_ids(const struct llama_model * model) {
 uint32_t llama_model_target_layer_ids_n(const struct llama_model * model) {
     return (uint32_t) model->target_layer_ids.size();
 }
+
+uint32_t llama_model_get_tok_embd(const struct llama_model * model, float * out) {
+    if (model->vocab.n_tokens() == 0 || model->tok_embd == nullptr) {
+        return 0;
+    }
+
+    const ggml_tensor * tensor = model->tok_embd;
+    const size_t nelements = ggml_nelements(tensor);
+    GGML_ASSERT(nelements <= UINT32_MAX); // for the return type
+
+    if (out == nullptr) {
+        return (uint32_t) nelements;
+    }
+
+    if (tensor->type == GGML_TYPE_F32) {
+        ggml_backend_tensor_get(tensor, out, 0, nelements * sizeof(float));
+        return (uint32_t) nelements;
+    }
+
+    std::vector<uint8_t> buf(ggml_nbytes(tensor));
+    ggml_backend_tensor_get(tensor, buf.data(), 0, buf.size());
+
+    const ggml_type_traits * traits = ggml_get_type_traits(tensor->type);
+    if (tensor->type == GGML_TYPE_F16) {
+        ggml_fp16_to_fp32_row((const ggml_fp16_t *) buf.data(), out, nelements);
+    } else if (tensor->type == GGML_TYPE_BF16) {
+        ggml_bf16_to_fp32_row((const ggml_bf16_t *) buf.data(), out, nelements);
+    } else if (ggml_is_quantized(tensor->type) && traits->to_float != nullptr) {
+        traits->to_float(buf.data(), out, nelements);
+    } else {
+        GGML_ABORT("unsupported tensor type for dequantization: %s", ggml_type_name(tensor->type));
+    }
+
+    return (uint32_t) nelements;
+}
index 5f206621d579eee455dca25fd80616745f28d36c..ad3dadaf39320a6871a94499f32ecb0709d9319e 100644 (file)
@@ -596,6 +596,11 @@ struct llama_model_qwen3vlmoe : public llama_model_base {
 };
 
 
+struct llama_model_qwen3tts : public llama_model_qwen3vl {
+    llama_model_qwen3tts(const struct llama_model_params & params) : llama_model_qwen3vl(params) {}
+};
+
+
 struct llama_model_phi2 : public llama_model_base {
     llama_model_phi2(const struct llama_model_params & params) : llama_model_base(params) {}
     void load_arch_hparams(llama_model_loader & ml) override;
diff --git a/src/models/qwen3tts.cpp b/src/models/qwen3tts.cpp
new file mode 100644 (file)
index 0000000..3604f84
--- /dev/null
@@ -0,0 +1,3 @@
+#include "models.h"
+
+// llama_model_qwen3tts reuses llama_model_qwen3vl's hparams/tensors/graph logic
index 724d6140d193655f7b477317c38bccd7ee840078..5596620f078272399b6293c6c298ffc640a16bc6 100644 (file)
@@ -16,11 +16,16 @@ void llama_model_qwen3vl::load_arch_hparams(llama_model_loader & ml) {
 void llama_model_qwen3vl::load_arch_tensors(llama_model_loader &) {
     LLAMA_LOAD_LOCALS;
 
+    int64_t n_vocab_out = n_vocab;
+    if (arch == LLM_ARCH_QWEN3TTS) {
+        n_vocab_out = 3072;
+    }
+
     tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0);
 
     // output
     output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0);
-    output      = create_tensor(tn(LLM_TENSOR_OUTPUT,      "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED);
+    output      = create_tensor(tn(LLM_TENSOR_OUTPUT,      "weight"), {n_embd, n_vocab_out}, TENSOR_NOT_REQUIRED);
     // if output is NULL, init from the input tok embed
     if (output == NULL) {
         output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED);
@@ -166,6 +171,24 @@ llama_model_qwen3vl::graph::graph(const llama_model & model, const llm_graph_par
     // lm_head
     cur = build_lora_mm(model.output, cur, model.output_s);
 
+    int64_t n_vocab_in  = model.tok_embd->ne[1];
+    int64_t n_vocab_out = model.output->ne[1];
+    if (n_vocab_in > n_vocab_out) {
+        // case: Qwen3TTS model with codec_head as output
+        GGML_ASSERT(model.output_norm);
+        int64_t pad = n_vocab_in - n_vocab_out;
+
+        // using this trick to get a scalar -inf tensor to pad the output
+        ggml_tensor * neg_inf = ggml_scale_bias(ctx0,
+                ggml_view_1d(ctx0, model.output_norm, 1, 0),
+                0.0f, -INFINITY);
+        neg_inf = ggml_repeat_4d(ctx0, neg_inf, pad, cur->ne[1], 1, 1);
+        cur = ggml_concat(ctx0, neg_inf, cur, 0); // [padded .. n_vocab_out, n_stream]
+
+    } else if (n_vocab_in < n_vocab_out) {
+        GGML_ABORT("invalid case");
+    }
+
     cb(cur, "result_output", -1);
     res->t_logits = cur;
 
index 4336e4e13d4f6b82da1b94802d21dadae831b9b9..1654f122a731cab3a45eb9c44fd2a585516695aa 100644 (file)
@@ -113,6 +113,8 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
         n_layer = 3;
     } else if (arch == LLM_ARCH_CHAMELEON) {
         n_vocab = 10240;
+    } else if (arch == LLM_ARCH_QWEN3TTS) {
+        n_vocab = 4096; // must be >= the hard-coded codec head size (3072)
     }
 
     const uint32_t n_embd_head = n_embd / n_head;
index 15040e4af5f9820cb40ce1245ef0838995e0696a..4675fb9a97b6a91399724400276a352a927e268b 100644 (file)
@@ -18,6 +18,8 @@ add_library(mtmd
             mtmd-image.cpp
             mtmd.h
             mtmd-helper.cpp
+            mtmd-helper-gen.cpp
+            mtmd-helper-common.h
             mtmd-helper.h
             clip.cpp
             clip.h
@@ -52,6 +54,8 @@ add_library(mtmd
             models/mimovl.cpp
             models/qwen3a.cpp
             models/mimo-audio.cpp
+            models/qwen3tts-spkenc.cpp
+            models/qwen3tts-gen.cpp
             models/step3vl.cpp
             models/siglip.cpp
             models/whisper-enc.cpp
index 3a08915876c113ea7673ebac8d4802b3559767ad..3cddd085ec6c2b3c220eb485d6a20fe30e7681a0 100644 (file)
@@ -33,3 +33,52 @@ A typical pipeline of the core libmtmd is as follows:
 We provide a set of helper functions via `mtmd_helper` to make using libmtmd easier. The helper provides:
 - Image, audio and video file decoding (for example, decode raw JPEG into RGB bitmap)
 - Manage `llama_batch` and calls to `llama_decode`
+
+## Audio generation support
+
+Audio generation is added to mtmd in PR [#26254](https://github.com/ggml-org/llama.cpp/pull/26254)
+
+Currently, we support the 3-stage pipeline below which should cover most TTS models:
+- Stage 1: Backbone / Semantic Stage: Backbone model accepts text prompt and reference voice as input
+- Stage 2: Acoustic Detail Generator: A model takes the hidden state from backbone and generate audio details (usually as audio codes or mel-spectrogram)
+- Stage 3: Waveform Reconstruction: Convert the semantic and acoustic data from previous stages to the final waveform
+
+For example, Qwen3-TTS:
+- Reference voice is encoded using ECAPA-TDNN speaker encoder (`speaker_encoder`)
+- Text prompt and reference voice are processed via a backbone (`talker.model`)
+- A model converts sampled semantic token and hidden state from stage 2 into a list of 15 acoustic codes (`talker.code_predictor`)
+- 16 generated codes are converted into waveform (`code2wav`)
+
+### API design constraints
+
+Due to wide variety of audio generation pipelines, the `mtmd_gen_audio` system is designed to be flexible and reusable by new models.
+
+`mtmd_gen_audio` is split into 2 main API:
+- Core API `mtmd.h`: handles main inference. Important: the API surface must be stateless; caller must handle state management and audio frame accumulation.
+- Helper API `mtmd-helper.h`: provides a model-agnostic stateful API. Usage example can be found in the `tools/tts` directory.
+
+### Checklist for porting new audio generation models to mtmd
+
+1. Establish a list of reusable and missing components from the current mtmd implementation.
+2. For GGUF conversion:
+    - Backbone model should be converted to a normal text model (loadable via `libllama`)
+        - If model used hard-coded embedding row ID, append them to token embeddings and assign token name for them (see `qwen3tts.py`)
+        - If model have a specific output logits head for audio codes (usually semantic code), keep the head as-is and pad the logits at inference time (see `src/models/qwen3vl.cpp`)
+    - Sidecar models (code2wav, bigvgan, etc) must live inside the mmproj GGUF (but can be in different `clip_context` if necessary)
+        - Note: it should use `ggml_build_forward_select` to select graphs if multiple graphs living in the same context
+    - Reuse existing GGUF metadata key name and tensor name whenever possible; think twice before adding extensive changes to GGUF writer. For example, Qwen3-TTS hard-code part of the hparams to `clip.cpp` as they won't likely to change.
+    - For tensor naming:
+        - Prefixed with `a.*` for tensors used by speaker encoder pipeline
+        - Prefixed with `a.gen.*` for generation stages (code / mel-spectrogram / PCM generation)
+3. Make sure most of the changes happen inside `mtmd-helper-gen.cpp`. A good PR looks like this:
+    - 10-20% changes is to add new backbone (text) model and conversion
+    - 60% changes inside `mtmd-helper-gen.cpp`
+    - 10% changes inside `libmtmd` and `clip.cpp` systems
+    - The rest downstream code (CLI, server) should have no changes at all
+4. Update usage documentation in `tools/tts/README.md`
+
+IMPORTANT: If your model needs changes that don't fit the existing infrastructure, **open an issue first for discussion**.
+
+No-go checklist (these will get the PR rejected and require discussion before proceeding):
+- Violating the API design constraints stated above
+- Adding a new model-specific binary: the API and binary surface must stay model-agnostic
index 29352abb4c0bc443757a9685e99c20355dda1636..e12140ba009d3201f94a423137e787e1e22c091c 100644 (file)
@@ -54,6 +54,9 @@ struct clip_graph {
 
     clip_graph(clip_ctx * ctx, const clip_image_f32 & img);
 
+    // build sub-graph, reuse buf from parent
+    clip_graph(const clip_graph & parent);
+
     virtual ~clip_graph() = default;
     virtual ggml_cgraph * build() = 0;
 
index d42b38222c2749d5b01eb587c42e3eed47b8f6d4..e1567ee5bafbe19e968d40ffb28d1e2ae53f4577 100644 (file)
@@ -32,6 +32,7 @@
 #define KEY_PROJ_TYPE           "clip.projector_type"
 #define KEY_HAS_AUDIO_ENC       "clip.has_audio_encoder"
 #define KEY_HAS_VISION_ENC      "clip.has_vision_encoder"
+#define KEY_HAS_GEN_AUDIO_ENC   "clip.has_gen_audio_encoder"
 #define KEY_USE_GELU            "clip.use_gelu"
 #define KEY_USE_SILU            "clip.use_silu"
 
@@ -89,6 +90,8 @@
 #define KEY_A_ATTN_WINDOW_SIZE     "clip.audio.window_size"          // mimo-audio-tokenizer: sliding-window radius
 #define KEY_A_LOCAL_BLOCK_COUNT    "clip.audio.local_block_count"    // mimo-v2.5: input_local_transformer layer count
 #define KEY_A_LOCAL_GROUP_SIZE     "clip.audio.local_group_size"     // mimo-v2.5: input_local_transformer grouping size
+// audio generation (gen-audio)-specific
+#define KEY_GEN_AUDIO_PROJ_TYPE    "clip.gen.audio.projector_type" // for models with mixed modalities
 #define KEY_AUDIO_SUBSAMPLING_FACTOR "clip.audio.subsampling_factor"
 
 //
 #define TN_MM_A_LOCAL_LN2      "mm.a.local_blk.%d.ln2.%s"
 #define TN_MM_A_LOCAL_NORM     "mm.a.local_norm.%s"
 
+// qwen3tts speaker encoder (ECAPA-TDNN)
+#define TN_A_SE_CONV1  "a.blk.%d.se_conv1.%s"
+#define TN_A_SE_CONV2  "a.blk.%d.se_conv2.%s"
+#define TN_A_CONV_RES2 "a.blk.%d.res2.%d.%s"
+#define TN_A_ASP_ATTN  "a.asp_attn.%s"
+#define TN_A_ASP_TDNN  "a.asp_tdnn.%s"
+
+// qwen3tts code_predictor
+#define TN_A_GEN_CODE_PROJ_IN  "a.gen.code.proj_in.%s"
+#define TN_A_GEN_CODE_EMBD     "a.gen.code.embd.%s"
+#define TN_A_GEN_CODE_HEAD     "a.gen.code.head.%s"
+#define TN_A_GEN_CODE_OUT_EMBD "a.gen.code.out_embd.%s"
+#define TN_A_GEN_CODE_NORM     "a.gen.code.output_norm.%s"
+
+// qwen3tts code2wav (RVQ codes -> raw PCM)
+// pre_transformer layers use the generic TN_ATTN_*/TN_FFN_*/TN_LN_*/TN_LS_* macros, prefix "a.gen.wav.tfm"
+#define TN_A_GEN_WAV_QUANT_FIRST_IN  "a.gen.wav.quant.first.in_proj.%s"
+#define TN_A_GEN_WAV_QUANT_FIRST_OUT "a.gen.wav.quant.first.out_proj.%s"
+#define TN_A_GEN_WAV_QUANT_FIRST_CB  "a.gen.wav.quant.first.codebook.%s"
+#define TN_A_GEN_WAV_QUANT_REST_IN   "a.gen.wav.quant.rest.in_proj.%s"
+#define TN_A_GEN_WAV_QUANT_REST_OUT  "a.gen.wav.quant.rest.out_proj.%s"
+#define TN_A_GEN_WAV_QUANT_REST_CB   "a.gen.wav.quant.rest.codebook.%s"
+#define TN_A_GEN_WAV_PRE_CONV        "a.gen.wav.pre_conv.%s"
+#define TN_A_GEN_WAV_TFM_IN_PROJ     "a.gen.wav.tfm.in_proj.%s"
+#define TN_A_GEN_WAV_TFM_OUT_PROJ    "a.gen.wav.tfm.out_proj.%s"
+#define TN_A_GEN_WAV_TFM_OUT_NORM    "a.gen.wav.tfm.output_norm.%s"
+#define TN_A_GEN_WAV_UP_CONV         "a.gen.wav.up.blk.%d.conv.%s"
+#define TN_A_GEN_WAV_UP_DWCONV       "a.gen.wav.up.blk.%d.dwconv.%s"
+#define TN_A_GEN_WAV_UP_NORM         "a.gen.wav.up.blk.%d.norm.%s"
+#define TN_A_GEN_WAV_UP_PW1          "a.gen.wav.up.blk.%d.pw1.%s"
+#define TN_A_GEN_WAV_UP_PW2          "a.gen.wav.up.blk.%d.pw2.%s"
+#define TN_A_GEN_WAV_UP_GAMMA        "a.gen.wav.up.blk.%d.gamma"
+#define TN_A_GEN_WAV_DAC_ENTRY       "a.gen.wav.dac.entry.%s"
+#define TN_A_GEN_WAV_DAC_SNAKE       "a.gen.wav.dac.blk.%d.snake.%s"
+#define TN_A_GEN_WAV_DAC_CONV        "a.gen.wav.dac.blk.%d.conv.%s"
+#define TN_A_GEN_WAV_DAC_RES_ACT1    "a.gen.wav.dac.blk.%d.res.%d.act1.%s"
+#define TN_A_GEN_WAV_DAC_RES_CONV1   "a.gen.wav.dac.blk.%d.res.%d.conv1.%s"
+#define TN_A_GEN_WAV_DAC_RES_ACT2    "a.gen.wav.dac.blk.%d.res.%d.act2.%s"
+#define TN_A_GEN_WAV_DAC_RES_CONV2   "a.gen.wav.dac.blk.%d.res.%d.conv2.%s"
+#define TN_A_GEN_WAV_DAC_POST_SNAKE  "a.gen.wav.dac.post_snake.%s"
+#define TN_A_GEN_WAV_DAC_POST_CONV   "a.gen.wav.dac.post_conv.%s"
+
 // cogvlm
 #define TN_MM_POST_FC_NORM "mm.post_fc_norm.%s"
 #define TN_MM_H_TO_4H      "mm.up.%s"
@@ -408,6 +453,8 @@ enum projector_type {
     PROJECTOR_TYPE_MINIMAX_M3,
     PROJECTOR_TYPE_GRANITE4_VISION,
     PROJECTOR_TYPE_MIMO_AUDIO,
+    PROJECTOR_TYPE_QWEN3TTS_SPKENC,
+    PROJECTOR_TYPE_QWEN3TTS_GEN,
     PROJECTOR_TYPE_UNKNOWN,
 };
 
@@ -465,6 +512,8 @@ static std::map<projector_type, std::string> PROJECTOR_TYPE_NAMES = {
     { PROJECTOR_TYPE_GRANITE4_VISION,   "granite4_vision"},
     { PROJECTOR_TYPE_MIMO_AUDIO,        "mimo_audio"},
     { PROJECTOR_TYPE_PARAKEET,          "parakeet"},
+    { PROJECTOR_TYPE_QWEN3TTS_SPKENC,   "qwen3tts_spkenc"},
+    { PROJECTOR_TYPE_QWEN3TTS_GEN,      "qwen3tts_gen"},
 };
 
 static projector_type clip_projector_type_from_string(const std::string & str) {
index 8b9db5101d2ce04a6d75b5d9ddaef9bcfd422384..101f49cd1849d1af1da705e7ef7377e3c2bfe763 100644 (file)
@@ -136,6 +136,19 @@ struct clip_hparams {
     int32_t rvq_num_quantizers = 0;
     std::vector<int32_t> rvq_codebook_size; // per-quantizer bin count (ragged, e.g. 1024/1024/256/128x17)
 
+    // qwen3tts code2wav
+    int32_t wav_tfm_n_layer      = 0;
+    int32_t wav_tfm_n_embd       = 0;
+    int32_t wav_tfm_n_ff         = 0;
+    int32_t wav_tfm_n_head       = 0;
+    int32_t wav_tfm_n_head_kv    = 0;
+    float   wav_tfm_eps          = 1e-5f;
+    float   wav_tfm_rope_theta   = 10000.0f;
+    int32_t wav_upsample_n_block = 0;
+    int32_t wav_dac_n_block      = 0;
+    int32_t wav_dac_n_res        = 0;
+    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;
     int32_t audio_local_group_size = 0;
@@ -286,6 +299,14 @@ struct clip_layer {
     ggml_tensor * cross_attn_norm_w = nullptr;
     ggml_tensor * cross_attn_norm_b = nullptr;
 
+    // qwen3tts speaker encoder: SE-Res2Net block, tdnn1/tdnn2 reuse conv_pw1_w/b and conv_pw2_w/b above
+    ggml_tensor * se_conv1_w = nullptr;
+    ggml_tensor * se_conv1_b = nullptr;
+    ggml_tensor * se_conv2_w = nullptr;
+    ggml_tensor * se_conv2_b = nullptr;
+    std::vector<ggml_tensor *> res2_conv_w; // Res2Net hierarchical branches
+    std::vector<ggml_tensor *> res2_conv_b;
+
     bool has_deepstack() const {
         return deepstack_fc1_w != nullptr;
     }
@@ -365,6 +386,73 @@ struct qf_block {
     std::vector<clip_layer> qf_proj_layers;
 };
 
+// qwen3tts code2wav: RVQ codes -> raw PCM
+struct clip_code2wav {
+    // "upsample" stage: one ConvNeXt block plus the causal ConvTranspose1d before it
+    struct upsample_block {
+        ggml_tensor * conv_w   = nullptr; // causal ConvTranspose1d, 2x
+        ggml_tensor * conv_b   = nullptr;
+        ggml_tensor * dwconv_w = nullptr; // depthwise causal conv, k=7
+        ggml_tensor * dwconv_b = nullptr;
+        ggml_tensor * norm_w   = nullptr; // LayerNorm
+        ggml_tensor * norm_b   = nullptr;
+        ggml_tensor * pw1_w    = nullptr; // pointwise expand
+        ggml_tensor * pw1_b    = nullptr;
+        ggml_tensor * pw2_w    = nullptr; // pointwise project
+        ggml_tensor * pw2_b    = nullptr;
+        ggml_tensor * gamma    = nullptr; // layer scale
+    };
+
+    // one DAC residual unit: SnakeBeta -> dilated causal conv -> SnakeBeta -> pointwise causal conv
+    struct dac_res {
+        ggml_tensor * act1_alpha = nullptr;
+        ggml_tensor * act1_beta  = nullptr;
+        ggml_tensor * conv1_w    = nullptr;
+        ggml_tensor * conv1_b    = nullptr;
+        ggml_tensor * act2_alpha = nullptr;
+        ggml_tensor * act2_beta  = nullptr;
+        ggml_tensor * conv2_w    = nullptr;
+        ggml_tensor * conv2_b    = nullptr;
+    };
+
+    // one DAC upsample block (SnakeBeta -> causal ConvTranspose1d -> 3 residual units)
+    struct dac_block {
+        ggml_tensor * snake_alpha = nullptr;
+        ggml_tensor * snake_beta  = nullptr;
+        ggml_tensor * conv_w      = nullptr; // causal ConvTranspose1d
+        ggml_tensor * conv_b      = nullptr;
+        std::vector<dac_res> res;
+    };
+
+    // quantizer: RVQ codebook decode
+    ggml_tensor * quant_first_in_w  = nullptr; // semantic RVQ, in_proj (1x1 conv, loaded as 2D)
+    ggml_tensor * quant_first_out_w = nullptr;
+    ggml_tensor * quant_first_cb_w  = nullptr; // codebook (1 layer)
+    ggml_tensor * quant_rest_in_w   = nullptr; // acoustic RVQ
+    ggml_tensor * quant_rest_out_w  = nullptr;
+    ggml_tensor * quant_rest_cb_w   = nullptr; // codebooks, merged 3D [15, vocab, dim]
+
+    ggml_tensor * pre_conv_w = nullptr;
+    ggml_tensor * pre_conv_b = nullptr;
+
+    ggml_tensor * tfm_in_proj_w     = nullptr;
+    ggml_tensor * tfm_in_proj_b     = nullptr;
+    ggml_tensor * tfm_out_proj_w    = nullptr;
+    ggml_tensor * tfm_out_proj_b    = nullptr;
+    ggml_tensor * tfm_output_norm_w = nullptr;
+    std::vector<clip_layer> tfm_layers; // reuses the generic block fields (ln_1/attn/ln_2/ffn/ls_1/ls_2)
+
+    std::vector<upsample_block> upsample;
+
+    ggml_tensor * dac_entry_w = nullptr;
+    ggml_tensor * dac_entry_b = nullptr;
+    std::vector<dac_block> dac;
+    ggml_tensor * dac_post_snake_alpha = nullptr;
+    ggml_tensor * dac_post_snake_beta  = nullptr;
+    ggml_tensor * dac_post_conv_w      = nullptr;
+    ggml_tensor * dac_post_conv_b      = nullptr;
+};
+
 struct clip_model {
     clip_modality modality = CLIP_MODALITY_VISION;
     projector_type proj_type = PROJECTOR_TYPE_MLP;
@@ -577,6 +665,24 @@ struct clip_model {
     ggml_tensor * conv2d_3_w = nullptr;
     ggml_tensor * conv2d_3_b = nullptr;
 
+    // qwen3tts speaker encoder (ECAPA-TDNN)
+    // reused tensors: stem conv is conv1d_1_w/b, feature aggregation is conv_out_w/b, output proj is mm_fc_w/b
+    ggml_tensor * spk_asp_attn_w = nullptr;
+    ggml_tensor * spk_asp_attn_b = nullptr;
+    ggml_tensor * spk_asp_tdnn_w = nullptr;
+    ggml_tensor * spk_asp_tdnn_b = nullptr;
+
+    // qwen3tts code_predictor
+    ggml_tensor * gen_code_proj_in_w  = nullptr; // small_to_mtp_projection
+    ggml_tensor * gen_code_proj_in_b  = nullptr;
+    ggml_tensor * gen_code_embd_w     = nullptr; // per-codebook embedding, merged 3D
+    ggml_tensor * gen_code_head_w     = nullptr; // per-codebook output head, merged 3D
+    ggml_tensor * gen_code_out_embd_w = nullptr; // codebook-0 embedding, fed back into the talker
+    ggml_tensor * gen_code_norm_w     = nullptr; // final norm
+
+    // qwen3tts code2wav: RVQ codes -> raw PCM
+    clip_code2wav c2w;
+
     // cogvlm
     ggml_tensor * mm_post_fc_norm_w = nullptr;
     ggml_tensor * mm_post_fc_norm_b = nullptr;
index c1870813fb932a64a21fd093ab4be90659420b67..d6670030ff79fd67550389c1e94678ab25606f97 100644 (file)
@@ -17,6 +17,7 @@
 #include <cstring>
 #include <fstream>
 #include <map>
+#include <random>
 #include <stdexcept>
 #include <unordered_set>
 #include <vector>
@@ -269,6 +270,29 @@ clip_graph::clip_graph(clip_ctx * ctx, const clip_image_f32 & img) :
     gf = ggml_new_graph_custom(ctx0, ctx->max_nodes, false);
 }
 
+clip_graph::clip_graph(const clip_graph & parent) :
+        model(parent.model),
+        hparams(parent.hparams),
+        proj_type(parent.proj_type),
+        img(parent.img),
+        patch_size(parent.patch_size),
+        n_patches_x(parent.n_patches_x),
+        n_patches_y(parent.n_patches_y),
+        n_patches(parent.n_patches),
+        n_embd(parent.n_embd),
+        n_head(parent.n_head),
+        n_head_kv(parent.n_head_kv),
+        d_head(parent.d_head),
+        n_layer(parent.n_layer),
+        n_mmproj_embd(parent.n_mmproj_embd),
+        eps(parent.eps),
+        kq_scale(parent.kq_scale),
+        flash_attn_type(parent.flash_attn_type) {
+    // reuse from parent
+    ctx0 = parent.ctx0;
+    gf   = parent.gf;
+}
+
 ggml_tensor * clip_graph::build_mm(ggml_tensor * w, ggml_tensor * x) const {
     return ggml_mul_mat(ctx0, w, x);
 }
@@ -873,7 +897,8 @@ ggml_tensor * clip_graph::build_patch_merge_permute(ggml_tensor * cur, int scale
     return cur;
 }
 
-static std::unique_ptr<clip_graph> clip_get_graph_builder(clip_ctx * ctx, const clip_image_f32_batch & imgs) {
+static std::unique_ptr<clip_graph> clip_get_graph_builder(clip_ctx * ctx, const clip_image_f32_batch & imgs,
+                                                            const clip_encode_params * params = nullptr) {
     const clip_image_f32 & img = imgs.entries[0];
     std::unique_ptr<clip_graph> builder;
 
@@ -1025,6 +1050,17 @@ static std::unique_ptr<clip_graph> clip_get_graph_builder(clip_ctx * ctx, const
             {
                 builder = std::make_unique<clip_graph_mimo_audio>(ctx, img);
             } break;
+        case PROJECTOR_TYPE_QWEN3TTS_SPKENC:
+            {
+                builder = std::make_unique<clip_graph_qwen3tts_spkenc>(ctx, img);
+            } break;
+        case PROJECTOR_TYPE_QWEN3TTS_GEN:
+            {
+                const auto  gen_process = params ? params->gen_process : CLIP_GEN_PROCESS_GEN_CODE;
+                const int   top_k = params ? params->top_k : 50;
+                const float top_p = params ? params->top_p : 1.0f;
+                builder = std::make_unique<clip_graph_qwen3tts_gen>(ctx, img, gen_process, top_k, top_p);
+            } break;
         case PROJECTOR_TYPE_YOUTUVL:
             {
                 builder = std::make_unique<clip_graph_youtuvl>(ctx, img);
@@ -1065,8 +1101,9 @@ struct clip_model_loader {
 
     size_t model_size = 0; // in bytes
 
-    bool has_vision = false;
-    bool has_audio  = false;
+    bool has_vision    = false;
+    bool has_audio     = false;
+    bool has_gen_audio = false;
 
     mtmd_progress_callback progress_callback = nullptr;
     void * progress_callback_user_data = nullptr;
@@ -1112,8 +1149,9 @@ struct clip_model_loader {
 
         // modalities
         {
-            get_bool(KEY_HAS_VISION_ENC, has_vision, false);
-            get_bool(KEY_HAS_AUDIO_ENC,  has_audio,  false);
+            get_bool(KEY_HAS_VISION_ENC,    has_vision,    false);
+            get_bool(KEY_HAS_AUDIO_ENC,     has_audio,     false);
+            get_bool(KEY_HAS_GEN_AUDIO_ENC, has_gen_audio, false);
 
             if (has_vision) {
                 LOG_INF("%s: has vision encoder\n", __func__);
@@ -1121,6 +1159,9 @@ struct clip_model_loader {
             if (has_audio) {
                 LOG_INF("%s: has audio encoder\n", __func__);
             }
+            if (has_gen_audio) {
+                LOG_INF("%s: has audio generation (gen) encoder\n", __func__);
+            }
         }
 
         // tensors
@@ -1147,6 +1188,8 @@ struct clip_model_loader {
             GGML_ASSERT(has_vision);
         } else if (modality == CLIP_MODALITY_AUDIO) {
             GGML_ASSERT(has_audio);
+        } else if (modality == CLIP_MODALITY_GEN_AUDIO) {
+            GGML_ASSERT(has_gen_audio);
         }
         model.modality = modality;
 
@@ -1163,6 +1206,8 @@ struct clip_model_loader {
                     get_string(KEY_VISION_PROJ_TYPE, proj_type, false);
                 } else if (modality == CLIP_MODALITY_AUDIO) {
                     get_string(KEY_AUDIO_PROJ_TYPE, proj_type, false);
+                } else if (modality == CLIP_MODALITY_GEN_AUDIO) {
+                    get_string(KEY_GEN_AUDIO_PROJ_TYPE, proj_type, false);
                 } else {
                     GGML_ABORT("unknown modality");
                 }
@@ -1182,12 +1227,13 @@ struct clip_model_loader {
             }
         }
 
-        const bool is_vision = model.modality == CLIP_MODALITY_VISION;
-        const bool is_audio  = model.modality == CLIP_MODALITY_AUDIO;
+        const bool is_vision    = model.modality == CLIP_MODALITY_VISION;
+        const bool is_audio     = model.modality == CLIP_MODALITY_AUDIO;
+        const bool is_gen_audio = model.modality == CLIP_MODALITY_GEN_AUDIO;
 
         // other hparams
         {
-            const char * prefix = is_vision ? "vision" : "audio";
+            const char * prefix = is_vision ? "vision" : (is_audio ? "audio" : "gen.audio");
             get_u32(string_format(KEY_N_EMBD,         prefix), hparams.n_embd);
             get_u32(string_format(KEY_N_HEAD,         prefix), hparams.n_head);
             get_u32(string_format(KEY_N_EMBD_HEAD,    prefix), hparams.n_embd_head, false);
@@ -1198,6 +1244,7 @@ struct clip_model_loader {
 
             // n_head_kv is optional (for GQA), default to n_head
             hparams.n_head_kv = hparams.n_head;
+            get_u32(string_format(KEY_N_HEAD_KV, prefix), hparams.n_head_kv, false);
 
             if (is_vision) {
                 get_u32(KEY_IMAGE_SIZE, hparams.image_size);
@@ -1226,6 +1273,11 @@ struct clip_model_loader {
                 hparams.image_size = 0;
                 hparams.patch_size = 1;
 
+            } else if (is_gen_audio) {
+                // these are unused, but still need to be set to avoid issues
+                hparams.image_size = 0;
+                hparams.patch_size = 1;
+
             } else {
                 GGML_ASSERT(false && "unknown modality");
             }
@@ -1647,6 +1699,33 @@ struct clip_model_loader {
                                 "%s: mimo_audio: %s must be > 0\n", __func__, KEY_A_LOCAL_GROUP_SIZE));
                         }
                     } break;
+                case PROJECTOR_TYPE_QWEN3TTS_SPKENC:
+                    {
+                        // ECAPA-TDNN speaker encoder, mel front-end uses the Slaney default (fmin=0, fmax=sr/2)
+                        hparams.audio_sample_rate = 24000;
+                        hparams.audio_n_fft       = 1024;
+                        hparams.audio_window_len  = 1024;
+                        hparams.audio_hop_len     = 256;
+                    } break;
+                case PROJECTOR_TYPE_QWEN3TTS_GEN:
+                    {
+                        // TODO: hardcoded for now, read from code_predictor_config instead
+                        hparams.rope_theta = 1000000.0f;
+
+                        // code2wav params
+                        hparams.wav_tfm_n_layer      = 8;
+                        hparams.wav_tfm_n_embd       = 512;
+                        hparams.wav_tfm_n_ff         = 1024;
+                        hparams.wav_tfm_n_head       = 16;
+                        hparams.wav_tfm_n_head_kv    = 16;
+                        hparams.wav_tfm_eps          = 1e-5f;
+                        hparams.wav_tfm_rope_theta   = 10000.0f;
+                        hparams.wav_upsample_n_block = 2;
+                        hparams.wav_dac_n_block      = 4;
+                        hparams.wav_dac_n_res        = 3;
+                        // matches the reference decoder's sliding_window (speech_tokenizer/config.json)
+                        hparams.wav_tfm_swa = 72;
+                    } break;
                 case PROJECTOR_TYPE_PADDLEOCR:
                     {
                         hparams.n_merge = 2;
@@ -1871,7 +1950,9 @@ struct clip_model_loader {
         }
 
         // TODO @ngxson : support both audio and video in the future
-        const char * prefix = model.modality == CLIP_MODALITY_AUDIO ? "a" : "v";
+        const char * prefix = model.modality == CLIP_MODALITY_AUDIO ? "a"
+                             : model.modality == CLIP_MODALITY_GEN_AUDIO ? "a.gen.code"
+                             : "v";
 
         // get offsets
         for (int64_t i = 0; i < gguf_get_n_tensors(ctx_gguf.get()); ++i) {
@@ -1973,7 +2054,8 @@ struct clip_model_loader {
         model.position_embeddings = get_tensor(string_format(TN_POS_EMBD, prefix), false);
 
         const bool has_standard_layers = (
-            model.proj_type != PROJECTOR_TYPE_GEMMA3NV);
+            model.proj_type != PROJECTOR_TYPE_GEMMA3NV &&
+            model.proj_type != PROJECTOR_TYPE_QWEN3TTS_SPKENC);
 
         // layers
         const int n_layers_to_load = has_standard_layers ? hparams.n_layer : 0;
@@ -2599,6 +2681,144 @@ struct clip_model_loader {
                     model.mm_1_w = get_tensor(string_format(TN_MM_AUDIO_MLP, 1, "weight"));
                     model.mm_2_w = get_tensor(string_format(TN_MM_AUDIO_MLP, 2, "weight"));
                 } break;
+            case PROJECTOR_TYPE_QWEN3TTS_SPKENC:
+                {
+                    // stem TDNN (block 0)
+                    model.conv1d_1_w = get_tensor(string_format(TN_CONV1D, 0, "weight"));
+                    model.conv1d_1_b = get_tensor(string_format(TN_CONV1D, 0, "bias"));
+
+                    // SE-Res2Net blocks (GGUF bid 1..3, one per hparams.n_layer)
+                    model.layers.resize(hparams.n_layer);
+                    for (int il = 0; il < hparams.n_layer; il++) {
+                        auto & layer = model.layers[il];
+                        int bid = il + 1;
+                        layer.conv_pw1_w = get_tensor(string_format(TN_CONV_PW1, prefix, bid, "weight"));
+                        layer.conv_pw1_b = get_tensor(string_format(TN_CONV_PW1, prefix, bid, "bias"));
+                        layer.conv_pw2_w = get_tensor(string_format(TN_CONV_PW2, prefix, bid, "weight"));
+                        layer.conv_pw2_b = get_tensor(string_format(TN_CONV_PW2, prefix, bid, "bias"));
+                        layer.se_conv1_w = get_tensor(string_format(TN_A_SE_CONV1, bid, "weight"));
+                        layer.se_conv1_b = get_tensor(string_format(TN_A_SE_CONV1, bid, "bias"));
+                        layer.se_conv2_w = get_tensor(string_format(TN_A_SE_CONV2, bid, "weight"));
+                        layer.se_conv2_b = get_tensor(string_format(TN_A_SE_CONV2, bid, "bias"));
+                        layer.res2_conv_w.resize(7);
+                        layer.res2_conv_b.resize(7);
+                        for (int xid = 0; xid < 7; xid++) {
+                            layer.res2_conv_w[xid] = get_tensor(string_format(TN_A_CONV_RES2, bid, xid, "weight"));
+                            layer.res2_conv_b[xid] = get_tensor(string_format(TN_A_CONV_RES2, bid, xid, "bias"));
+                        }
+                    }
+
+                    // multi-layer feature aggregation
+                    model.conv_out_w = get_tensor(string_format(TN_CONV_OUT, "weight"));
+                    model.conv_out_b = get_tensor(string_format(TN_CONV_OUT, "bias"));
+
+                    // attentive statistics pooling
+                    model.spk_asp_attn_w = get_tensor(string_format(TN_A_ASP_ATTN, "weight"));
+                    model.spk_asp_attn_b = get_tensor(string_format(TN_A_ASP_ATTN, "bias"));
+                    model.spk_asp_tdnn_w = get_tensor(string_format(TN_A_ASP_TDNN, "weight"));
+                    model.spk_asp_tdnn_b = get_tensor(string_format(TN_A_ASP_TDNN, "bias"));
+
+                    // final speaker embedding projection
+                    model.mm_fc_w = get_tensor(string_format(TN_MM_AUDIO_FC, "weight"));
+                    model.mm_fc_b = get_tensor(string_format(TN_MM_AUDIO_FC, "bias"));
+                } break;
+            case PROJECTOR_TYPE_QWEN3TTS_GEN:
+                {
+                    // code_predictor
+                    model.gen_code_proj_in_w = get_tensor(string_format(TN_A_GEN_CODE_PROJ_IN, "weight"));
+                    model.gen_code_proj_in_b = get_tensor(string_format(TN_A_GEN_CODE_PROJ_IN, "bias"));
+                    model.gen_code_embd_w     = get_tensor(string_format(TN_A_GEN_CODE_EMBD,     "weight"));
+                    model.gen_code_head_w     = get_tensor(string_format(TN_A_GEN_CODE_HEAD,     "weight"));
+                    model.gen_code_out_embd_w = get_tensor(string_format(TN_A_GEN_CODE_OUT_EMBD, "weight"));
+                    model.gen_code_norm_w     = get_tensor(string_format(TN_A_GEN_CODE_NORM,     "weight"));
+
+                    // code2wav: RVQ codes -> raw PCM, lives in the same ctx as code_predictor
+                    {
+                        auto & c2w = model.c2w;
+
+                        c2w.quant_first_in_w  = get_tensor(string_format(TN_A_GEN_WAV_QUANT_FIRST_IN,  "weight"));
+                        c2w.quant_first_out_w = get_tensor(string_format(TN_A_GEN_WAV_QUANT_FIRST_OUT, "weight"));
+                        c2w.quant_first_cb_w  = get_tensor(string_format(TN_A_GEN_WAV_QUANT_FIRST_CB,  "weight"));
+                        c2w.quant_rest_in_w   = get_tensor(string_format(TN_A_GEN_WAV_QUANT_REST_IN,   "weight"));
+                        c2w.quant_rest_out_w  = get_tensor(string_format(TN_A_GEN_WAV_QUANT_REST_OUT,  "weight"));
+                        c2w.quant_rest_cb_w   = get_tensor(string_format(TN_A_GEN_WAV_QUANT_REST_CB,   "weight"));
+
+                        c2w.pre_conv_w = get_tensor(string_format(TN_A_GEN_WAV_PRE_CONV, "weight"));
+                        c2w.pre_conv_b = get_tensor(string_format(TN_A_GEN_WAV_PRE_CONV, "bias"));
+
+                        c2w.tfm_in_proj_w     = get_tensor(string_format(TN_A_GEN_WAV_TFM_IN_PROJ,  "weight"));
+                        c2w.tfm_in_proj_b     = get_tensor(string_format(TN_A_GEN_WAV_TFM_IN_PROJ,  "bias"));
+                        c2w.tfm_out_proj_w    = get_tensor(string_format(TN_A_GEN_WAV_TFM_OUT_PROJ, "weight"));
+                        c2w.tfm_out_proj_b    = get_tensor(string_format(TN_A_GEN_WAV_TFM_OUT_PROJ, "bias"));
+                        c2w.tfm_output_norm_w = get_tensor(string_format(TN_A_GEN_WAV_TFM_OUT_NORM, "weight"));
+
+                        // loaded manually, the generic model.layers loop is taken by code_predictor
+                        c2w.tfm_layers.resize(hparams.wav_tfm_n_layer);
+                        for (int il = 0; il < hparams.wav_tfm_n_layer; il++) {
+                            auto & layer = c2w.tfm_layers[il];
+                            const char * p = "a.gen.wav.tfm";
+                            layer.q_w      = get_tensor(string_format(TN_ATTN_Q,      p, il, "weight"));
+                            layer.k_w      = get_tensor(string_format(TN_ATTN_K,      p, il, "weight"));
+                            layer.v_w      = get_tensor(string_format(TN_ATTN_V,      p, il, "weight"));
+                            layer.o_w      = get_tensor(string_format(TN_ATTN_OUTPUT, p, il, "weight"));
+                            layer.ln_1_w   = get_tensor(string_format(TN_LN_1,        p, il, "weight"));
+                            layer.ln_2_w   = get_tensor(string_format(TN_LN_2,        p, il, "weight"));
+                            layer.ls_1_w   = get_tensor(string_format(TN_LS_1,        p, il, "weight"));
+                            layer.ls_2_w   = get_tensor(string_format(TN_LS_2,        p, il, "weight"));
+                            layer.ff_gate_w = get_tensor(string_format(TN_FFN_GATE,   p, il, "weight"));
+                            layer.ff_up_w   = get_tensor(string_format(TN_FFN_UP,     p, il, "weight"));
+                            layer.ff_down_w = get_tensor(string_format(TN_FFN_DOWN,   p, il, "weight"));
+                        }
+
+                        // upsample: 2x (causal ConvTranspose1d + ConvNeXt block)
+                        c2w.upsample.resize(hparams.wav_upsample_n_block);
+                        for (int il = 0; il < hparams.wav_upsample_n_block; il++) {
+                            auto & up = c2w.upsample[il];
+                            up.conv_w   = get_tensor(string_format(TN_A_GEN_WAV_UP_CONV,   il, "weight"));
+                            up.conv_b   = get_tensor(string_format(TN_A_GEN_WAV_UP_CONV,   il, "bias"));
+                            up.dwconv_w = get_tensor(string_format(TN_A_GEN_WAV_UP_DWCONV, il, "weight"));
+                            up.dwconv_b = get_tensor(string_format(TN_A_GEN_WAV_UP_DWCONV, il, "bias"));
+                            up.norm_w   = get_tensor(string_format(TN_A_GEN_WAV_UP_NORM,   il, "weight"));
+                            up.norm_b   = get_tensor(string_format(TN_A_GEN_WAV_UP_NORM,   il, "bias"));
+                            up.pw1_w    = get_tensor(string_format(TN_A_GEN_WAV_UP_PW1,    il, "weight"));
+                            up.pw1_b    = get_tensor(string_format(TN_A_GEN_WAV_UP_PW1,    il, "bias"));
+                            up.pw2_w    = get_tensor(string_format(TN_A_GEN_WAV_UP_PW2,    il, "weight"));
+                            up.pw2_b    = get_tensor(string_format(TN_A_GEN_WAV_UP_PW2,    il, "bias"));
+                            up.gamma    = get_tensor(string_format(TN_A_GEN_WAV_UP_GAMMA,  il));
+                        }
+
+                        // DAC decoder: conv_pre + n upsample blocks (each with n_res residual units) + conv_post
+                        c2w.dac_entry_w = get_tensor(string_format(TN_A_GEN_WAV_DAC_ENTRY, "weight"));
+                        c2w.dac_entry_b = get_tensor(string_format(TN_A_GEN_WAV_DAC_ENTRY, "bias"));
+
+                        c2w.dac.resize(hparams.wav_dac_n_block);
+                        for (int il = 0; il < hparams.wav_dac_n_block; il++) {
+                            auto & blk = c2w.dac[il];
+                            blk.snake_alpha = get_tensor(string_format(TN_A_GEN_WAV_DAC_SNAKE, il, "alpha"));
+                            blk.snake_beta  = get_tensor(string_format(TN_A_GEN_WAV_DAC_SNAKE, il, "beta"));
+                            blk.conv_w      = get_tensor(string_format(TN_A_GEN_WAV_DAC_CONV,  il, "weight"));
+                            blk.conv_b      = get_tensor(string_format(TN_A_GEN_WAV_DAC_CONV,  il, "bias"));
+
+                            blk.res.resize(hparams.wav_dac_n_res);
+                            for (int ir = 0; ir < hparams.wav_dac_n_res; ir++) {
+                                auto & res = blk.res[ir];
+                                res.act1_alpha = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_ACT1,  il, ir, "alpha"));
+                                res.act1_beta  = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_ACT1,  il, ir, "beta"));
+                                res.conv1_w    = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_CONV1, il, ir, "weight"));
+                                res.conv1_b    = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_CONV1, il, ir, "bias"));
+                                res.act2_alpha = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_ACT2,  il, ir, "alpha"));
+                                res.act2_beta  = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_ACT2,  il, ir, "beta"));
+                                res.conv2_w    = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_CONV2, il, ir, "weight"));
+                                res.conv2_b    = get_tensor(string_format(TN_A_GEN_WAV_DAC_RES_CONV2, il, ir, "bias"));
+                            }
+                        }
+
+                        c2w.dac_post_snake_alpha = get_tensor(string_format(TN_A_GEN_WAV_DAC_POST_SNAKE, "alpha"));
+                        c2w.dac_post_snake_beta  = get_tensor(string_format(TN_A_GEN_WAV_DAC_POST_SNAKE, "beta"));
+                        c2w.dac_post_conv_w      = get_tensor(string_format(TN_A_GEN_WAV_DAC_POST_CONV,  "weight"));
+                        c2w.dac_post_conv_b      = get_tensor(string_format(TN_A_GEN_WAV_DAC_POST_CONV,  "bias"));
+                    }
+                } break;
             case PROJECTOR_TYPE_VOXTRAL:
                 {
                     model.conv1d_1_w = get_tensor(string_format(TN_CONV1D, 1, "weight"));
@@ -3427,6 +3647,7 @@ struct clip_model_loader {
 struct clip_init_result clip_init(const char * fname, struct clip_context_params ctx_params) {
     clip_ctx * ctx_vision = nullptr;
     clip_ctx * ctx_audio = nullptr;
+    clip_ctx * ctx_gen_audio = nullptr;
 
     try {
         clip_model_loader loader(fname,
@@ -3459,16 +3680,25 @@ struct clip_init_result clip_init(const char * fname, struct clip_context_params
             }
         }
 
+        if (loader.has_gen_audio) {
+            ctx_gen_audio = new clip_ctx(ctx_params);
+            loader.load_hparams(ctx_gen_audio->model, CLIP_MODALITY_GEN_AUDIO);
+            loader.load_tensors(*ctx_gen_audio);
+            // TODO: fix warmup
+            ctx_gen_audio->buf_compute_meta.resize(ctx_gen_audio->max_nodes * ggml_tensor_overhead() + ggml_graph_overhead());
+        }
+
     } catch (const std::exception & e) {
         LOG_ERR("%s: failed to load model '%s': %s\n", __func__, fname, e.what());
 
         delete ctx_vision;
         delete ctx_audio;
+        delete ctx_gen_audio;
 
-        return {nullptr, nullptr};
+        return {nullptr, nullptr, nullptr};
     }
 
-    return {ctx_vision, ctx_audio};
+    return {ctx_vision, ctx_audio, ctx_gen_audio};
 }
 
 struct clip_cap clip_get_cap(const char * fname) {
@@ -3784,6 +4014,16 @@ int clip_n_output_tokens(const clip_ctx * ctx, const clip_image_f32 * img) {
                 const int ds = ctx->model.hparams.audio_proj_downsample_rate;
                 n_patches = ((img->nx() + ws - 1) / ws) * (ws / ds);
             } break;
+        case PROJECTOR_TYPE_QWEN3TTS_SPKENC:
+            {
+                // pooling gives one speaker embedding, whatever the clip length is
+                n_patches = 1;
+            } break;
+        case PROJECTOR_TYPE_QWEN3TTS_GEN:
+            {
+                // one hidden-state vector fed back to the talker per call
+                n_patches = 1;
+            } break;
         case PROJECTOR_TYPE_GRANITE4_VISION:
             {
                 // Per-tile output token count: each projector block outputs
@@ -3817,7 +4057,16 @@ bool clip_image_encode(struct clip_ctx * ctx, int n_threads, const clip_image_f3
 }
 
 bool clip_image_batch_encode(clip_ctx * ctx, int n_threads, const clip_image_f32_batch * imgs_c_ptr, std::vector<float> & out_batch_embd) {
-    const clip_image_f32_batch & imgs = *imgs_c_ptr;
+    clip_encode_params params;
+    params.imgs = imgs_c_ptr;
+    params.n_threads = n_threads;
+    params.out_embd = &out_batch_embd;
+
+    return clip_encode(ctx, &params);
+}
+
+bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) {
+    const clip_image_f32_batch & imgs = *params->imgs;
     int n_batch_cur = imgs.entries.size();
 
     // [QWEN_VIDEO] for video models, the batch dimension is used as temporal dimension for merged frames
@@ -3828,12 +4077,12 @@ bool clip_image_batch_encode(clip_ctx * ctx, int n_threads, const clip_image_f32
 
     // if buffers are not allocated, we need to do a warmup run to allocate them
     if (!ctx->is_allocated) {
-        clip_model_loader::warmup(*ctx, *imgs_c_ptr);
+        clip_model_loader::warmup(*ctx, *params->imgs);
     }
 
     // build the inference graph
     ggml_backend_sched_reset(ctx->sched.get());
-    ggml_cgraph * gf = clip_get_graph_builder(ctx, imgs)->build();
+    ggml_cgraph * gf = clip_get_graph_builder(ctx, imgs, params)->build();
     ggml_backend_sched_alloc_graph(ctx->sched.get(), gf);
 
     // set inputs
@@ -3918,8 +4167,8 @@ bool clip_image_batch_encode(clip_ctx * ctx, int n_threads, const clip_image_f32
         }
         set_input_f32("inp_raw", inp_raw);
 
-    } else {
-        // audio input
+    } else if (!(ctx->proj_type() == PROJECTOR_TYPE_QWEN3TTS_GEN && params->gen_process == CLIP_GEN_PROCESS_GEN_WAV)) {
+        // audio input, code2wav is not here: its only input is "inp_codes", set in the switch below
         GGML_ASSERT(imgs.entries.size() == 1);
 
         const auto & mel_inp = imgs.entries[0];
@@ -4475,9 +4724,77 @@ bool clip_image_batch_encode(clip_ctx * ctx, int n_threads, const clip_image_f32
         case PROJECTOR_TYPE_COGVLM:
         case PROJECTOR_TYPE_YASA2:
         case PROJECTOR_TYPE_GEMMA4UA:
+        case PROJECTOR_TYPE_QWEN3TTS_SPKENC:
             {
                 // do nothing
             } break;
+        case PROJECTOR_TYPE_QWEN3TTS_GEN:
+            {
+                if (params->gen_process == CLIP_GEN_PROCESS_GEN_WAV) {
+                    GGML_ASSERT(params->codes != nullptr);
+
+                    // frame-major input to group-major, rear-padded with code 0 up to one window
+                    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;
+                    GGML_ASSERT(n_frames > 0 && n_frames <= n_frames_w);
+
+                    // codes are used as ggml_get_rows indices, so check them against the codebook vocab
+                    const int64_t vocab_first = model.c2w.quant_first_cb_w->ne[1];
+                    const int64_t vocab_rest  = model.c2w.quant_rest_cb_w->ne[1];
+                    for (int64_t f = 0; f < n_frames; f++) {
+                        for (int64_t g = 0; g < n_codes; g++) {
+                            const int32_t c = (*params->codes)[f * n_codes + g];
+                            const int64_t vocab = (g == 0) ? vocab_first : vocab_rest;
+                            if (c < 0 || (int64_t) c >= vocab) {
+                                LOG_ERR("%s: code out of range (frame %lld, group %lld, code %d, vocab %lld)\n",
+                                        __func__, (long long) f, (long long) g, c, (long long) vocab);
+                                return false;
+                            }
+                        }
+                    }
+
+                    std::vector<int32_t> codes(n_frames_w * n_codes, 0);
+                    for (int64_t f = 0; f < n_frames; f++) {
+                        for (int64_t g = 0; g < n_codes; g++) {
+                            codes[g * n_frames_w + f] = (*params->codes)[f * n_codes + g];
+                        }
+                    }
+                    set_input_i32("inp_codes", codes);
+
+                    // upload the state from the previous call, or zero-fill on a cold start
+                    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<uint8_t> zeros(nb, 0);
+                            ggml_backend_tensor_set(t, zeros.data(), 0, nb);
+                        }
+                        offset += nb;
+                    }
+                } else {
+                    // code0 indexes gen_code_out_embd_w via ggml_get_rows; bound it
+                    const int64_t vocab0 = model.gen_code_out_embd_w->ne[1];
+                    if (params->code0 < 0 || (int64_t) params->code0 >= vocab0) {
+                        LOG_ERR("%s: code0 out of range (%d, vocab %lld)\n", __func__, params->code0, (long long) vocab0);
+                        return false;
+                    }
+                    std::vector<int32_t> code0 = { params->code0 };
+                    set_input_i32("inp_code0", code0);
+
+                    // one uniform(0,1) draw per codebook, used by do_sampling()
+                    static std::mt19937 rng{ std::random_device{}() };
+                    std::uniform_real_distribution<float> dist(0.0f, 1.0f);
+                    const int64_t n_acoustic = model.gen_code_head_w->ne[2];
+                    for (int64_t g = 0; g < n_acoustic; g++) {
+                        std::vector<float> r = { dist(rng) };
+                        set_input_f32(("inp_rand_" + std::to_string(g)).c_str(), r);
+                    }
+                }
+            } break;
         case PROJECTOR_TYPE_HUNYUANVL:
             {
                 // Compute the HunyuanVL 2D position embedding on CPU (with the
@@ -4883,7 +5200,7 @@ bool clip_image_batch_encode(clip_ctx * ctx, int n_threads, const clip_image_f32
     if (reg) {
         auto ggml_backend_set_n_threads_fn = (ggml_backend_set_n_threads_t) ggml_backend_reg_get_proc_address(reg, "ggml_backend_set_n_threads");
         if (ggml_backend_set_n_threads_fn) {
-            ggml_backend_set_n_threads_fn(ctx->backend_cpu, n_threads);
+            ggml_backend_set_n_threads_fn(ctx->backend_cpu, params->n_threads);
         }
     }
 
@@ -4893,34 +5210,90 @@ bool clip_image_batch_encode(clip_ctx * ctx, int n_threads, const clip_image_f32
         return false;
     }
 
-    // the last node is the embedding tensor
-    ggml_tensor * embeddings = ggml_graph_node(gf, -1);
+    // the last node is the embedding tensor, code2wav has no out_embd
+    ggml_tensor * embeddings = params->out_embd ? ggml_graph_node(gf, -1) : nullptr;
+
+    if (embeddings != nullptr) {
+        // sanity check (assuming that all images in batch have the same number of tokens, so we only check the first one)
+        const int n_tokens_out = embeddings->ne[1];
+        const int expected_n_tokens_out = clip_n_output_tokens(ctx, &imgs.entries[0]);
+        if (n_tokens_out != expected_n_tokens_out) {
+            LOG_ERR("%s: expected output %d tokens, got %d\n", __func__, expected_n_tokens_out, n_tokens_out);
+            GGML_ABORT("Invalid number of output tokens");
+        }
+
+        LOG_DBG("%s: output embedding shape [%d, %d, %d]\n", __func__,
+            (int)embeddings->ne[0], (int)embeddings->ne[1], (int)embeddings->ne[2]);
 
-    // sanity check (assuming that all images in batch have the same number of tokens, so we only check the first one)
-    const int n_tokens_out = embeddings->ne[1];
-    const int expected_n_tokens_out = clip_n_output_tokens(ctx, &imgs.entries[0]);
-    if (n_tokens_out != expected_n_tokens_out) {
-        LOG_ERR("%s: expected output %d tokens, got %d\n", __func__, expected_n_tokens_out, n_tokens_out);
-        GGML_ABORT("Invalid number of output tokens");
+        // copy output to user buffer if provided
+        // if output is empty, skip the copy
+        auto & out_batch_embd = *params->out_embd;
+        if (!out_batch_embd.empty()) {
+            if (out_batch_embd.size() != (size_t)ggml_nelements(embeddings)) {
+                LOG_ERR("%s: output buffer has %zu elements but expected %zu\n", __func__, out_batch_embd.size(), (size_t)ggml_nelements(embeddings));
+                GGML_ABORT("Output buffer size mismatch");
+            }
+            ggml_backend_tensor_get(embeddings, out_batch_embd.data(), 0, ggml_nbytes(embeddings));
+        } else {
+            LOG_WRN("%s: output buffer is empty, skipping copy\n", __func__);
+        }
     }
 
-    LOG_DBG("%s: output embedding shape [%d, %d, %d]\n", __func__,
-        (int)embeddings->ne[0], (int)embeddings->ne[1], (int)embeddings->ne[2]);
+    //
+    // for audio gen models
+    //
 
-    // copy output to user buffer if provided
-    // if output is empty, skip the copy
-    if (!out_batch_embd.empty()) {
-        if (out_batch_embd.size() != (size_t)ggml_nelements(embeddings)) {
-            LOG_ERR("%s: output buffer has %zu elements but expected %zu\n", __func__, out_batch_embd.size(), (size_t)ggml_nelements(embeddings));
-            GGML_ABORT("Output buffer size mismatch");
+    if (params->out_codes != nullptr) {
+        ggml_tensor * codes = ggml_graph_get_tensor(gf, "out_codes");
+        if (codes == nullptr) {
+            GGML_ABORT("out_codes requested but graph has no \"out_codes\" tensor");
+        }
+        auto & out_codes = *params->out_codes;
+        out_codes.resize(ggml_nelements(codes));
+        ggml_backend_tensor_get(codes, out_codes.data(), 0, ggml_nbytes(codes));
+    }
+    if (params->out_audio != nullptr) {
+        ggml_tensor * audio = ggml_graph_get_tensor(gf, "out_audio");
+        if (audio == nullptr) {
+            GGML_ABORT("out_audio requested but graph has no \"out_audio\" tensor");
+        }
+        auto & out_audio = *params->out_audio;
+        out_audio.resize(ggml_nelements(audio));
+        ggml_backend_tensor_get(audio, out_audio.data(), 0, ggml_nbytes(audio));
+
+        // drop the tail audio that comes from the code-0 rear padding
+        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;
         }
-        ggml_backend_tensor_get(embeddings, out_batch_embd.data(), 0, ggml_nbytes(embeddings));
-    } else {
-        LOG_WRN("%s: output buffer is empty, skipping copy\n", __func__);
     }
 
+    //
     // Debug: dump final embeddings if MTMD_DEBUG_EMBEDDINGS is set
-    if (ctx->debug_output_embeddings) {
+    //
+
+    if (ctx->debug_output_embeddings && embeddings != nullptr) {
         const int64_t n_embd = embeddings->ne[0];
         const int64_t n_tokens = embeddings->ne[1];
         std::vector<float> emb_data(ggml_nelements(embeddings));
@@ -5047,6 +5420,10 @@ int clip_n_mmproj_embd(const struct clip_ctx * ctx) {
             return ctx->model.mm_ffn_down_w->ne[1];
         case PROJECTOR_TYPE_MIMO_AUDIO:
             return ctx->model.mm_2_w->ne[1];
+        case PROJECTOR_TYPE_QWEN3TTS_SPKENC:
+            return ctx->model.mm_fc_w->ne[2];
+        case PROJECTOR_TYPE_QWEN3TTS_GEN:
+            return ctx->model.gen_code_out_embd_w->ne[0];
         case PROJECTOR_TYPE_PARAKEET:
             return ctx->model.mm_1_w->ne[1];
         default:
index 967093a812d68fd3bb2325769d71b13cf443760c..7f706d976eb3dfa7eddbee43cb10b51c87295056 100644 (file)
@@ -37,6 +37,7 @@ struct clip_image_f32_batch;
 enum clip_modality {
     CLIP_MODALITY_VISION,
     CLIP_MODALITY_AUDIO,
+    CLIP_MODALITY_GEN_AUDIO,
 };
 
 enum clip_flash_attn_type {
@@ -61,6 +62,7 @@ struct clip_context_params {
 struct clip_init_result {
     struct clip_ctx * ctx_v; // vision context
     struct clip_ctx * ctx_a; // audio context
+    struct clip_ctx * ctx_gen_a; // audio generation context
 };
 
 struct clip_init_result clip_init(const char * fname, struct clip_context_params ctx_params);
@@ -84,6 +86,33 @@ int clip_n_mmproj_embd(const struct clip_ctx * ctx);
 bool clip_image_encode      (struct clip_ctx * ctx, int n_threads, const clip_image_f32 * img, std::vector<float> & out_vec);
 bool clip_image_batch_encode(struct clip_ctx * ctx, int n_threads, const struct clip_image_f32_batch * imgs, std::vector<float> & out_batch_embd);
 
+enum clip_gen_process_type {
+    CLIP_GEN_PROCESS_GEN_UNKNOWN,
+    CLIP_GEN_PROCESS_GEN_CODE, // h_state to codes
+    CLIP_GEN_PROCESS_GEN_WAV,  // codes to raw PCM audio
+};
+struct clip_encode_params {
+    int n_threads = 1;
+    const clip_image_f32_batch * imgs = nullptr;
+    std::vector<float> * out_embd = nullptr;
+
+    // for audio gen, imgs has exactly one entry: hidden state from backbone (GEN_CODE) or unused (GEN_WAV)
+    clip_gen_process_type gen_process = CLIP_GEN_PROCESS_GEN_UNKNOWN;
+
+    // GEN_CODE: out_embd receives the embd to feed back to the backbone
+    int32_t code0 = 0; // semantic code sampled by the backbone
+    int32_t top_k = 50;
+    float   top_p = 1.0f;
+    std::vector<int32_t> * out_codes = nullptr; // this frame's 16 sampled codes
+
+    // GEN_WAV
+    const std::vector<int32_t> * codes = nullptr;     // this frame's 16 RVQ codes
+    std::vector<float> * out_audio = nullptr;         // decoded PCM samples, F32
+    const std::vector<uint8_t> * state_in  = nullptr; // state from previous call, null or wrong size means cold start
+    std::vector<uint8_t> *       state_out = nullptr; // state for the next call
+};
+bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params);
+
 bool clip_is_llava(const struct clip_ctx * ctx);
 // note for contributor: this clip_is_(model) pattern is deprecated
 //                       do NOT add new functions like this
index e54366a086f4dd0e25ccda1ec26757e90b7982f5..eb924972bf5c01e93d918ae1a70e17b249ae264c 100644 (file)
@@ -2,6 +2,11 @@
 
 #include "../clip-graph.h"
 
+#include <map>
+#include <string>
+#include <utility>
+#include <vector>
+
 /*
  * 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.
@@ -215,6 +220,111 @@ struct clip_graph_mimo_audio : clip_graph {
     ggml_cgraph * build() override;
 };
 
+struct clip_graph_qwen3tts_spkenc : clip_graph {
+    clip_graph_qwen3tts_spkenc(clip_ctx * ctx, const clip_image_f32 & img) : clip_graph(ctx, img) {}
+    ggml_cgraph * build() override;
+
+    ggml_tensor * conv1d_same(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int dilation) const;
+    ggml_tensor * res2net(ggml_tensor * x, const clip_layer & layer, int dilation, int scale) const;
+    ggml_tensor * se_block(ggml_tensor * x, const clip_layer & layer) const;
+    ggml_tensor * se_res2net_block(ggml_tensor * x, const clip_layer & layer, int dilation, int scale) const;
+    ggml_tensor * attentive_stats_pool(ggml_tensor * x) const;
+};
+
+struct clip_graph_qwen3tts_gen : clip_graph {
+    clip_graph_qwen3tts_gen(clip_ctx * ctx, const clip_image_f32 & img, clip_gen_process_type gen_process, int top_k, float top_p)
+        : clip_graph(ctx, img), gen_process(gen_process), top_k(top_k), top_p(top_p) {}
+    ggml_cgraph * build() override;
+
+    // which sub-graph build() constructs, fixed at graph-build time
+    clip_gen_process_type gen_process;
+
+    // sampling params, fixed at graph-build time (GEN_CODE only)
+    int   top_k;
+    float top_p;
+
+    //
+    // code_gen: backbone hidden state + sampled code0 -> 16 RVQ codes
+    // MTP-style code predictor, one token per codebook
+    //
+    struct code_gen : clip_graph {
+        code_gen(const clip_graph & parent, int top_k, float top_p)
+            : clip_graph(parent), top_k(top_k), top_p(top_p) {}
+        ggml_cgraph * build() override { GGML_ABORT("call prefill()/step() instead"); }
+
+        int   top_k;
+        float top_p;
+
+        ggml_tensor * cache_set(ggml_tensor * cache, int row_idx, ggml_tensor * value) const;
+        ggml_tensor * do_sampling(ggml_tensor * logits, ggml_tensor * inp_rand) const;
+
+        ggml_tensor * const_i32(ggml_tensor * anchor, float value) const;
+        ggml_tensor * causal_mask_row(int64_t n_kv_pad, int pos) const;
+        ggml_tensor * project_in(ggml_tensor * cur) const;
+
+        ggml_tensor * layer_forward(
+                ggml_tensor * cur,
+                const clip_layer & layer,
+                ggml_tensor * inp_pos,
+                ggml_tensor * kq_mask,
+                ggml_tensor *& k_cache_layer,
+                ggml_tensor *& v_cache_layer,
+                int64_t n_kv_pad,
+                int pos,
+                int il) const;
+
+        void prefill(
+                std::vector<ggml_tensor *> & k_cache,
+                std::vector<ggml_tensor *> & v_cache,
+                ggml_tensor *& out_code_cache,
+                ggml_tensor * h_state,
+                ggml_tensor * code0_embd,
+                ggml_tensor * inp_rand) const;
+
+        ggml_tensor * step(
+                std::vector<ggml_tensor *> & k_cache,
+                std::vector<ggml_tensor *> & v_cache,
+                ggml_tensor * out_code_cache,
+                ggml_tensor * inp_rand,
+                int step_idx) const;
+    };
+
+    //
+    // code2wav: RVQ codes -> raw PCM (quantizer + pre_conv + pre_transformer + upsample + DAC).
+    //
+    struct code2wav : clip_graph {
+        code2wav(const clip_graph & parent) : clip_graph(parent) {}
+        ggml_cgraph * build() override { GGML_ABORT("call decode() instead"); }
+
+        // state_in: previous call's persisted state, by slot name (see list_c2w_state_slots())
+        std::map<std::string, ggml_tensor *> state_in;
+        // state_out: this call's state to persist, added to the graph outputs by build()
+        mutable std::vector<std::pair<std::string, ggml_tensor *>> state_out;
+
+        // stateful conv ops: read/update their state via state_in/state_out[state_name]
+        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, 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 [1, n_codes] I32 -> 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, see qwen3tts-gen.cpp
+struct c2w_state_slot {
+    std::string name;
+    int64_t     ne0;
+    int64_t     ne1;
+};
+std::vector<c2w_state_slot> 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
new file mode 100644 (file)
index 0000000..b6c95ef
--- /dev/null
@@ -0,0 +1,766 @@
+#include "models.h"
+
+#include <string>
+
+// on-device sampling: top-k, top-p, then a random draw
+ggml_tensor * clip_graph_qwen3tts_gen::code_gen::do_sampling(ggml_tensor * logits, ggml_tensor * inp_rand) const {
+    logits = ggml_reshape_1d(ctx0, logits, ggml_nelements(logits));
+    const int64_t n_vocab = logits->ne[0];
+
+    // sort a's rows by idx
+    auto sort_by = [this](ggml_tensor * a, ggml_tensor * idx) {
+        ggml_tensor * a2d = ggml_reshape_2d(ctx0, a, 1, a->ne[0]);
+        return ggml_reshape_1d(ctx0, ggml_get_rows(ctx0, a2d, idx), idx->ne[0]);
+    };
+
+    ggml_tensor * cur        = logits;
+    ggml_tensor * candidates = nullptr; // maps row index back to vocab id
+
+    if (top_k > 0 && top_k < n_vocab) {
+        ggml_tensor * idx = ggml_top_k(ctx0, cur, top_k);
+        candidates = idx;
+        cur        = sort_by(cur, idx);
+        cb(cur, "sample_top_k_logits", -1);
+    }
+
+    if (top_p < 1.0f) {
+        ggml_tensor * sorted_idx    = ggml_argsort(ctx0, cur, GGML_SORT_ORDER_DESC);
+        ggml_tensor * sorted_logits = sort_by(cur, sorted_idx);
+        candidates = candidates ? sort_by(candidates, sorted_idx) : sorted_idx;
+
+        ggml_tensor * probs = ggml_soft_max(ctx0, sorted_logits);
+        ggml_tensor * cdf   = ggml_cumsum(ctx0, probs);
+
+        // keep_mask[i] = 1 once cdf[i] crosses top_p
+        ggml_tensor * cdf_scaled = ggml_scale_bias(ctx0, cdf, -1.0f, top_p);
+        ggml_tensor * keep_mask  = ggml_step(ctx0, cdf_scaled);
+        ggml_tensor * idxf       = ggml_sum(ctx0, keep_mask);
+        idxf = ggml_clamp(ctx0, idxf, 0.0f, (float) keep_mask->ne[0] - 1);
+        ggml_tensor * ones = ggml_scale_bias(ctx0, idxf, 0.0f, 1.0f);
+
+        // top-p must include the crossing element, so force it to 1
+        ggml_tensor * keep_mask_2d = ggml_reshape_2d(ctx0, keep_mask, 1, keep_mask->ne[0]);
+        keep_mask_2d = ggml_set_rows(ctx0, keep_mask_2d, ones, ggml_cast(ctx0, idxf, GGML_TYPE_I32));
+        keep_mask    = ggml_reshape_1d(ctx0, keep_mask_2d, keep_mask->ne[0]);
+
+        // log(1) = 0 (keep), log(0) = -inf (drop)
+        ggml_tensor * bias = ggml_log(ctx0, keep_mask);
+        cur = ggml_add(ctx0, sorted_logits, bias);
+        cb(cur, "sample_top_p_logits", -1);
+    }
+
+    // draw one token: find where the cdf crosses inp_rand
+    ggml_tensor * probs  = ggml_soft_max(ctx0, cur);
+    ggml_tensor * cumsum = ggml_cumsum(ctx0, probs);
+
+    ggml_tensor * diff       = ggml_sub(ctx0, cumsum, inp_rand);
+    ggml_tensor * cross_mask = ggml_step(ctx0, diff);
+    ggml_tensor * idxf       = ggml_sum(ctx0, cross_mask);
+    ggml_tensor * idx        = ggml_cast(ctx0, ggml_scale_bias(ctx0, idxf, -1.0f, (float) cross_mask->ne[0]), GGML_TYPE_I32);
+
+    if (candidates) {
+        ggml_tensor * cand_2d = ggml_reshape_2d(ctx0, candidates, 1, candidates->ne[0]);
+        idx = ggml_get_rows(ctx0, cand_2d, idx);
+    }
+    cb(idx, "sample_token_id", -1);
+
+    return idx;
+}
+
+// returns a new cache with row row_idx set to value
+ggml_tensor * clip_graph_qwen3tts_gen::code_gen::cache_set(ggml_tensor * cache, int row_idx, ggml_tensor * value) const {
+    const int64_t n_embd  = cache->ne[0];
+    const int64_t n_cache = cache->ne[1];
+    GGML_ASSERT(row_idx >= 0 && row_idx < n_cache);
+
+    // append value as the last row, then gather it back into place
+    ggml_tensor * value_2d  = ggml_reshape_2d(ctx0, value, n_embd, 1);
+    ggml_tensor * cache_ext = ggml_concat(ctx0, cache, value_2d, 1); // [n_embd, n_cache + 1]
+
+    // gather indices [0..row_idx-1, n_cache, row_idx+1..n_cache-1]
+    // built via concat, since ggml_set_rows needs F32/F16 values, not an I32 index array
+    ggml_tensor * idx = const_i32(cache, (float) n_cache);
+    if (row_idx > 0) {
+        ggml_tensor * prefix = ggml_cast(ctx0, ggml_arange(ctx0, 0.0f, (float) row_idx, 1.0f), GGML_TYPE_I32);
+        idx = ggml_concat(ctx0, prefix, idx, 0);
+    }
+    if (row_idx < n_cache - 1) {
+        ggml_tensor * suffix = ggml_cast(ctx0, ggml_arange(ctx0, (float) (row_idx + 1), (float) n_cache, 1.0f), GGML_TYPE_I32);
+        idx = ggml_concat(ctx0, idx, suffix, 0);
+    }
+
+    ggml_tensor * result = ggml_get_rows(ctx0, cache_ext, idx);
+    cb(result, "cache_set_out", -1);
+    return result;
+}
+
+// builds a const i32 with no host upload: view a tensor, zero it via scale, add value, cast to i32
+ggml_tensor * clip_graph_qwen3tts_gen::code_gen::const_i32(ggml_tensor * anchor, float value) const {
+    ggml_tensor * v = ggml_view_1d(ctx0, anchor, 1, 0);
+    if (v->type != GGML_TYPE_F32) {
+        v = ggml_cast(ctx0, v, GGML_TYPE_F32);
+    }
+    return ggml_cast(ctx0, ggml_scale_bias(ctx0, v, 0.0f, value), GGML_TYPE_I32);
+}
+
+// causal keep-mask row for a query at position pos, window size n_kv_pad
+ggml_tensor * clip_graph_qwen3tts_gen::code_gen::causal_mask_row(int64_t n_kv_pad, int pos) const {
+    ggml_tensor * ones = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_kv_pad, n_kv_pad), 1.0f);
+    ggml_tensor * keep = ggml_tri(ctx0, ones, GGML_TRI_TYPE_LOWER_DIAG);
+    ggml_tensor * row  = ggml_view_1d(ctx0, keep, n_kv_pad, (size_t) pos * keep->nb[1]);
+    ggml_tensor * mask = ggml_log(ctx0, row); // 0 = keep, -inf = masked
+    return ggml_reshape_4d(ctx0, mask, n_kv_pad, 1, 1, 1);
+}
+
+// talker hidden size -> predictor hidden size (small_to_mtp_projection)
+ggml_tensor * clip_graph_qwen3tts_gen::code_gen::project_in(ggml_tensor * cur) const {
+    if (!model.gen_code_proj_in_w) {
+        return cur;
+    }
+    cur = ggml_mul_mat(ctx0, model.gen_code_proj_in_w, cur);
+    if (model.gen_code_proj_in_b) {
+        cur = ggml_add(ctx0, cur, model.gen_code_proj_in_b);
+    }
+    return cur;
+}
+
+// one transformer layer at position pos; writes k/v into k_cache_layer/v_cache_layer at row pos
+ggml_tensor * clip_graph_qwen3tts_gen::code_gen::layer_forward(
+        ggml_tensor * cur,
+        const clip_layer & layer,
+        ggml_tensor * inp_pos,
+        ggml_tensor * kq_mask,
+        ggml_tensor *& k_cache_layer,
+        ggml_tensor *& v_cache_layer,
+        int64_t n_kv_pad,
+        int pos,
+        int il) const {
+    const int     n_head    = hparams.n_head;
+    const int     n_head_kv = hparams.n_head_kv;
+    const int64_t d_head    = layer.q_w->ne[1] / n_head; // real head_dim, not n_embd / n_head
+    const float   kq_scale  = 1.0f / sqrtf((float) d_head);
+
+    ggml_tensor * residual = cur;
+
+    ggml_tensor * h = ggml_rms_norm(ctx0, cur, hparams.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);
+
+    q = ggml_reshape_3d(ctx0, q, d_head, n_head, 1);
+    k = ggml_reshape_3d(ctx0, k, d_head, n_head_kv, 1);
+
+    q = ggml_rms_norm(ctx0, q, hparams.eps);
+    q = ggml_mul(ctx0, q, layer.q_norm);
+    k = ggml_rms_norm(ctx0, k, hparams.eps);
+    k = ggml_mul(ctx0, k, layer.k_norm);
+
+    q = ggml_rope_ext(ctx0, q, inp_pos, nullptr, (int) d_head, GGML_ROPE_TYPE_NEOX, 0,
+                      hparams.rope_theta, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f);
+    k = ggml_rope_ext(ctx0, k, inp_pos, nullptr, (int) d_head, GGML_ROPE_TYPE_NEOX, 0,
+                      hparams.rope_theta, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f);
+
+    // write k/v into the cache at row pos, flat layout
+    ggml_tensor * k_flat = ggml_reshape_1d(ctx0, k, d_head * n_head_kv);
+    k_cache_layer = cache_set(k_cache_layer, pos, k_flat);
+    v_cache_layer = cache_set(v_cache_layer, pos, v);
+
+    ggml_tensor * q_cur = ggml_reshape_4d(ctx0, q, d_head, n_head, 1, 1);
+    ggml_tensor * k_cur = ggml_reshape_4d(ctx0, k_cache_layer, d_head, n_head_kv, n_kv_pad, 1);
+    ggml_tensor * v_cur = ggml_reshape_4d(ctx0, v_cache_layer, d_head, n_head_kv, n_kv_pad, 1);
+
+    ggml_tensor * attn_out = build_attn(layer.o_w, layer.o_b, q_cur, k_cur, v_cur, kq_mask, kq_scale, il);
+
+    cur = ggml_add(ctx0, residual, attn_out);
+
+    ggml_tensor * h2 = ggml_rms_norm(ctx0, cur, hparams.eps);
+    h2 = ggml_mul(ctx0, h2, layer.ln_2_w);
+
+    ggml_tensor * gate = ggml_mul_mat(ctx0, layer.ff_gate_w, h2);
+    ggml_tensor * up   = ggml_mul_mat(ctx0, layer.ff_up_w, h2);
+    ggml_tensor * gu   = ggml_swiglu_split(ctx0, gate, up);
+    ggml_tensor * down = ggml_mul_mat(ctx0, layer.ff_down_w, gu);
+
+    return ggml_add(ctx0, cur, down);
+}
+
+// position 0: hidden bridge, seeds the k/v cache, no sampling
+// position 1: embed(code0), sample with lm_head[0], write out_code_cache[1]
+void clip_graph_qwen3tts_gen::code_gen::prefill(
+        std::vector<ggml_tensor *> & k_cache,
+        std::vector<ggml_tensor *> & v_cache,
+        ggml_tensor *& out_code_cache,
+        ggml_tensor * h_state,
+        ggml_tensor * code0_embd,
+        ggml_tensor * inp_rand) const {
+    const int64_t n_kv_pad = k_cache[0]->ne[1];
+
+    {
+        ggml_tensor * cur     = project_in(h_state);
+        ggml_tensor * kq_mask = causal_mask_row(n_kv_pad, 0);
+        ggml_tensor * inp_pos = const_i32(k_cache[0], 0.0f);
+        for (size_t il = 0; il < model.layers.size(); il++) {
+            cur = layer_forward(cur, model.layers[il], inp_pos, kq_mask, k_cache[il], v_cache[il], n_kv_pad, 0, (int) il);
+        }
+        // position 0's output is unused, it only seeded the cache
+    }
+
+    {
+        ggml_tensor * cur     = project_in(code0_embd);
+        ggml_tensor * kq_mask = causal_mask_row(n_kv_pad, 1);
+        ggml_tensor * inp_pos = const_i32(k_cache[0], 1.0f);
+        for (size_t il = 0; il < model.layers.size(); il++) {
+            cur = layer_forward(cur, model.layers[il], inp_pos, kq_mask, k_cache[il], v_cache[il], n_kv_pad, 1, (int) il);
+        }
+
+        cur = ggml_rms_norm(ctx0, cur, hparams.eps);
+        cur = ggml_mul(ctx0, cur, model.gen_code_norm_w);
+
+        ggml_tensor * head_w = model.gen_code_head_w;
+        ggml_tensor * head_g = ggml_view_2d(ctx0, head_w, head_w->ne[0], head_w->ne[1], head_w->nb[1], 0); // lm_head[0]
+        ggml_tensor * logits = ggml_mul_mat(ctx0, head_g, cur);
+
+        ggml_tensor * sampled = do_sampling(logits, inp_rand);
+        out_code_cache = cache_set(out_code_cache, 1, sampled);
+    }
+}
+
+// one decode step of code_predictor
+// at step_idx g:
+// - read code from out_code_cache[g], then embed it with codebook table g-1
+// - write new kv at cache row g+1, sample with lm_head[g]
+// - write result to out_code_cache[g+1]
+// step_idx must be in [1, n_acoustic - 1]
+ggml_tensor * clip_graph_qwen3tts_gen::code_gen::step(
+        std::vector<ggml_tensor *> & k_cache,
+        std::vector<ggml_tensor *> & v_cache,
+        ggml_tensor * out_code_cache,
+        ggml_tensor * inp_rand,
+        int step_idx) const {
+    const int64_t n_acoustic = model.gen_code_head_w->ne[2];
+    GGML_ASSERT(step_idx >= 1 && step_idx < n_acoustic);
+    GGML_ASSERT(k_cache.size() == model.layers.size());
+    GGML_ASSERT(v_cache.size() == model.layers.size());
+
+    const int64_t n_kv_pad = k_cache[0]->ne[1];
+    const int     pos      = step_idx + 1; // new cache row and RoPE position
+
+    // embed the previous code via this step's codebook table (rows are already scalars)
+    ggml_tensor * code_in = ggml_view_1d(ctx0, out_code_cache, 1, (size_t) step_idx * out_code_cache->nb[1]);
+
+    ggml_tensor * embd_w = model.gen_code_embd_w; // [n_embd_talker, vocab, n_acoustic]
+    ggml_tensor * embd_g = ggml_view_2d(ctx0, embd_w, embd_w->ne[0], embd_w->ne[1], embd_w->nb[1],
+                                        (size_t) (step_idx - 1) * embd_w->nb[2]);
+    ggml_tensor * cur = ggml_get_rows(ctx0, embd_g, code_in);
+    cur = ggml_reshape_1d(ctx0, cur, cur->ne[0]);
+    cb(cur, "step_embd_in", step_idx);
+
+    cur = project_in(cur);
+    cb(cur, "step_proj_in", step_idx);
+
+    ggml_tensor * kq_mask = causal_mask_row(n_kv_pad, pos);
+    ggml_tensor * inp_pos = const_i32(k_cache[0], (float) pos);
+
+    for (size_t il = 0; il < model.layers.size(); il++) {
+        cur = layer_forward(cur, model.layers[il], inp_pos, kq_mask, k_cache[il], v_cache[il], n_kv_pad, pos, (int) il);
+        cb(cur, "step_layer_out", (int) il);
+    }
+
+    // final norm, this step's lm_head, sample, write the result
+    cur = ggml_rms_norm(ctx0, cur, hparams.eps);
+    cur = ggml_mul(ctx0, cur, model.gen_code_norm_w);
+
+    ggml_tensor * head_w = model.gen_code_head_w; // [n_embd_pred, vocab, n_acoustic]
+    ggml_tensor * head_g = ggml_view_2d(ctx0, head_w, head_w->ne[0], head_w->ne[1], head_w->nb[1],
+                                        (size_t) step_idx * head_w->nb[2]);
+    ggml_tensor * logits = ggml_mul_mat(ctx0, head_g, cur);
+    cb(logits, "step_logits", step_idx);
+
+    ggml_tensor * sampled = do_sampling(logits, inp_rand);
+    cb(sampled, "step_sampled", step_idx);
+
+    return cache_set(out_code_cache, pos, sampled);
+}
+
+// causal conv1d, stride 1: prepend persisted left-context instead of zero-padding, then a plain conv
+// x: [T, IC] (T-first). w: [K, IC, OC]. state_name empty means K == 1 (no left-context). 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_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]. 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_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, the (kernel - stride) overlap tail is kept as state for the next call
+// x: [T, IC], w: [K, OC, IC]. state_name empty means K == stride (no overlap). 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     OC       = (int) w->ne[1];
+    const int     trim     = K - stride;
+    const int64_t emit_len = x->ne[0] * stride;
+
+    // transposed conv as GEMM + col2im scatter-add, y: [emit_len + trim, OC]
+    ggml_tensor * w2  = ggml_reshape_2d(ctx0, w, (int64_t) K * OC, w->ne[2]);
+    w2                = ggml_cont(ctx0, ggml_transpose(ctx0, w2));
+    ggml_tensor * xt  = ggml_cont(ctx0, ggml_transpose(ctx0, x));
+    ggml_tensor * col = ggml_mul_mat(ctx0, w2, xt);
+    ggml_tensor * y   = ggml_col2im_1d(ctx0, col, stride, OC, 0);
+
+    ggml_tensor * out = y;
+    if (trim > 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) {
+        out = ggml_add(ctx0, out, ggml_reshape_2d(ctx0, b, 1, b->ne[0]));
+    }
+    return out;
+}
+
+// SnakeBeta activation: y = x + sin(alpha*x)^2 * inv_beta (alpha/inv_beta folded via exp/reciprocal at conversion time)
+// x: [T, C]. alpha/beta: [C], broadcasts over T
+ggml_tensor * clip_graph_qwen3tts_gen::code2wav::snake(ggml_tensor * x, ggml_tensor * alpha, ggml_tensor * beta) const {
+    ggml_tensor * a = ggml_reshape_2d(ctx0, alpha, 1, alpha->ne[0]);
+    ggml_tensor * b = ggml_reshape_2d(ctx0, beta,  1, beta->ne[0]);
+
+    // expand reshapes first so mul/sin/sqr/mul/add lands as consecutive nodes, letting backends fuse them
+    ggml_build_forward_expand(gf, a);
+    ggml_build_forward_expand(gf, b);
+
+    ggml_tensor * s = ggml_sin(ctx0, ggml_mul(ctx0, x, a));
+    s = ggml_sqr(ctx0, s);
+    s = ggml_mul(ctx0, s, b);
+    return ggml_add(ctx0, x, s);
+}
+
+// RVQ codebook decode: T frames of 16 codes -> 512-dim hidden (C-first, [512, T])
+// codebook 0 (semantic) and 1..15 (acoustic) sum within their group, project separately, then add
+ggml_tensor * clip_graph_qwen3tts_gen::code2wav::quant_decode(ggml_tensor * inp_codes) const {
+    const auto & c2w = model.c2w;
+    const int64_t T = inp_codes->ne[0];
+
+    // ids for codebook group g over all T frames, [T] I32
+    auto group_ids = [&](int g) {
+        return ggml_view_1d(ctx0, inp_codes, T, (size_t) g * inp_codes->nb[1]);
+    };
+
+    ggml_tensor * sem     = ggml_get_rows(ctx0, c2w.quant_first_cb_w, group_ids(0)); // [256, T]
+    ggml_tensor * sem_out = ggml_mul_mat(ctx0, c2w.quant_first_out_w, sem);          // [512, T]
+
+    ggml_tensor * acc = nullptr;
+    const int64_t n_acoustic = c2w.quant_rest_cb_w->ne[2];
+    for (int g = 1; g <= n_acoustic; g++) {
+        ggml_tensor * cb_g  = ggml_view_2d(ctx0, c2w.quant_rest_cb_w, c2w.quant_rest_cb_w->ne[0], c2w.quant_rest_cb_w->ne[1],
+                                           c2w.quant_rest_cb_w->nb[1], (size_t) (g - 1) * c2w.quant_rest_cb_w->nb[2]);
+        ggml_tensor * embd = ggml_get_rows(ctx0, cb_g, group_ids(g)); // [256, T]
+        acc = acc ? ggml_add(ctx0, acc, embd) : embd;
+    }
+    ggml_tensor * ac_out = ggml_mul_mat(ctx0, c2w.quant_rest_out_w, acc); // [512, T]
+
+    ggml_tensor * hidden = ggml_add(ctx0, sem_out, ac_out);
+    cb(hidden, "wav_quant_hidden", -1);
+    return hidden;
+}
+
+// one pre_transformer layer over a batch of N = sliding_window new frames
+// attention runs over [(W-1)-frame prefix from the last batch] + [N new frames]
+// RoPE positions come from a persisted counter, so phases line up across batches
+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 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); // [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, 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);
+
+    // the position counter is the same for all layers, 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 * 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: key j is visible to query i 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);
+
+    // on a cold start, key j is real state only when j >= prefix - tfm_pos, mask out the rest
+    ggml_tensor * warm = ggml_step(ctx0, ggml_scale_bias(ctx0, ggml_add(ctx0, pos_k, base),
+                                                         1.0f, 0.5f - (float) prefix)); // j + pos > prefix - 0.5
+    keep = ggml_mul(ctx0, keep, warm);
+
+    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);
+    }
+    cur = ggml_add(ctx0, residual, attn_out);
+
+    ggml_tensor * residual2 = cur;
+    ggml_tensor * h2 = ggml_rms_norm(ctx0, cur, hparams.wav_tfm_eps);
+    h2 = ggml_mul(ctx0, h2, layer.ln_2_w);
+
+    ggml_tensor * gate = ggml_mul_mat(ctx0, layer.ff_gate_w, h2);
+    ggml_tensor * up   = ggml_mul_mat(ctx0, layer.ff_up_w, h2);
+    ggml_tensor * gu   = ggml_swiglu_split(ctx0, gate, up);
+    ggml_tensor * down = ggml_mul_mat(ctx0, layer.ff_down_w, gu);
+    if (layer.ls_2_w) {
+        down = ggml_mul(ctx0, down, layer.ls_2_w);
+    }
+    return ggml_add(ctx0, residual2, down);
+}
+
+// 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
+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, state_prefix + "_dwconv"); // [T, C]
+    ggml_tensor * hc = ggml_cont(ctx0, ggml_transpose(ctx0, h)); // [C, T]
+
+    hc = ggml_norm(ctx0, hc, 1e-6f);
+    hc = ggml_mul(ctx0, hc, blk.norm_w);
+    hc = ggml_add(ctx0, hc, blk.norm_b);
+
+    ggml_tensor * g = ggml_mul_mat(ctx0, blk.pw1_w, hc);
+    g = ggml_add(ctx0, g, blk.pw1_b);
+    g = ggml_gelu(ctx0, g);
+    g = ggml_mul_mat(ctx0, blk.pw2_w, g);
+    g = ggml_add(ctx0, g, blk.pw2_b);
+    g = ggml_mul(ctx0, g, blk.gamma);
+
+    ggml_tensor * g_t = ggml_cont(ctx0, ggml_transpose(ctx0, g)); // back to [T, C]
+    return ggml_add(ctx0, residual, g_t);
+}
+
+// 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 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, state_name);
+    h = snake(h, res.act2_alpha, res.act2_beta);
+    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 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;
+
+    // 1. quantizer decode: N frames of 16 codes -> [512, N] (C-first)
+    ggml_tensor * hidden = quant_decode(inp_codes);
+
+    // 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, N], project down, run the layers, 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), N]
+
+    for (int il = 0; il < hparams.wav_tfm_n_layer; il++) {
+        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, N]
+    cb(cur, "wav_tfm_out", -1);
+
+    // 4. upsample: 2x (causal ConvTranspose1d, stride 2 + ConvNeXt block), back to T-first
+    // kernel == stride here, so there is no overlap tail to persist
+    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, "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_DILATIONS[3] = { 1, 3, 9 };
+
+    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, 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], 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, "dac_post_conv"); // [n_samples, 1]
+
+    x = ggml_clamp(ctx0, x, -1.0f, 1.0f);
+    x = ggml_reshape_1d(ctx0, x, x->ne[0]);
+    cb(x, "wav_audio_out", -1);
+    return x;
+}
+
+// code2wav's persisted state buffers: RoPE position counter, K/V per pre_transformer layer,
+// left-context/tail per stateful conv. shape lookup only, no graph needed
+std::vector<c2w_state_slot> list_c2w_state_slots(const clip_hparams & hparams, const clip_model & model) {
+    const auto & c2w = model.c2w;
+    std::vector<c2w_state_slot> slots;
+
+    slots.push_back({"tfm_pos", 1, 1});
+
+    // prefix is (W-1) frames, the batch itself gives the other N=W 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;
+}
+
+// both sub-graphs are always built, so the topology stays constant
+// ggml_build_forward_select() then picks the one that actually runs
+ggml_cgraph * clip_graph_qwen3tts_gen::build() {
+    GGML_ASSERT(n_batch == 1); // this module only ever processes one frame at a time
+
+    int idx;
+    switch (gen_process) {
+        case CLIP_GEN_PROCESS_GEN_CODE: idx = 0; break;
+        case CLIP_GEN_PROCESS_GEN_WAV:  idx = 1; break;
+        default: GGML_ABORT("unknown gen_process");
+    }
+
+    // ---- CLIP_GEN_PROCESS_GEN_CODE: backbone hidden state -> 16 RVQ codes + next-step embd ----
+    // not build_inp_raw(), a GEN_WAV call's `img` has no hidden-state data
+    ggml_tensor * h_state = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, n_mmproj_embd);
+    ggml_set_name(h_state, "inp_raw"); // must keep this exact name, clip_encode() sets it by name
+    ggml_set_input(h_state);
+
+    ggml_tensor * code0 = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 1);
+    ggml_set_name(code0, "inp_code0");
+    ggml_set_input(code0);
+
+    ggml_tensor * code0_embd = ggml_get_rows(ctx0, model.gen_code_out_embd_w, code0);
+    code0_embd = ggml_reshape_1d(ctx0, code0_embd, code0_embd->ne[0]);
+    cb(code0_embd, "code0_embd", -1);
+
+    const int64_t n_acoustic = model.gen_code_head_w->ne[2]; // 15
+    const int     n_codes    = (int) n_acoustic + 1;         // 16
+    const int64_t n_kv_pad   = n_codes;
+    const int     n_layer    = (int) model.layers.size();
+    const int     n_head     = hparams.n_head;
+    const int     n_head_kv  = hparams.n_head_kv;
+    const int64_t d_head     = model.layers[0].q_w->ne[1] / n_head;
+
+    // zero-filled per layer k/v caches, so masked-out rows can't hold garbage
+    std::vector<ggml_tensor *> k_cache(n_layer), v_cache(n_layer);
+    for (int il = 0; il < n_layer; il++) {
+        k_cache[il] = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, d_head * n_head_kv, n_kv_pad), 0.0f);
+        v_cache[il] = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, d_head * n_head_kv, n_kv_pad), 0.0f);
+    }
+
+    code_gen cg(*this, top_k, top_p);
+
+    ggml_tensor * out_code_cache = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, 1, n_codes);
+    out_code_cache = cg.cache_set(out_code_cache, 0, code0);
+
+    ggml_tensor * inp_rand0 = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, 1);
+    ggml_set_name(inp_rand0, "inp_rand_0");
+    ggml_set_input(inp_rand0);
+
+    cg.prefill(k_cache, v_cache, out_code_cache, h_state, code0_embd, inp_rand0);
+
+    for (int g = 1; g < n_acoustic; g++) {
+        ggml_tensor * inp_rand = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, 1);
+        ggml_set_name(inp_rand, ("inp_rand_" + std::to_string(g)).c_str());
+        ggml_set_input(inp_rand);
+        out_code_cache = cg.step(k_cache, v_cache, out_code_cache, inp_rand, g);
+    }
+
+    // output 1: this frame's 16 sampled codes, for the caller's code2wav window
+    ggml_tensor * out_codes = ggml_cont(ctx0, out_code_cache);
+    ggml_set_name(out_codes, "out_codes");
+    ggml_set_output(out_codes);
+
+    // output 2: sum of all 16 codebook embeddings, fed back to the talker for the next frame
+    ggml_tensor * out_embd = code0_embd;
+    for (int g = 1; g <= n_acoustic; g++) {
+        ggml_tensor * code_g = ggml_view_1d(ctx0, out_code_cache, 1, (size_t) g * out_code_cache->nb[1]);
+
+        ggml_tensor * embd_g = ggml_view_2d(ctx0, model.gen_code_embd_w, model.gen_code_embd_w->ne[0], model.gen_code_embd_w->ne[1],
+                                            model.gen_code_embd_w->nb[1], (size_t) (g - 1) * model.gen_code_embd_w->nb[2]);
+        ggml_tensor * e = ggml_get_rows(ctx0, embd_g, code_g);
+        e = ggml_reshape_1d(ctx0, e, e->ne[0]);
+
+        out_embd = ggml_add(ctx0, out_embd, e);
+    }
+    out_embd = ggml_reshape_2d(ctx0, out_embd, out_embd->ne[0], 1);
+    cb(out_embd, "gen_audio_out", -1);
+
+    // ---- CLIP_GEN_PROCESS_GEN_WAV: 16 RVQ codes -> raw PCM ----
+    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);
+
+    code2wav c2w(*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);
+        c2w.state_in[slot.name] = t;
+    }
+
+    ggml_tensor * out_audio = c2w.decode(inp_codes);
+    ggml_set_name(out_audio, "out_audio");
+    ggml_set_output(out_audio);
+
+    for (auto & slot : c2w.state_out) {
+        ggml_set_name(slot.second, ("state_out_" + slot.first).c_str());
+        ggml_set_output(slot.second);
+    }
+
+    // out_embd goes last, clip_encode() reads it back via ggml_graph_node(gf, -1)
+    ggml_tensor * outs[2];
+    outs[0] = out_codes; outs[1] = out_audio;
+    ggml_build_forward_select(gf, outs, 2, idx);
+    for (auto & slot : c2w.state_out) {
+        outs[0] = out_codes; outs[1] = slot.second;
+        ggml_build_forward_select(gf, outs, 2, idx);
+    }
+    outs[0] = out_embd; outs[1] = out_audio;
+    ggml_build_forward_select(gf, outs, 2, idx);
+
+    return gf;
+}
diff --git a/tools/mtmd/models/qwen3tts-spkenc.cpp b/tools/mtmd/models/qwen3tts-spkenc.cpp
new file mode 100644 (file)
index 0000000..d4659fd
--- /dev/null
@@ -0,0 +1,197 @@
+#include "models.h"
+
+static constexpr int SPK_RES2NET_SCALE = 8; // enc_res2net_scale
+static constexpr int SPK_DILATIONS[3]  = { 2, 3, 4 }; // enc_dilations[1..3]
+
+// conv1d, kernel K, padding "same" (reflect), dilation d
+// x: [C, T] (ne[0]=C, ne[1]=T) -> [out_c, T]
+ggml_tensor * clip_graph_qwen3tts_spkenc::conv1d_same(ggml_tensor * x, ggml_tensor * w, ggml_tensor * b, int dilation) const {
+    const int K   = (int) w->ne[0];
+    const int IC  = (int) w->ne[1];
+    const int OC  = (int) w->ne[2];
+    const int pad = ((K - 1) * dilation) / 2;
+
+    // ggml_pad_reflect_1d pads ne[0], so bring T onto ne[0] first, same layout as im2col wants
+    ggml_tensor * x_t = ggml_cont(ctx0, ggml_transpose(ctx0, x)); // [T, IC]
+    if (pad > 0) {
+        x_t = ggml_pad_reflect_1d(ctx0, x_t, pad, pad); // [T + 2*pad, IC]
+    }
+    ggml_tensor * x4d = ggml_reshape_4d(ctx0, x_t, x_t->ne[0], IC, 1, 1);
+
+    // dummy F32 kernel, im2col only reads its shape, so a quantized w does not assert
+    ggml_tensor * dummy = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, K, IC, 1, 1);
+
+    ggml_tensor * col = ggml_im2col(ctx0, dummy, x4d, 1, 1, 0, 0, dilation, 1, false, GGML_TYPE_F32);
+    const int64_t T_out = col->ne[1];
+    col = ggml_reshape_2d(ctx0, col, (int64_t) K * IC, T_out);
+
+    ggml_tensor * w2d = ggml_reshape_2d(ctx0, w, (int64_t) K * IC, OC);
+    ggml_tensor * y   = ggml_mul_mat(ctx0, w2d, col); // [OC, T_out]
+    ggml_mul_mat_set_prec(y, GGML_PREC_F32);
+
+    ggml_tensor * b2d = ggml_reshape_2d(ctx0, b, OC, 1);
+    y = ggml_add(ctx0, y, b2d);
+    return y;
+}
+
+// Res2Net: split channel axis into `scale` chunks, chain dilated conv1d branches
+// x: [C, T] -> [C, T]
+ggml_tensor * clip_graph_qwen3tts_spkenc::res2net(ggml_tensor * x, const clip_layer & layer, int dilation, int scale) const {
+    const int64_t C  = x->ne[0];
+    const int64_t T  = x->ne[1];
+    const int64_t Cs = C / scale;
+
+    std::vector<ggml_tensor *> outs;
+    outs.reserve(scale);
+
+    auto chunk = [&](int i) -> ggml_tensor * {
+        return ggml_view_2d(ctx0, x, Cs, T, x->nb[1], (size_t) i * Cs * x->nb[0]);
+    };
+
+    ggml_tensor * prev = nullptr;
+    for (int i = 0; i < scale; i++) {
+        ggml_tensor * c = ggml_cont(ctx0, chunk(i));
+        if (i == 0) {
+            outs.push_back(c);
+            continue;
+        }
+        ggml_tensor * inp = (i >= 2) ? ggml_add(ctx0, c, prev) : c;
+        ggml_tensor * y   = conv1d_same(inp, layer.res2_conv_w[i - 1], layer.res2_conv_b[i - 1], dilation);
+        y                 = ggml_relu(ctx0, y);
+        outs.push_back(y);
+        prev = y;
+    }
+
+    ggml_tensor * acc = outs[0];
+    for (int i = 1; i < scale; i++) {
+        acc = ggml_concat(ctx0, acc, outs[i], 0);
+    }
+    return acc;
+}
+
+// squeeze-and-excitation gate. x: [C, T] -> [C, T]
+ggml_tensor * clip_graph_qwen3tts_spkenc::se_block(ggml_tensor * x, const clip_layer & layer) const {
+    // temporal mean, keepdim: transpose so T is on ne[0], reduce, transpose back
+    ggml_tensor * x_t  = ggml_cont(ctx0, ggml_transpose(ctx0, x));    // [T, C]
+    ggml_tensor * mean = ggml_mean(ctx0, x_t);                        // [1, C]
+    mean               = ggml_cont(ctx0, ggml_transpose(ctx0, mean)); // [C, 1]
+
+    ggml_tensor * h = conv1d_same(mean, layer.se_conv1_w, layer.se_conv1_b, 1);
+    h = ggml_relu(ctx0, h);
+    h = conv1d_same(h, layer.se_conv2_w, layer.se_conv2_b, 1);
+    h = ggml_sigmoid(ctx0, h); // [C, 1]
+
+    return ggml_mul(ctx0, x, h); // broadcast gate over T
+}
+
+// tdnn1 -> res2net -> tdnn2 -> se, plus residual. x: [C, T] -> [C, T]
+ggml_tensor * clip_graph_qwen3tts_spkenc::se_res2net_block(ggml_tensor * x, const clip_layer & layer, int dilation, int scale) const {
+    ggml_tensor * residual = x;
+    ggml_tensor * h  = conv1d_same(x, layer.conv_pw1_w, layer.conv_pw1_b, 1); // tdnn1
+    h = ggml_relu(ctx0, h);
+    h = res2net(h, layer, dilation, scale);
+    h = conv1d_same(h, layer.conv_pw2_w, layer.conv_pw2_b, 1); // tdnn2
+    h = ggml_relu(ctx0, h);
+    h = se_block(h, layer);
+    return ggml_add(ctx0, h, residual);
+}
+
+// attentive statistics pooling. x: [C, T] -> [2*C, 1]
+ggml_tensor * clip_graph_qwen3tts_spkenc::attentive_stats_pool(ggml_tensor * x) const {
+    const int64_t T = x->ne[1];
+
+    // mean over T: [C, 1]
+    ggml_tensor * x_t  = ggml_cont(ctx0, ggml_transpose(ctx0, x));
+    ggml_tensor * mean = ggml_mean(ctx0, x_t);
+    mean               = ggml_cont(ctx0, ggml_transpose(ctx0, mean));
+
+    // std over T: sqrt(clamp(mean((x - mean)^2), eps))
+    ggml_tensor * mean_rep = ggml_repeat(ctx0, mean, x);
+    ggml_tensor * centered = ggml_sub(ctx0, x, mean_rep);
+    ggml_tensor * var_t    = ggml_cont(ctx0, ggml_transpose(ctx0, ggml_sqr(ctx0, centered)));
+    ggml_tensor * var      = ggml_mean(ctx0, var_t);
+    var                    = ggml_cont(ctx0, ggml_transpose(ctx0, var));
+    var                    = ggml_scale_bias(ctx0, var, 1.0f, 1e-12f);
+    ggml_tensor * std      = ggml_sqrt(ctx0, var);
+
+    // attention input: cat([x, mean, std]) along channel axis -> [3C, T]
+    ggml_tensor * std_rep = ggml_repeat(ctx0, std, x);
+    ggml_tensor * cat     = ggml_concat(ctx0, x, mean_rep, 0);
+    cat                   = ggml_concat(ctx0, cat, std_rep, 0);
+
+    // attention TDNN (3C -> attn_c) + ReLU, tanh, then 1x1 conv (attn_c -> C)
+    ggml_tensor * a = conv1d_same(cat, model.spk_asp_tdnn_w, model.spk_asp_tdnn_b, 1);
+    a = ggml_relu(ctx0, a);
+    a = ggml_tanh(ctx0, a);
+    a = conv1d_same(a, model.spk_asp_attn_w, model.spk_asp_attn_b, 1);
+
+    // softmax over T
+    ggml_tensor * a_t = ggml_cont(ctx0, ggml_transpose(ctx0, a));   // [T, C]
+    ggml_tensor * w_t = ggml_soft_max(ctx0, a_t);
+    ggml_tensor * w   = ggml_cont(ctx0, ggml_transpose(ctx0, w_t)); // [C, T]
+
+    // weighted mean: sum(w * x) over T, multiply by T to undo ggml_mean's 1/T scaling
+    ggml_tensor * wx     = ggml_mul(ctx0, w, x);
+    ggml_tensor * wx_t   = ggml_cont(ctx0, ggml_transpose(ctx0, wx));
+    ggml_tensor * w_mean = ggml_mean(ctx0, wx_t);
+    w_mean = ggml_scale(ctx0, w_mean, (float) T);
+    w_mean = ggml_cont(ctx0, ggml_transpose(ctx0, w_mean)); // [C, 1]
+
+    // weighted std: sum(w * (x - w_mean)^2) over T
+    ggml_tensor * w_mean_rep = ggml_repeat(ctx0, w_mean, x);
+    ggml_tensor * dev        = ggml_sub(ctx0, x, w_mean_rep);
+    ggml_tensor * w_var_in   = ggml_mul(ctx0, w, ggml_sqr(ctx0, dev));
+    ggml_tensor * w_var_t    = ggml_cont(ctx0, ggml_transpose(ctx0, w_var_in));
+    ggml_tensor * w_var      = ggml_mean(ctx0, w_var_t);
+    w_var                    = ggml_scale(ctx0, w_var, (float) T);
+    w_var                    = ggml_cont(ctx0, ggml_transpose(ctx0, w_var));
+    w_var                    = ggml_scale_bias(ctx0, w_var, 1.0f, 1e-12f);
+    ggml_tensor * w_std      = ggml_sqrt(ctx0, w_var);
+
+    return ggml_concat(ctx0, w_mean, w_std, 0); // [2C, 1]
+}
+
+ggml_cgraph * clip_graph_qwen3tts_spkenc::build() {
+    // inp_raw: [T, n_mel, 1, 1], from mtmd_audio_preprocessor_qwen3tts_spk
+    ggml_tensor * inp = build_inp_raw(1);
+    inp = ggml_reshape_2d(ctx0, inp, inp->ne[0], inp->ne[1]);
+
+    // this file's convention is [C, T]; the preprocessor delivers [T, C]
+    ggml_tensor * mel = ggml_cont(ctx0, ggml_transpose(ctx0, inp)); // [n_mel, T]
+    cb(mel, "mel", -1);
+
+    // frontend conv0 TDNN k=5, dilation=1: 128 -> 512
+    ggml_tensor * cur = conv1d_same(mel, model.conv1d_1_w, model.conv1d_1_b, 1);
+    cur = ggml_relu(ctx0, cur);
+    cb(cur, "frontend", -1);
+
+    // 3 SE-Res2Net blocks at dilations 2, 3, 4
+    GGML_ASSERT((int) model.layers.size() == 3);
+    std::vector<ggml_tensor *> blk_out(3);
+    for (int il = 0; il < 3; il++) {
+        cur = se_res2net_block(cur, model.layers[il], SPK_DILATIONS[il], SPK_RES2NET_SCALE);
+        blk_out[il] = cur;
+        cb(cur, "block_out", il);
+    }
+
+    // multi-layer feature aggregation: cat blk[0..2] then TDNN k=1 + ReLU
+    ggml_tensor * cat = ggml_concat(ctx0, blk_out[0], blk_out[1], 0);
+    cat = ggml_concat(ctx0, cat, blk_out[2], 0); // [1536, T]
+    ggml_tensor * mfa = conv1d_same(cat, model.conv_out_w, model.conv_out_b, 1);
+    mfa = ggml_relu(ctx0, mfa);
+    cb(mfa, "mfa", -1);
+
+    // attentive statistics pooling: [1536, T] -> [3072, 1]
+    ggml_tensor * stats = attentive_stats_pool(mfa);
+    cb(stats, "asp", -1);
+
+    // final FC k=1: [3072, 1] -> [enc_dim, 1]
+    ggml_tensor * emb = conv1d_same(stats, model.mm_fc_w, model.mm_fc_b, 1);
+
+    emb = ggml_reshape_1d(ctx0, emb, emb->ne[0]);
+    emb = ggml_cont(ctx0, emb);
+    cb(emb, "spk_embedding", -1);
+
+    ggml_build_forward_expand(gf, emb);
+    return gf;
+}
index fea03557d05c018e1ebc9be6d3c7c1ad8f6f927c..7fbc18ea939676dea5f72f040249e82ddff8b7b8 100644 (file)
@@ -791,6 +791,66 @@ bool mtmd_audio_preprocessor_mimo_audio::preprocess(const float *
     return true;
 }
 
+//
+// mtmd_audio_preprocessor_qwen3tts_spk
+//
+// same as mel_spectrogram() in modeling_qwen3_tts.py
+// ECAPA-TDNN takes the whole clip in one pass, so no Whisper-style chunking or normalization
+//
+
+void mtmd_audio_preprocessor_qwen3tts_spk::initialize() {
+    cache.fill_sin_cos_table(hparams.audio_n_fft);
+    cache.fill_hann_window(hparams.audio_window_len, true);
+    cache.fill_mel_filterbank_matrix(hparams.n_mel_bins, hparams.audio_n_fft, hparams.audio_sample_rate);
+}
+
+bool mtmd_audio_preprocessor_qwen3tts_spk::preprocess(const float *                 samples,
+                                                      size_t                        n_samples,
+                                                      std::vector<mtmd_audio_mel> & output) {
+    if (n_samples == 0) {
+        return false;
+    }
+
+    GGML_ASSERT(!cache.sin_vals.empty());
+    GGML_ASSERT(!cache.cos_vals.empty());
+    GGML_ASSERT(!cache.filters.data.empty());
+
+    // reflect pad by (n_fft - hop) / 2 = 384, matching center=False STFT framing
+    const int pad = (hparams.audio_n_fft - hparams.audio_hop_len) / 2;
+    if (n_samples < (size_t) pad + 1) {
+        return false;
+    }
+
+    std::vector<float> padded(n_samples + 2 * pad, 0.0f);
+    for (int i = 0; i < pad; i++) {
+        padded[i] = samples[pad - i];
+    }
+    std::copy(samples, samples + n_samples, padded.begin() + pad);
+    for (int i = 0; i < pad; i++) {
+        padded[n_samples + pad + i] = samples[n_samples - 2 - i];
+    }
+
+    filter_params params;
+    params.n_mel            = hparams.n_mel_bins;
+    params.n_fft_bins       = 1 + (hparams.audio_n_fft / 2);
+    params.hann_window_size = hparams.audio_window_len;
+    params.hop_length       = hparams.audio_hop_len;
+    params.sample_rate      = hparams.audio_sample_rate;
+    params.no_padding       = true; // reflect padding already applied above
+    params.use_natural_log  = true;
+    params.use_magnitude    = true;
+    params.mel_floor        = 1e-5f;
+
+    mtmd_audio_mel out;
+    bool ok = log_mel_spectrogram(padded.data(), (int) padded.size(), 4, params, cache, out);
+    if (!ok) {
+        return false;
+    }
+
+    output.push_back(std::move(out));
+    return true;
+}
+
 //
 // mtmd_audio_preprocessor_conformer
 //
index f65f282d96e24ca9caec5324eb38ce4dca2b7c55..b4d6f725980852b6f42b4779716c38d416b28a9d 100644 (file)
@@ -120,6 +120,15 @@ struct mtmd_audio_preprocessor_mimo_audio : mtmd_audio_preprocessor {
     mtmd_audio_cache cache;
 };
 
+struct mtmd_audio_preprocessor_qwen3tts_spk : mtmd_audio_preprocessor {
+    mtmd_audio_preprocessor_qwen3tts_spk(const clip_ctx * ctx) : mtmd_audio_preprocessor(ctx) {}
+    void initialize() override;
+    bool preprocess(const float * samples, size_t n_samples, std::vector<mtmd_audio_mel> & output) override;
+
+  private:
+    mtmd_audio_cache cache;
+};
+
 struct mtmd_audio_preprocessor_parakeet : mtmd_audio_preprocessor {
     mtmd_audio_preprocessor_parakeet(clip_ctx * ctx) : mtmd_audio_preprocessor(ctx) { }
     void initialize() override;
index 08288c868139e3463208a4f0eae0b4ec975e9b36..07b45b6440738639c567ec6c4b31107c02b535f2 100644 (file)
@@ -116,6 +116,14 @@ struct mtmd_cli_context {
             exit(1);
         }
 
+        init_vision_context(params);
+
+        if (!mtmd_helper_model_can_chat(lctx, ctx_vision.get())) {
+            LOG_ERR("Model does not support chat mode\n");
+            LOG_ERR("Hint: for TTS models, please use llama-tts\n");
+            exit(1);
+        }
+
         if (!llama_model_chat_template(model, nullptr) && params.chat_template.empty()) {
             LOG_ERR("Model does not have chat template.\n");
             LOG_ERR("  For old llava models, you may need to use '--chat-template vicuna'\n");
@@ -129,8 +137,6 @@ struct mtmd_cli_context {
         chat_history.clear();
         LOG_INF("%s: chat template example:\n%s\n", __func__, common_chat_format_example(tmpls.get(), params.use_jinja, params.default_template_kwargs).c_str());
 
-        init_vision_context(params);
-
         // load antiprompt tokens for legacy templates
         if (params.chat_template == "vicuna") {
             antiprompt_tokens = common_tokenize(lctx, "ASSISTANT:", false, true);
diff --git a/tools/mtmd/mtmd-helper-common.h b/tools/mtmd/mtmd-helper-common.h
new file mode 100644 (file)
index 0000000..968b4df
--- /dev/null
@@ -0,0 +1,180 @@
+#pragma once
+
+// shared internal utilities for the mtmd-helper-*.cpp translation units
+// (mtmd-helper.cpp, mtmd-helper-gen.cpp)
+// NOT part of the public mtmd-helper.h API
+
+#include "ggml.h"
+#include "llama.h"
+#include "mtmd.h"
+
+#include <cstdarg>
+#include <cstdio>
+#include <cstdlib>
+#include <vector>
+
+//
+// logging
+//
+
+struct mtmd_helper_logger {
+    ggml_log_callback default_callback = [](ggml_log_level level, const char * text, void * user_data) {
+        (void) level;
+        (void) user_data;
+        fputs(text, stderr);
+        fflush(stderr);
+    };
+
+    ggml_log_callback log_callback = default_callback;
+    void * log_callback_user_data;
+
+    void log_v(enum ggml_log_level level, const char * format, va_list args) {
+        if (format == NULL) {
+            return;
+        }
+        va_list args_copy;
+        va_copy(args_copy, args);
+        char buffer[128];
+        int len = vsnprintf(buffer, 128, format, args);
+        if (len < 128) {
+            log_callback(level, buffer, log_callback_user_data);
+        } else {
+            char * buffer2 = (char *) calloc(len + 1, sizeof(char));
+            vsnprintf(buffer2, len + 1, format, args_copy);
+            buffer2[len] = 0;
+            log_callback(level, buffer2, log_callback_user_data);
+            free(buffer2);
+        }
+        va_end(args_copy);
+    }
+
+    void log(enum ggml_log_level level, const char * format, ...) {
+        va_list args;
+        va_start(args, format);
+        log_v(level, format, args);
+        va_end(args);
+    }
+};
+
+// inline, so all TUs including this header share one instance
+inline mtmd_helper_logger g_logger;
+
+#define LOG_DBG(...) g_logger.log(GGML_LOG_LEVEL_DEBUG, __VA_ARGS__)
+#define LOG_INF(...) g_logger.log(GGML_LOG_LEVEL_INFO,  __VA_ARGS__)
+#define LOG_WRN(...) g_logger.log(GGML_LOG_LEVEL_WARN,  __VA_ARGS__)
+#define LOG_ERR(...) g_logger.log(GGML_LOG_LEVEL_ERROR, __VA_ARGS__)
+
+//
+// embd batch
+//
+
+// helper struct to make working with embd batch easier
+// note: this will be removed after llama_batch_ext refactoring
+struct decode_embd_batch {
+    int n_pos_per_embd;
+    int n_mmproj_embd;
+    std::vector<llama_pos>      pos;
+    std::vector<llama_pos>      pos_view; // used by mrope
+    std::vector<int32_t>        n_seq_id;
+    std::vector<llama_seq_id>   seq_id_0;
+    std::vector<llama_seq_id *> seq_ids;
+    std::vector<int8_t>         logits;
+    llama_batch batch;
+    decode_embd_batch(float * embd, int32_t n_tokens, int n_pos_per_embd, int n_mmproj_embd) : n_pos_per_embd(n_pos_per_embd), n_mmproj_embd(n_mmproj_embd) {
+        GGML_ASSERT(n_tokens > 0 && n_pos_per_embd > 0 && n_mmproj_embd > 0);
+        pos     .resize(n_tokens * n_pos_per_embd);
+        n_seq_id.resize(n_tokens);
+        seq_ids .resize(n_tokens + 1);
+        logits  .resize(n_tokens);
+        seq_id_0.resize(1);
+        seq_ids [n_tokens] = nullptr;
+        batch = {
+            /*n_tokens       =*/ n_tokens,
+            /*tokens         =*/ nullptr,
+            /*embd           =*/ embd,
+            /*pos            =*/ pos.data(),
+            /*n_seq_id       =*/ n_seq_id.data(),
+            /*seq_id         =*/ seq_ids.data(),
+            /*logits         =*/ logits.data(),
+        };
+    }
+
+    void set_position_normal(llama_pos pos_0, llama_seq_id seq_id) {
+        seq_id_0[0] = seq_id;
+        for (int i = 0; i < batch.n_tokens; i++) {
+            batch.pos     [i] = pos_0 + i;
+            batch.n_seq_id[i] = 1;
+            batch.seq_id  [i] = seq_id_0.data();
+            batch.logits  [i] = false;
+        }
+    }
+
+    // M-RoPE for image
+    void set_position_mrope_2d(const std::vector<mtmd_decoder_pos> & rel_pos, llama_seq_id seq_id) {
+        GGML_ASSERT(n_pos_per_embd == 4);
+        GGML_ASSERT(!rel_pos.empty() && (int32_t)rel_pos.size() == batch.n_tokens);
+        seq_id_0[0] = seq_id;
+        for (int32_t i = 0; i < batch.n_tokens; i++) {
+            pos[i                     ] = rel_pos[i].t;
+            pos[i + batch.n_tokens    ] = rel_pos[i].y;
+            pos[i + batch.n_tokens * 2] = rel_pos[i].x;
+            pos[i + batch.n_tokens * 3] = rel_pos[i].z;
+        }
+        for (int i = 0; i < batch.n_tokens; i++) {
+            batch.n_seq_id[i] = 1;
+            batch.seq_id  [i] = seq_id_0.data();
+            batch.logits  [i] = false;
+        }
+    }
+
+    // M-RoPE for audio
+    void set_position_mrope_1d(llama_pos pos_0, llama_seq_id seq_id) {
+        GGML_ASSERT(n_pos_per_embd == 4);
+        seq_id_0[0] = seq_id;
+        for (int i = 0; i < batch.n_tokens; i++) {
+            pos[i                     ] = pos_0 + i;
+            pos[i + batch.n_tokens    ] = pos_0 + i;
+            pos[i + batch.n_tokens * 2] = pos_0 + i;
+            pos[i + batch.n_tokens * 3] = pos_0 + i;
+        }
+        for (int i = 0; i < batch.n_tokens; i++) {
+            batch.n_seq_id[i] = 1;
+            batch.seq_id  [i] = seq_id_0.data();
+            batch.logits  [i] = false;
+        }
+    }
+
+    llama_batch get_view(int offset, int n_tokens) {
+        GGML_ASSERT(offset >= 0 && n_tokens > 0 && offset + n_tokens <= batch.n_tokens);
+        llama_pos * pos_ptr;
+        pos_view.clear();
+        pos_view.reserve(n_tokens * n_pos_per_embd);
+        if (n_pos_per_embd > 1) {
+            // mrope
+            // for example, with layout of src: 1234...1234...1234...1234...
+            //       offset 2 will give us dst: 34...34...34...34...
+            for (int i = 0; i < n_pos_per_embd; i++) {
+                // assume n_tokens is less than or equal to batch.n_tokens
+                // batch.n_tokens is number of **total** tokens
+                // n_tokens is number of viewed token
+                size_t src_idx = i * batch.n_tokens + offset;
+                pos_view.insert(pos_view.end(),
+                    pos.data() + src_idx,
+                    pos.data() + src_idx + n_tokens);
+            }
+            pos_ptr = pos_view.data();
+        } else {
+            // normal
+            pos_ptr = pos.data() + offset;
+        }
+        return {
+            /*n_tokens       =*/ n_tokens,
+            /*tokens         =*/ nullptr,
+            /*embd           =*/ batch.embd     + offset * n_mmproj_embd,
+            /*pos            =*/ pos_ptr,
+            /*n_seq_id       =*/ batch.n_seq_id + offset,
+            /*seq_id         =*/ batch.seq_id   + offset,
+            /*logits         =*/ batch.logits   + offset,
+        };
+    }
+};
diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp
new file mode 100644 (file)
index 0000000..b52dc8e
--- /dev/null
@@ -0,0 +1,505 @@
+#include "mtmd.h"
+#include "mtmd-helper.h"
+#include "mtmd-helper-common.h"
+#include "llama.h"
+#include "../src/llama-ext.h"
+
+#include <algorithm>
+#include <cstring>
+#include <memory>
+#include <string>
+#include <unordered_map>
+#include <vector>
+
+#ifdef MTMD_INTERNAL_HEADER
+#error "mtmd-helper is a public library outside of mtmd. it must not include internal headers"
+#endif
+
+//
+// Audio generation helpers
+//
+
+// --tts-lang codes -> language names used by the codec_language special tokens
+static const std::unordered_map<std::string, std::string> tts_lang_codes = {
+    { "zh", "chinese"    },
+    { "en", "english"    },
+    { "de", "german"     },
+    { "it", "italian"    },
+    { "pt", "portuguese" },
+    { "es", "spanish"    },
+    { "ja", "japanese"   },
+    { "ko", "korean"     },
+    { "fr", "french"     },
+    { "ru", "russian"    },
+};
+
+static std::string tts_resolve_lang(const std::string & lang) {
+    auto it = tts_lang_codes.find(lang);
+    return it != tts_lang_codes.end() ? it->second : lang;
+}
+
+static llama_token find_special_token(const llama_vocab * vocab, const std::string & piece) {
+    const int32_t n = llama_vocab_n_tokens(vocab);
+    for (llama_token t = 0; t < n; t++) {
+        if (piece == llama_vocab_get_text(vocab, t)) {
+            return t;
+        }
+    }
+    return LLAMA_TOKEN_NULL;
+}
+
+static bool write_wav16(std::vector<char> & buf, const std::vector<float> & pcm, int32_t rate) {
+    // RIFF chunk sizes are 32-bit; refuse to emit a file with a truncated header
+    if (pcm.size() > ((size_t) UINT32_MAX - 36) / 2) {
+        return false;
+    }
+    const uint32_t data_sz   = (uint32_t) (pcm.size() * 2);
+    const uint32_t riff_sz   = 36 + data_sz;
+    const uint32_t fmt_sz    = 16, byte_rate = (uint32_t) rate * 2;
+    const uint16_t fmt = 1, ch = 1, align = 2, bits = 16;
+    const uint32_t rate32    = (uint32_t) rate;
+    auto put = [&](const void * p, size_t n) {
+        const char * c = (const char *) p;
+        buf.insert(buf.end(), c, c + n);
+    };
+    put("RIFF", 4); put(&riff_sz, 4); put("WAVE", 4);
+    put("fmt ", 4); put(&fmt_sz, 4);
+    put(&fmt, 2); put(&ch, 2); put(&rate32, 4);
+    put(&byte_rate, 4); put(&align, 2); put(&bits, 2);
+    put("data", 4); put(&data_sz, 4);
+    for (float v : pcm) {
+        int16_t s = (int16_t) (std::max(-1.0f, std::min(1.0f, v)) * 32767.0f);
+        put(&s, 2);
+    }
+    return true;
+}
+
+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;
+    // decodes at most n_batch prompt tokens; returns remaining count (0 = done), <0 on error
+    virtual int32_t step_prompt(int32_t n_batch) = 0;
+    // sampled can be LLAMA_TOKEN_NULL for pipelines with no discrete backbone token,
+    // those read what they need from h_state_in instead
+    virtual int32_t step_gen(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, int64_t * out_n_samples) = 0;
+
+protected:
+    llama_context * lctx;
+    mtmd_context  * mctx;
+    const llama_model * model;
+    const llama_vocab  * vocab;
+    int n_embd;
+    mtmd_gen_audio_info info;
+};
+
+// Qwen3-TTS: backbone samples codec_0, code_predictor gives the other 15 codebooks,
+// then code2wav decodes them to PCM
+class qwen3tts_gen_audio_pipeline : public mtmd_gen_audio_pipeline {
+public:
+    using mtmd_gen_audio_pipeline::mtmd_gen_audio_pipeline;
+
+    void reset() override {
+        seq_id = 0;
+        pos = 0;
+        codes_buf.clear();
+        c2w_state.clear();
+        audio_pcm.clear();
+        overlay.clear();
+        overlay_idx = 0;
+        h_state_buf.clear();
+        out_buf.clear();
+        prompt_embd_buf.clear();
+        prompt_batch.reset();
+        n_prompt = 0;
+        prompt_pos = 0;
+    }
+
+    int32_t set_input(const mtmd_helper_gen_audio_inp * inp) override {
+        reset();
+        seq_id = inp->seq_id;
+
+        if (!ensure_cache()) {
+            return 1;
+        }
+
+        const std::string lang   = tts_resolve_lang((inp->lang && inp->lang[0]) ? 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<float> 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<float>(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<float> 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<float> & vb) {
+            std::vector<float> 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<llama_token> 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<std::vector<float>> 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));
+
+        n_prompt = (int) prompt.size();
+
+        // the talker uses the qwen3vl interleaved mrope, all sections are equal for a 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;
+
+        prompt_embd_buf.resize((size_t) n_prompt * (size_t) n_e);
+        for (int i = 0; i < n_prompt; i++) {
+            memcpy(prompt_embd_buf.data() + (size_t) i * n_e, prompt[(size_t) i].data(), (size_t) n_e * sizeof(float));
+        }
+
+        prompt_batch.reset(new decode_embd_batch(prompt_embd_buf.data(), n_prompt, n_pos_per_embd, n_e));
+        if (mrope) prompt_batch->set_position_mrope_1d(0, seq_id);
+        else       prompt_batch->set_position_normal  (0, seq_id);
+        prompt_pos = 0;
+
+        pos = 0;
+        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: after frame k, the input adds
+        // trailing text row k on top of the codes embedding, then tts_eos, then tts_pad
+        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_prompt(int32_t n_batch) override {
+        GGML_ASSERT(n_batch > 0);
+        if (prompt_pos >= n_prompt) {
+            return 0;
+        }
+        const int32_t n_tokens_batch = std::min(n_batch, n_prompt - prompt_pos);
+        llama_batch batch_view = prompt_batch->get_view(prompt_pos, n_tokens_batch);
+
+        const bool is_last_batch = (prompt_pos + n_tokens_batch) == n_prompt;
+        if (is_last_batch) {
+            batch_view.logits[n_tokens_batch - 1] = 1;
+        }
+
+        if (llama_decode(lctx, batch_view) != 0) {
+            LOG_ERR("mtmd_helper_gen_audio: prompt decode failed\n");
+            return -1;
+        }
+
+        pos        += n_tokens_batch;
+        prompt_pos += n_tokens_batch;
+
+        if (prompt_pos >= n_prompt) {
+            // prompt fully processed, its embedding buffer is no longer needed
+            prompt_batch.reset();
+            prompt_embd_buf.clear();
+            return 0;
+        }
+        return n_prompt - prompt_pos;
+    }
+
+    int32_t step_gen(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<float *>(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_gen_wav()) {
+                return 1;
+            }
+        }
+
+        std::vector<float> 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, seq_id);
+        else       batch_embd.set_position_normal  (pos, seq_id);
+        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, int64_t * out_n_samples) override {
+        if (!flush_gen_wav()) {
+            return 1;
+        }
+
+        *out_sample_rate = info.sample_rate;
+        if (out_n_samples) {
+            *out_n_samples = (int64_t) audio_pcm.size();
+        }
+
+        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();
+        if (!write_wav16(out_buf, audio_pcm, info.sample_rate)) {
+            LOG_ERR("mtmd_helper_gen_audio: output too large for WAV\n");
+            return 1;
+        }
+        *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_pad>");
+        tts_bos   = find_special_token(vocab, "<tts_text_bos>");
+        tts_eos   = find_special_token(vocab, "<tts_text_eod>");
+        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);
+        if (llama_model_get_tok_embd(model, tok_embd.data()) != n_tok_embd) {
+            LOG_ERR("mtmd_helper_gen_audio: token embedding copy failed\n");
+            return false;
+        }
+        specials_ok = true;
+        return true;
+    }
+
+    // runs the reference wav through the speaker encoder, returns one x-vector embedding row
+    bool encode_speaker(mtmd_bitmap * bitmap, std::vector<float> & 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;
+    }
+
+    // one GEN_WAV process() call over the buffered codes, state is carried across batches
+    bool flush_gen_wav() {
+        if (codes_buf.empty()) {
+            return true;
+        }
+        mtmd_gen_inp inp{};
+        inp.type       = MTMD_GEN_PROCESS_TYPE_GEN_WAV;
+        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: gen_wav 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;
+    llama_token codec_eos  = LLAMA_TOKEN_NULL;
+    llama_token codec_pad  = LLAMA_TOKEN_NULL;
+    llama_token c_think    = LLAMA_TOKEN_NULL;
+    llama_token c_think_b  = LLAMA_TOKEN_NULL;
+    llama_token c_think_e  = LLAMA_TOKEN_NULL;
+    llama_token tts_pad    = LLAMA_TOKEN_NULL;
+    llama_token tts_bos    = LLAMA_TOKEN_NULL;
+    llama_token tts_eos    = LLAMA_TOKEN_NULL;
+    std::vector<float> tok_embd; // whole token embedding matrix, n_vocab * n_embd
+
+    // must match hparams.wav_tfm_swa hardcoded in clip.cpp
+    size_t window_frames = 72;
+
+    // per-generation state, cleared by reset()
+    llama_seq_id seq_id = 0;
+    bool mrope = false;
+    int pos = 0;
+    // prompt decode state, consumed batch-by-batch by step_prompt()
+    std::vector<float> prompt_embd_buf;
+    std::unique_ptr<decode_embd_batch> prompt_batch;
+    int n_prompt = 0;
+    int prompt_pos = 0;
+    int32_t top_k = 50;
+    float   top_p = 1.0f;
+    std::vector<int32_t> codes_buf;
+    std::vector<uint8_t> c2w_state;
+    std::vector<float>   audio_pcm;
+    std::vector<std::vector<float>> overlay;
+    size_t overlay_idx = 0;
+    std::vector<float> h_state_buf;
+    mtmd_helper_gen_audio_outtype out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
+    std::vector<char> out_buf;
+};
+
+static std::unique_ptr<mtmd_gen_audio_pipeline> 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<mtmd_gen_audio_pipeline>(new qwen3tts_gen_audio_pipeline(lctx, mctx));
+        default:
+            return nullptr;
+    }
+}
+
+struct mtmd_helper_gen_audio {
+    std::unique_ptr<mtmd_gen_audio_pipeline> 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->pipeline = make_pipeline(lctx, mctx);
+    return ctx;
+}
+
+void mtmd_helper_gen_audio_free(mtmd_helper_gen_audio * ctx) {
+    delete ctx;
+}
+
+void mtmd_helper_gen_audio_reset(mtmd_helper_gen_audio * ctx) {
+    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) {
+    if (!ctx->pipeline) {
+        LOG_ERR("mtmd_helper_gen_audio: unsupported or missing gen-audio pipeline\n");
+        return 1;
+    }
+    return ctx->pipeline->set_input(inp);
+}
+
+int32_t mtmd_helper_gen_audio_step_prompt(mtmd_helper_gen_audio * ctx, int32_t n_batch) {
+    if (!ctx->pipeline) {
+        return -1;
+    }
+    return ctx->pipeline->step_prompt(n_batch);
+}
+
+int32_t mtmd_helper_gen_audio_step_gen(mtmd_helper_gen_audio * ctx, llama_token sampled,
+                                       const float * h_state_in, const float ** h_state_out) {
+    if (!ctx->pipeline) {
+        return 1;
+    }
+    return ctx->pipeline->step_gen(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, int64_t * out_n_samples) {
+    if (!ctx->pipeline) {
+        return 1;
+    }
+    return ctx->pipeline->get_output(out_sample_rate, out_data, out_data_len, out_n_samples);
+}
index 90451d02ebd85b17899ffb0d912d6952d7335c07..d77c93966471b27709c2c9fa7ffcf933ba02c930 100644 (file)
@@ -9,6 +9,7 @@
 
 #include "mtmd.h"
 #include "mtmd-helper.h"
+#include "mtmd-helper-common.h"
 #include "llama.h"
 
 #include <algorithm>
 // internal logging functions
 //
 
-struct mtmd_helper_logger {
-    ggml_log_callback default_callback = [](ggml_log_level level, const char * text, void * user_data) {
-        (void) level;
-        (void) user_data;
-        fputs(text, stderr);
-        fflush(stderr);
-    };
-
-    ggml_log_callback log_callback = default_callback;
-    void * log_callback_user_data;
-
-    void log_v(enum ggml_log_level level, const char * format, va_list args) {
-        if (format == NULL) {
-            return;
-        }
-        va_list args_copy;
-        va_copy(args_copy, args);
-        char buffer[128];
-        int len = vsnprintf(buffer, 128, format, args);
-        if (len < 128) {
-            log_callback(level, buffer, log_callback_user_data);
-        } else {
-            char * buffer2 = (char *) calloc(len + 1, sizeof(char));
-            vsnprintf(buffer2, len + 1, format, args_copy);
-            buffer2[len] = 0;
-            log_callback(level, buffer2, log_callback_user_data);
-            free(buffer2);
-        }
-        va_end(args_copy);
-    }
-
-    void log(enum ggml_log_level level, const char * format, ...) {
-        va_list args;
-        va_start(args, format);
-        log_v(level, format, args);
-        va_end(args);
-    }
-} g_logger;
-
-#define LOG_DBG(...) g_logger.log(GGML_LOG_LEVEL_DEBUG, __VA_ARGS__)
-#define LOG_INF(...) g_logger.log(GGML_LOG_LEVEL_INFO,  __VA_ARGS__)
-#define LOG_WRN(...) g_logger.log(GGML_LOG_LEVEL_WARN,  __VA_ARGS__)
-#define LOG_ERR(...) g_logger.log(GGML_LOG_LEVEL_ERROR, __VA_ARGS__)
-
 void mtmd_helper_log_set(ggml_log_callback log_callback, void * user_data) {
     if (log_callback == nullptr) {
         log_callback = g_logger.default_callback;
@@ -127,117 +84,6 @@ void mtmd_helper_image_get_decoder_pos(const mtmd_image_tokens * chunks, llama_p
     }
 }
 
-// helper struct to make working with embd batch easier
-// note: this will be removed after llama_batch_ext refactoring
-struct decode_embd_batch {
-    int n_pos_per_embd;
-    int n_mmproj_embd;
-    std::vector<llama_pos>      pos;
-    std::vector<llama_pos>      pos_view; // used by mrope
-    std::vector<int32_t>        n_seq_id;
-    std::vector<llama_seq_id>   seq_id_0;
-    std::vector<llama_seq_id *> seq_ids;
-    std::vector<int8_t>         logits;
-    llama_batch batch;
-    decode_embd_batch(float * embd, int32_t n_tokens, int n_pos_per_embd, int n_mmproj_embd) : n_pos_per_embd(n_pos_per_embd), n_mmproj_embd(n_mmproj_embd) {
-        GGML_ASSERT(n_tokens > 0 && n_pos_per_embd > 0 && n_mmproj_embd > 0);
-        pos     .resize(n_tokens * n_pos_per_embd);
-        n_seq_id.resize(n_tokens);
-        seq_ids .resize(n_tokens + 1);
-        logits  .resize(n_tokens);
-        seq_id_0.resize(1);
-        seq_ids [n_tokens] = nullptr;
-        batch = {
-            /*n_tokens       =*/ n_tokens,
-            /*tokens         =*/ nullptr,
-            /*embd           =*/ embd,
-            /*pos            =*/ pos.data(),
-            /*n_seq_id       =*/ n_seq_id.data(),
-            /*seq_id         =*/ seq_ids.data(),
-            /*logits         =*/ logits.data(),
-        };
-    }
-
-    void set_position_normal(llama_pos pos_0, llama_seq_id seq_id) {
-        seq_id_0[0] = seq_id;
-        for (int i = 0; i < batch.n_tokens; i++) {
-            batch.pos     [i] = pos_0 + i;
-            batch.n_seq_id[i] = 1;
-            batch.seq_id  [i] = seq_id_0.data();
-            batch.logits  [i] = false;
-        }
-    }
-
-    // M-RoPE for image
-    void set_position_mrope_2d(const std::vector<mtmd_decoder_pos> & rel_pos, llama_seq_id seq_id) {
-        GGML_ASSERT(n_pos_per_embd == 4);
-        GGML_ASSERT(!rel_pos.empty() && (int32_t)rel_pos.size() == batch.n_tokens);
-        seq_id_0[0] = seq_id;
-        for (int32_t i = 0; i < batch.n_tokens; i++) {
-            pos[i                     ] = rel_pos[i].t;
-            pos[i + batch.n_tokens    ] = rel_pos[i].y;
-            pos[i + batch.n_tokens * 2] = rel_pos[i].x;
-            pos[i + batch.n_tokens * 3] = rel_pos[i].z;
-        }
-        for (int i = 0; i < batch.n_tokens; i++) {
-            batch.n_seq_id[i] = 1;
-            batch.seq_id  [i] = seq_id_0.data();
-            batch.logits  [i] = false;
-        }
-    }
-
-    // M-RoPE for audio
-    void set_position_mrope_1d(llama_pos pos_0, llama_seq_id seq_id) {
-        GGML_ASSERT(n_pos_per_embd == 4);
-        seq_id_0[0] = seq_id;
-        for (int i = 0; i < batch.n_tokens; i++) {
-            pos[i                     ] = pos_0 + i;
-            pos[i + batch.n_tokens    ] = pos_0 + i;
-            pos[i + batch.n_tokens * 2] = pos_0 + i;
-            pos[i + batch.n_tokens * 3] = pos_0 + i;
-        }
-        for (int i = 0; i < batch.n_tokens; i++) {
-            batch.n_seq_id[i] = 1;
-            batch.seq_id  [i] = seq_id_0.data();
-            batch.logits  [i] = false;
-        }
-    }
-
-    llama_batch get_view(int offset, int n_tokens) {
-        GGML_ASSERT(offset >= 0 && n_tokens > 0 && offset + n_tokens <= batch.n_tokens);
-        llama_pos * pos_ptr;
-        pos_view.clear();
-        pos_view.reserve(n_tokens * n_pos_per_embd);
-        if (n_pos_per_embd > 1) {
-            // mrope
-            // for example, with layout of src: 1234...1234...1234...1234...
-            //       offset 2 will give us dst: 34...34...34...34...
-            for (int i = 0; i < n_pos_per_embd; i++) {
-                // assume n_tokens is less than or equal to batch.n_tokens
-                // batch.n_tokens is number of **total** tokens
-                // n_tokens is number of viewed token
-                size_t src_idx = i * batch.n_tokens + offset;
-                pos_view.insert(pos_view.end(),
-                    pos.data() + src_idx,
-                    pos.data() + src_idx + n_tokens);
-            }
-            pos_ptr = pos_view.data();
-        } else {
-            // normal
-            pos_ptr = pos.data() + offset;
-        }
-        return {
-            /*n_tokens       =*/ n_tokens,
-            /*tokens         =*/ nullptr,
-            /*embd           =*/ batch.embd     + offset * n_mmproj_embd,
-            /*pos            =*/ pos_ptr,
-            /*n_seq_id       =*/ batch.n_seq_id + offset,
-            /*seq_id         =*/ batch.seq_id   + offset,
-            /*logits         =*/ batch.logits   + offset,
-        };
-    }
-};
-
 // Helper class to set non-causal attention via RAII
 class scope_non_causal {
 public:
@@ -1084,3 +930,18 @@ int32_t mtmd_helper_video_read_next(mtmd_helper_video * ctx,
     GGML_ASSERT(false && "video is not supported in this build (MTMD_VIDEO is set to OFF)");
 #endif
 }
+
+bool mtmd_helper_model_can_chat(llama_context * lctx, mtmd_context * mctx) {
+    if (!mctx) {
+        return true;
+    }
+
+    auto * model = llama_get_model(lctx);
+    auto * tmpl = llama_model_chat_template(model, nullptr);
+    auto info = mtmd_gen_audio_get_info(mctx);
+
+    // tts-only model cannot be used for chat (no chat template)
+    bool is_tts_only = info.type != MTMD_GEN_AUDIO_TYPE_NONE && tmpl == nullptr;
+
+    return !is_tts_only;
+}
index 680a2317df0e040bf5f7b3666c5226c955d5b19d..7e5cf9b5098c657960f7ecec54720a0b9f14ede0 100644 (file)
@@ -157,6 +157,73 @@ MTMD_API int32_t mtmd_helper_video_read_next(mtmd_helper_video * ctx,
             mtmd_bitmap ** out_bitmap,
             char ** out_text);
 
+// return true if model can be used for chat
+MTMD_API bool mtmd_helper_model_can_chat(struct llama_context * lctx, struct mtmd_context * mctx);
+
+//
+// Audio generation helpers
+// (early-stage experimental, subjected to breaking changes)
+//
+
+// audio generation helper context
+// contains accumulator for generated audio features and PCM audio
+struct mtmd_helper_gen_audio;
+typedef struct mtmd_helper_gen_audio mtmd_helper_gen_audio;
+
+enum mtmd_helper_gen_audio_outtype {
+    MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM, // raw PCM
+    MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV, // WAV PCM 16-bit LE, mono
+};
+struct mtmd_helper_gen_audio_inp {
+    llama_seq_id seq_id;
+
+    const char * prompt;
+    size_t       prompt_len;
+
+    mtmd_bitmap * speaker_ref; // optional, can be NULL
+    const char * lang; // optional, can be NULL
+
+    int32_t top_k;
+    float   top_p;
+
+    enum mtmd_helper_gen_audio_outtype out_type;
+};
+
+MTMD_API mtmd_helper_gen_audio * mtmd_helper_gen_audio_init(
+                                    struct llama_context * lctx,
+                                    struct mtmd_context * mctx);
+
+MTMD_API void mtmd_helper_gen_audio_free(mtmd_helper_gen_audio * ctx);
+
+MTMD_API void mtmd_helper_gen_audio_reset(mtmd_helper_gen_audio * ctx);
+
+MTMD_API int32_t mtmd_helper_gen_audio_set_input(
+                        mtmd_helper_gen_audio * ctx,
+                        const struct mtmd_helper_gen_audio_inp * inp);
+
+// processes at most n_batch prompt tokens per call
+// returns: >0 = number of prompt tokens remaining, 0 = done, <0 = error
+MTMD_API int32_t mtmd_helper_gen_audio_step_prompt(
+                        mtmd_helper_gen_audio * ctx,
+                        int32_t n_batch);
+
+// generates one frame; must only be called after step_prompt() has returned 0
+// h_state_out is valid until next step_gen() or reset() call
+MTMD_API int32_t mtmd_helper_gen_audio_step_gen(
+                        mtmd_helper_gen_audio * ctx,
+                        llama_token sampled,
+                        const float *  h_state_in,
+                        const float ** h_state_out);
+
+// out_data valid until next get_output() or reset() call
+// out_n_samples (optional, can be NULL) receives the number of generated PCM samples
+MTMD_API 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,
+                        int64_t * out_n_samples);
+
 #ifdef __cplusplus
 } // extern "C"
 #endif
@@ -177,6 +244,31 @@ struct mtmd_helper_video_deleter {
 };
 using video_ptr = std::unique_ptr<mtmd_helper_video, mtmd_helper_video_deleter>;
 
+// audio generation-related C++ wrappers
+struct mtmd_helper_gen_audio_deleter {
+    void operator()(mtmd_helper_gen_audio * val) { mtmd_helper_gen_audio_free(val); }
+};
+using gen_audio_ptr = std::unique_ptr<mtmd_helper_gen_audio, mtmd_helper_gen_audio_deleter>;
+struct gen_audio {
+    gen_audio_ptr ctx;
+    gen_audio(struct llama_context * lctx, struct mtmd_context * mctx) : ctx(mtmd_helper_gen_audio_init(lctx, mctx)) {}
+    void reset() {
+        mtmd_helper_gen_audio_reset(ctx.get());
+    }
+    int32_t set_input(const struct mtmd_helper_gen_audio_inp * inp) {
+        return mtmd_helper_gen_audio_set_input(ctx.get(), inp);
+    }
+    int32_t step_prompt(int32_t n_batch) {
+        return mtmd_helper_gen_audio_step_prompt(ctx.get(), n_batch);
+    }
+    int32_t step_gen(llama_token sampled, const float * h_state, const float ** h_state_out) {
+        return mtmd_helper_gen_audio_step_gen(ctx.get(), sampled, h_state, h_state_out);
+    }
+    int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples = nullptr) {
+        return mtmd_helper_gen_audio_get_output(ctx.get(), out_sample_rate, out_data, out_data_len, out_n_samples);
+    }
+};
+
 } // namespace mtmd_helper
 #endif
 
index d3899f5c853d243e169c03004e5c2251c17f9be4..ff90d6818cb5f4c1e8204defb9af4cd3a031b5b4 100644 (file)
@@ -262,6 +262,13 @@ struct mtmd_context {
     struct clip_ctx * ctx_a; // audio
     std::vector<float> out_embd; // image embedding vector
 
+    // generation context
+    struct clip_ctx * ctx_gen_a; // audio
+    std::vector<int32_t> gen_out_codes; // this frame's 16 sampled codes (GEN_CODE)
+    std::vector<float>   gen_out_embd;  // next-step hidden state fed back to backbone (GEN_CODE)
+    std::vector<float>   gen_out_audio; // decoded PCM samples for the current frame (GEN_WAV)
+    std::vector<uint8_t> gen_out_state; // state to feed into the next GEN_WAV call
+
     bool print_timings;
     int n_threads;
     std::string media_marker;
@@ -354,6 +361,7 @@ struct mtmd_context {
         auto res = clip_init(mmproj_fname, ctx_clip_params);
         ctx_v = res.ctx_v;
         ctx_a = res.ctx_a;
+        ctx_gen_a = res.ctx_gen_a;
         if (!ctx_v && !ctx_a) {
             throw std::runtime_error(string_format("Failed to load CLIP model from %s\n", mmproj_fname));
         }
@@ -378,6 +386,15 @@ struct mtmd_context {
                 "hint: you may be using wrong mmproj\n",
                 n_embd_text, n_embd_clip));
         }
+        if (ctx_gen_a) {
+            int n_embd_gen = clip_n_mmproj_embd(ctx_gen_a);
+            if (n_embd_text > 0 && n_embd_text != n_embd_gen) {
+                throw std::runtime_error(string_format(
+                    "mismatch between text model (n_embd = %d) and gen-audio mmproj (n_embd = %d)\n"
+                    "hint: you may be using wrong mmproj\n",
+                    n_embd_text, n_embd_gen));
+            }
+        }
         if (ctx_v) {
             init_vision();
         }
@@ -740,6 +757,10 @@ struct mtmd_context {
                     aud_end = "<|mimo_audio_end|>";
                     audio_preproc = std::make_unique<mtmd_audio_preprocessor_mimo_audio>(ctx_a);
                 } break;
+            case PROJECTOR_TYPE_QWEN3TTS_SPKENC:
+                {
+                    audio_preproc = std::make_unique<mtmd_audio_preprocessor_qwen3tts_spk>(ctx_a);
+                } break;
             default:
                 throw std::runtime_error(string_format("%s: unexpected audio projector type %d\n", __func__, proj));
         }
@@ -780,6 +801,7 @@ struct mtmd_context {
     ~mtmd_context() {
         clip_free(ctx_a);
         clip_free(ctx_v);
+        clip_free(ctx_gen_a);
     }
 
 private:
@@ -1553,6 +1575,125 @@ float * mtmd_get_output_embd(mtmd_context * ctx) {
     return ctx->out_embd.data();
 }
 
+//
+// audio generation
+//
+
+mtmd_gen_audio_info mtmd_gen_audio_get_info(const mtmd_context * ctx) {
+    mtmd_gen_audio_info info;
+    if (!ctx->ctx_gen_a) {
+        info.type = MTMD_GEN_AUDIO_TYPE_NONE;
+        return info;
+    }
+    switch (clip_get_projector_type(ctx->ctx_gen_a)) {
+        case PROJECTOR_TYPE_QWEN3TTS_GEN:
+            info.type = MTMD_GEN_AUDIO_TYPE_QWEN3TTS;
+            info.sample_rate = 24000;
+            break;
+        default:
+            info.type = MTMD_GEN_AUDIO_TYPE_NONE;
+            break;
+    }
+    return info;
+}
+
+static int32_t mtmd_gen_audio_process_impl(mtmd_context * ctx, const mtmd_gen_inp * inp, mtmd_gen_out * out) {
+    clip_ctx * ctx_clip = ctx->ctx_gen_a;
+    if (!ctx_clip) {
+        LOG_ERR("%s: model does not support audio generation\n", __func__);
+        return 1;
+    }
+
+    if (inp->type == MTMD_GEN_PROCESS_TYPE_GEN_CODE) {
+        const size_t n_embd = (size_t) clip_n_mmproj_embd(ctx_clip);
+
+        clip_image_f32 hidden_state;
+        hidden_state.set_size({(int) n_embd, 1}, false, true);
+        hidden_state.cpy_buf(std::vector<float>(inp->embd, inp->embd + n_embd));
+
+        clip_image_f32_batch batch;
+        batch.is_audio = true;
+        batch.entries.push_back(std::move(hidden_state));
+
+        std::vector<float>   out_embd(n_embd);
+        std::vector<int32_t> out_codes;
+
+        clip_encode_params params;
+        params.imgs        = &batch;
+        params.n_threads   = ctx->n_threads;
+        params.gen_process = CLIP_GEN_PROCESS_GEN_CODE;
+        params.out_embd    = &out_embd;
+        params.out_codes   = &out_codes;
+        params.code0       = inp->code0;
+        params.top_k       = inp->top_k;
+        params.top_p       = inp->top_p;
+
+        if (!clip_encode(ctx_clip, &params)) {
+            LOG_ERR("%s: clip_encode failed (gen_code)\n", __func__);
+            return 1;
+        }
+
+        ctx->gen_out_embd  = std::move(out_embd);
+        ctx->gen_out_codes = std::move(out_codes);
+
+        out->embd    = ctx->gen_out_embd.data();
+        out->codes   = ctx->gen_out_codes.data();
+        out->n_codes = ctx->gen_out_codes.size();
+        return 0;
+    }
+
+    // MTMD_GEN_PROCESS_TYPE_GEN_WAV
+    if (!inp->codes || inp->n_codes == 0) {
+        LOG_ERR("%s: codes required for gen_wav\n", __func__);
+        return 1;
+    }
+    std::vector<int32_t> in_codes(inp->codes, inp->codes + inp->n_codes);
+    std::vector<uint8_t> in_state;
+    if (inp->state_data) {
+        in_state.assign(inp->state_data, inp->state_data + inp->state_size);
+    }
+
+    // gen_wav has no hidden-state input, the batch entry is an unused placeholder
+    // TODO @ngxson : some models in the future may require hidden-state input, need to update this code later
+    clip_image_f32 dummy;
+    dummy.set_size({1, 1}, false, true);
+    dummy.cpy_buf(std::vector<float>(1, 0.0f));
+
+    clip_image_f32_batch batch;
+    batch.is_audio = true;
+    batch.entries.push_back(std::move(dummy));
+
+    clip_encode_params params;
+    params.imgs        = &batch;
+    params.n_threads   = ctx->n_threads;
+    params.gen_process = CLIP_GEN_PROCESS_GEN_WAV;
+    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, &params)) {
+        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->state_data = (const char *) ctx->gen_out_state.data();
+    out->state_size = ctx->gen_out_state.size();
+
+    return 0;
+}
+
+int32_t mtmd_gen_audio_process(mtmd_context * ctx, const struct mtmd_gen_inp * inp, struct mtmd_gen_out * out) {
+    try {
+        return mtmd_gen_audio_process_impl(ctx, inp, out);
+    } catch (const std::exception & e) {
+        LOG_ERR("%s: error: %s\n", __func__, e.what());
+        return 1;
+    }
+}
+
 mtmd_batch * mtmd_batch_init(mtmd_context * ctx) {
     return new mtmd_batch(ctx);
 }
index 3b8c1200b5665345823d25c5311dd4206c8ae3cc..84651f8dcd00b5ed613d92a2b179b119dabec8fe 100644 (file)
@@ -327,6 +327,60 @@ struct mtmd_caps {
 };
 MTMD_API struct mtmd_caps mtmd_get_cap_from_file(const char * mmproj_fname);
 
+/////////////////////////////////////////
+// EXPERIMENTAL API for audio generation, subjected to breaking changes
+
+// represent the pipeline type
+enum mtmd_gen_audio_type {
+    MTMD_GEN_AUDIO_TYPE_NONE, // not supported
+    MTMD_GEN_AUDIO_TYPE_QWEN3TTS,
+};
+struct mtmd_gen_audio_info {
+    enum mtmd_gen_audio_type type;
+    int32_t sample_rate; // in Hz, for example 24000 for qwen3tts
+};
+MTMD_API struct mtmd_gen_audio_info mtmd_gen_audio_get_info(const mtmd_context * ctx);
+
+enum mtmd_gen_process_type {
+    MTMD_GEN_PROCESS_TYPE_GEN_CODE, // h_state to semantic (codes, mel-spectrogram, etc.)
+    MTMD_GEN_PROCESS_TYPE_GEN_WAV,  // convert semantic to PCM audio
+                                    // for qwen3tts, this is code2wav
+};
+struct mtmd_gen_inp {
+    enum mtmd_gen_process_type type;
+
+    // for MTMD_GEN_PROCESS_TYPE_GEN_CODE
+    int32_t code0;  // the sampled codebook 0 entry from backbone
+    float * embd;   // the hidden state from backbone, must have n_text_embd elements
+    int32_t top_k;
+    float   top_p;
+
+    // for MTMD_GEN_PROCESS_TYPE_GEN_WAV
+    int32_t * 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
+
+    // for MTMD_GEN_PROCESS_TYPE_GEN_CODE
+    const int32_t * codes;
+    size_t n_codes;
+    const float * embd; // the generated hidden state, to be fed back to backbone
+                        // it must have n_text_embd elements
+
+    // for MTMD_GEN_PROCESS_TYPE_GEN_WAV
+    const float * audio;
+    size_t        n_samples;
+    const char * state_data;
+    size_t       state_size;
+};
+// note: this API is stateless, caller must handle state management and audio frame accumulation
+MTMD_API int32_t mtmd_gen_audio_process(mtmd_context * ctx,
+                                const struct mtmd_gen_inp * inp,
+                                struct mtmd_gen_out * out);
+
 /////////////////////////////////////////
 
 // test function, to be used in test-mtmd-c-api.c
index 26a8bb8f2d1fb2be153f35d84926725343ce0746..0a0b5730eaa6b57017e19f0e08c41d65ff5e3d79 100644 (file)
@@ -1,6 +1,6 @@
 set(TARGET llama-tts)
 add_executable(${TARGET} tts.cpp)
-target_link_libraries(${TARGET} PRIVATE llama llama-common ${CMAKE_THREAD_LIBS_INIT})
+target_link_libraries(${TARGET} PRIVATE llama llama-common mtmd ${CMAKE_THREAD_LIBS_INIT})
 target_compile_features(${TARGET} PRIVATE cxx_std_17)
 
 if(LLAMA_TOOLS_INSTALL)
index 4749bb9f5a7214f487504e92e8774648b7eb9ed6..dd84336c39900a580f3d7e25b6ad58f0d6120f4d 100644 (file)
-# llama.cpp/example/tts
-This example demonstrates the Text To Speech feature. It uses a
-[model](https://www.outeai.com/blog/outetts-0.2-500m) from
-[outeai](https://www.outeai.com/).
+# llama.cpp TTS
 
-## Quickstart
-If you have built llama.cpp with SSL support you can simply run the
-following command and the required models will be downloaded automatically:
-```console
-$ build/bin/llama-tts --tts-oute-default -p "Hello world" && aplay output.wav
-```
-For details about the models and how to convert them to the required format
-see the following sections.
+This is a tool to demonstrate audio generation capability in llama.cpp via `libmtmd`. It was added via PR [#26254](https://github.com/ggml-org/llama.cpp/pull/26254)
 
-### Model conversion
-Checkout or download the model that contains the LLM model:
-```console
-$ pushd models
-$ git clone --branch main --single-branch --depth 1 https://huggingface.co/OuteAI/OuteTTS-0.2-500M
-$ cd OuteTTS-0.2-500M && git lfs install && git lfs pull
-$ popd
-```
-Convert the model to .gguf format:
-```console
-(venv) python convert_hf_to_gguf.py models/OuteTTS-0.2-500M \
-    --outfile models/outetts-0.2-0.5B-f16.gguf --outtype f16
-```
-The generated model will be `models/outetts-0.2-0.5B-f16.gguf`.
+Note: this tool used to serve as a demo for OuteTTS, but it was converted to a more model-agnostic tool.
 
-We can optionally quantize this to Q8_0 using the following command:
-```console
-$ build/bin/llama-quantize models/outetts-0.2-0.5B-f16.gguf \
-    models/outetts-0.2-0.5B-q8_0.gguf q8_0
-```
-The quantized model will be `models/outetts-0.2-0.5B-q8_0.gguf`.
-
-Next we do something similar for the audio decoder. First download or checkout
-the model for the voice decoder:
-```console
-$ pushd models
-$ git clone --branch main --single-branch --depth 1 https://huggingface.co/novateur/WavTokenizer-large-speech-75token
-$ cd WavTokenizer-large-speech-75token && git lfs install && git lfs pull
-$ popd
-```
-This model file is a PyTorch checkpoint (.ckpt) and we first need to convert it to
-huggingface format:
-```console
-(venv) python tools/tts/convert_pt_to_hf.py \
-    models/WavTokenizer-large-speech-75token/wavtokenizer_large_speech_320_24k.ckpt
-...
-Model has been successfully converted and saved to models/WavTokenizer-large-speech-75token/model.safetensors
-Metadata has been saved to models/WavTokenizer-large-speech-75token/index.json
-Config has been saved to models/WavTokenizer-large-speech-75tokenconfig.json
-```
-Then we can convert the huggingface format to gguf:
-```console
-(venv) python convert_hf_to_gguf.py models/WavTokenizer-large-speech-75token \
-    --outfile models/wavtokenizer-large-75-f16.gguf --outtype f16
-...
-INFO:hf-to-gguf:Model successfully exported to models/wavtokenizer-large-75-f16.gguf
-```
+## Common usage
 
-### Running the example
+Simple usage:
 
-With both of the models generated, the LLM model and the voice decoder model,
-we can run the example:
-```console
-$ build/bin/llama-tts -m  ./models/outetts-0.2-0.5B-q8_0.gguf \
-    -mv ./models/wavtokenizer-large-75-f16.gguf \
-    -p "Hello world"
-...
-main: audio written to file 'output.wav'
-```
-The output.wav file will contain the audio of the prompt. This can be heard
-by playing the file with a media player. On Linux the following command will
-play the audio:
-```console
-$ aplay output.wav
+```sh
+llama-tts -hf ggml-org/Qwen3-TTS-12Hz-1.7B-Base-GGUF -p "Hello world" --output out.wav
 ```
 
-### Running the example with llama-server
-Running this example with `llama-server` is also possible and requires two
-server instances to be started. One will serve the LLM model and the other
-will serve the voice decoder model.
+Common params:
+- Sampling params such as `--top-k`, `--top-p`, `--temp`, etc.
+- `-n <number_of_frames>` limits the output length, e.g. `-n 500`. Note that how many milliseconds each frame represents varies by model
+- Core inference params such as `-ngl`, `-b`, `-ub`, etc.
 
-The LLM model server can be started with the following command:
-```console
-$ ./build/bin/llama-server -m ./models/outetts-0.2-0.5B-q8_0.gguf --port 8020
-```
+## Qwen3-TTS
 
-And the voice decoder model server can be started using:
-```console
-./build/bin/llama-server -m ./models/wavtokenizer-large-75-f16.gguf --port 8021 --embeddings --pooling none
-```
+Available params:
+- `--tts-lang` can be `zh`, `en`, `de`, `it`, `pt`, `es`, `ja`, `ko`, `fr`, `ru` (default: `en`)
+- `--tts-speaker-file` should point to a speaker reference audio file (wav, mp3)
 
-Then we can run [tts-outetts.py](tts-outetts.py) to generate the audio.
+Example usage:
 
-First create a virtual environment for python and install the required
-dependencies (this in only required to be done once):
-```console
-$ python3 -m venv venv
-$ source venv/bin/activate
-(venv) pip install requests numpy
-```
-
-And then run the python script using:
-```conole
-(venv) python ./tools/tts/tts-outetts.py http://localhost:8020 http://localhost:8021 "Hello world"
-spectrogram generated: n_codes: 90, n_embd: 1282
-converting to audio ...
-audio generated: 28800 samples
-audio written to file "output.wav"
-```
-And to play the audio we can again use aplay or any other media player:
-```console
-$ aplay output.wav
+```sh
+llama-tts -hf ggml-org/Qwen3-TTS-12Hz-1.7B-Base-GGUF \
+    -p "Hello world" \
+    --tts-lang english \
+    --tts-speaker-file speaker.mp3 \
+    --output out.wav
 ```
diff --git a/tools/tts/convert_pt_to_hf.py b/tools/tts/convert_pt_to_hf.py
deleted file mode 100644 (file)
index ebd55d9..0000000
+++ /dev/null
@@ -1,180 +0,0 @@
-# convert the https://huggingface.co/novateur/WavTokenizer-large-speech-75token to HF format
-# the goal is to be able to reuse the convert_hf_to_gguf.py after that to create a GGUF file with the WavTokenizer decoder
-#
-# TODO: this script is LLM-generated and probably very inefficient and should be rewritten
-
-import torch
-import json
-import os
-import sys
-import re
-
-from safetensors.torch import save_file
-
-# default
-model_path = './model.pt'
-
-# read from CLI
-if len(sys.argv) > 1:
-    model_path = sys.argv[1]
-
-# get the directory of the input model
-path_dst = os.path.dirname(model_path)
-
-print(f"Loading model from {model_path}")
-
-model = torch.load(model_path, map_location='cpu')
-
-#print(model)
-
-# print all keys
-for key in model.keys():
-    print(key)
-    if key == 'hyper_parameters':
-        #print(model[key])
-        # dump as json pretty
-        print(json.dumps(model[key], indent=4))
-    #if key != 'state_dict' and key != 'optimizer_states':
-    #    print(model[key])
-
-# Check if the loaded model is a state_dict or a model instance
-if isinstance(model, torch.nn.Module):
-    state_dict = model.state_dict()
-else:
-    state_dict = model
-
-# Print the structure of the state_dict to understand its format
-print("State dictionary keys:")
-for key in state_dict.keys():
-    print(key)
-
-# Ensure the state_dict is flat and contains only torch.Tensor objects
-def flatten_state_dict(state_dict, parent_key='', sep='.'):
-    items = []
-    items_new = []
-
-    for k, v in state_dict.items():
-        new_key = f"{parent_key}{sep}{k}" if parent_key else k
-        if isinstance(v, torch.Tensor):
-            items.append((new_key, v))
-        elif isinstance(v, dict):
-            items.extend(flatten_state_dict(v, new_key, sep=sep).items())
-            return dict(items)
-
-    size_total_mb = 0
-
-    for key, value in list(items):
-        # keep only what we need for inference
-        if not key.startswith('state_dict.feature_extractor.encodec.quantizer.') and \
-           not key.startswith('state_dict.backbone.') and \
-           not key.startswith('state_dict.head.out'):
-               print('Skipping key: ', key)
-               continue
-
-        new_key = key
-
-        new_key = new_key.replace('state_dict.', '')
-        new_key = new_key.replace('pos_net', 'posnet')
-
-        # check if matches "backbone.posnet.%d.bias" or "backbone.posnet.%d.weight"
-        if new_key.startswith("backbone.posnet."):
-            match = re.match(r"backbone\.posnet\.(\d+)\.(bias|weight)", new_key)
-            if match:
-               new_key = f"backbone.posnet.{match.group(1)}.norm.{match.group(2)}"
-
-        # "feature_extractor.encodec.quantizer.vq.layers.0._codebook.embed" -> "backbone.embedding.weight"
-        if new_key == "feature_extractor.encodec.quantizer.vq.layers.0._codebook.embed":
-            new_key = "backbone.embedding.weight"
-
-        # these are the only rows used
-        # ref: https://github.com/edwko/OuteTTS/blob/a613e79c489d8256dd657ea9168d78de75895d82/outetts/wav_tokenizer/audio_codec.py#L100
-        if new_key.endswith("norm.scale.weight"):
-            new_key = new_key.replace("norm.scale.weight", "norm.weight")
-            value = value[0]
-
-        if new_key.endswith("norm.shift.weight"):
-            new_key = new_key.replace("norm.shift.weight", "norm.bias")
-            value = value[0]
-
-        if new_key.endswith("gamma"):
-            new_key = new_key.replace("gamma", "gamma.weight")
-
-        # convert from 1D [768] to 2D [768, 1] so that ggml_add can broadcast the bias
-        if (new_key.endswith("norm.weight") or new_key.endswith("norm1.weight") or new_key.endswith("norm2.weight") or new_key.endswith(".bias")) and (new_key.startswith("backbone.posnet") or new_key.startswith("backbone.embed.bias")):
-            value = value.unsqueeze(1)
-
-        if new_key.endswith("dwconv.bias"):
-            value = value.unsqueeze(1)
-
-        size_mb = value.element_size() * value.nelement() / (1024 * 1024)
-        print(f"{size_mb:8.2f} MB - {new_key}: {value.shape}")
-
-        size_total_mb += size_mb
-
-        #print(key, '->', new_key, ': ', value)
-        #print(key, '->', new_key)
-
-        items_new.append((new_key, value))
-
-    print(f"Total size: {size_total_mb:8.2f} MB")
-
-    return dict(items_new)
-
-flattened_state_dict = flatten_state_dict(state_dict)
-
-
-# Convert the model to the safetensors format
-output_path = path_dst + '/model.safetensors'
-save_file(flattened_state_dict, output_path)
-
-print(f"Model has been successfully converted and saved to {output_path}")
-
-# Calculate the total size of the .safetensors file
-total_size = os.path.getsize(output_path)
-
-# Create the weight map
-weight_map = {
-    "model.safetensors": ["*"]  # Assuming all weights are in one file
-}
-
-# Create metadata for the index.json file
-metadata = {
-    "total_size": total_size,
-    "weight_map": weight_map
-}
-
-# Save the metadata to index.json
-index_path = path_dst + '/index.json'
-with open(index_path, 'w') as f:
-    json.dump(metadata, f, indent=4)
-
-print(f"Metadata has been saved to {index_path}")
-
-config = {
-    "architectures": [
-        "WavTokenizerDec"
-    ],
-    "hidden_size": 1282,
-    "n_embd_features": 512,
-    "n_ff": 2304,
-    "vocab_size": 4096,
-    "n_head": 1,
-    "layer_norm_epsilon": 1e-6,
-    "group_norm_epsilon": 1e-6,
-    "group_norm_groups": 32,
-    "max_position_embeddings": 8192, # ?
-    "n_layer": 12,
-    "posnet": {
-        "n_embd": 768,
-        "n_layer": 6
-    },
-    "convnext": {
-        "n_embd": 768,
-        "n_layer": 12
-    },
-}
-
-with open(path_dst + '/config.json', 'w') as f:
-    json.dump(config, f, indent=4)
-
-print(f"Config has been saved to {path_dst + 'config.json'}")
diff --git a/tools/tts/tts-outetts.py b/tools/tts/tts-outetts.py
deleted file mode 100644 (file)
index 3791f9f..0000000
+++ /dev/null
@@ -1,299 +0,0 @@
-import sys
-#import json
-#import struct
-import requests
-import re
-import struct
-import numpy as np
-from concurrent.futures import ThreadPoolExecutor
-
-
-def fill_hann_window(size, periodic=True):
-    if periodic:
-        return np.hanning(size + 1)[:-1]
-    return np.hanning(size)
-
-
-def irfft(n_fft, complex_input):
-    return np.fft.irfft(complex_input, n=n_fft)
-
-
-def fold(buffer, n_out, n_win, n_hop, n_pad):
-    result = np.zeros(n_out)
-    n_frames = len(buffer) // n_win
-
-    for i in range(n_frames):
-        start = i * n_hop
-        end = start + n_win
-        result[start:end] += buffer[i * n_win:(i + 1) * n_win]
-
-    return result[n_pad:-n_pad] if n_pad > 0 else result
-
-
-def process_frame(args):
-    l, n_fft, ST, hann = args
-    frame = irfft(n_fft, ST[l])
-    frame = frame * hann
-    hann2 = hann * hann
-    return frame, hann2
-
-
-def embd_to_audio(embd, n_codes, n_embd, n_thread=4):
-    embd = np.asarray(embd, dtype=np.float32).reshape(n_codes, n_embd)
-
-    n_fft = 1280
-    n_hop = 320
-    n_win = 1280
-    n_pad = (n_win - n_hop) // 2
-    n_out = (n_codes - 1) * n_hop + n_win
-
-    hann = fill_hann_window(n_fft, True)
-
-    E = np.zeros((n_embd, n_codes), dtype=np.float32)
-    for l in range(n_codes):
-        for k in range(n_embd):
-            E[k, l] = embd[l, k]
-
-    half_embd = n_embd // 2
-    S = np.zeros((n_codes, half_embd + 1), dtype=np.complex64)
-
-    for k in range(half_embd):
-        for l in range(n_codes):
-            mag = E[k, l]
-            phi = E[k + half_embd, l]
-
-            mag = np.clip(np.exp(mag), 0, 1e2)
-            S[l, k] = mag * np.exp(1j * phi)
-
-    res = np.zeros(n_codes * n_fft)
-    hann2_buffer = np.zeros(n_codes * n_fft)
-
-    with ThreadPoolExecutor(max_workers=n_thread) as executor:
-        args = [(l, n_fft, S, hann) for l in range(n_codes)]
-        results = list(executor.map(process_frame, args))
-
-        for l, (frame, hann2) in enumerate(results):
-            res[l*n_fft:(l+1)*n_fft] = frame
-            hann2_buffer[l*n_fft:(l+1)*n_fft] = hann2
-
-    audio = fold(res, n_out, n_win, n_hop, n_pad)
-    env = fold(hann2_buffer, n_out, n_win, n_hop, n_pad)
-
-    mask = env > 1e-10
-    audio[mask] /= env[mask]
-
-    return audio
-
-
-def save_wav(filename, audio_data, sample_rate):
-    num_channels = 1
-    bits_per_sample = 16
-    bytes_per_sample = bits_per_sample // 8
-    data_size = len(audio_data) * bytes_per_sample
-    byte_rate = sample_rate * num_channels * bytes_per_sample
-    block_align = num_channels * bytes_per_sample
-    chunk_size = 36 + data_size  # 36 = size of header minus first 8 bytes
-
-    header = struct.pack(
-        '<4sI4s4sIHHIIHH4sI',
-        b'RIFF',
-        chunk_size,
-        b'WAVE',
-        b'fmt ',
-        16,                # fmt chunk size
-        1,                 # audio format (PCM)
-        num_channels,
-        sample_rate,
-        byte_rate,
-        block_align,
-        bits_per_sample,
-        b'data',
-        data_size
-    )
-
-    audio_data = np.clip(audio_data * 32767, -32768, 32767)
-    pcm_data = audio_data.astype(np.int16)
-
-    with open(filename, 'wb') as f:
-        f.write(header)
-        f.write(pcm_data.tobytes())
-
-
-def process_text(text: str):
-    text = re.sub(r'\d+(\.\d+)?', lambda x: x.group(), text.lower()) # TODO this needs to be fixed
-    text = re.sub(r'[-_/,\.\\]', ' ', text)
-    text = re.sub(r'[^a-z\s]', '', text)
-    text = re.sub(r'\s+', ' ', text).strip()
-    return text.split()
-
-# usage:
-# python tts-outetts.py http://server-llm:port http://server-dec:port "text"
-
-if len(sys.argv) <= 3:
-    print("usage: python tts-outetts.py http://server-llm:port http://server-dec:port \"text\"")
-    exit(1)
-
-host_llm = sys.argv[1]
-host_dec = sys.argv[2]
-text = sys.argv[3]
-
-prefix = """<|im_start|>
-<|text_start|>the<|text_sep|>overall<|text_sep|>package<|text_sep|>from<|text_sep|>just<|text_sep|>two<|text_sep|>people<|text_sep|>is<|text_sep|>pretty<|text_sep|>remarkable<|text_sep|>sure<|text_sep|>i<|text_sep|>have<|text_sep|>some<|text_sep|>critiques<|text_sep|>about<|text_sep|>some<|text_sep|>of<|text_sep|>the<|text_sep|>gameplay<|text_sep|>aspects<|text_sep|>but<|text_sep|>its<|text_sep|>still<|text_sep|>really<|text_sep|>enjoyable<|text_sep|>and<|text_sep|>it<|text_sep|>looks<|text_sep|>lovely<|text_sep|>"""
-
-words = process_text(text)
-words = "<|text_sep|>".join([i.strip() for i in words])
-words += "<|text_end|>\n"
-
-# voice data
-# TODO: load from json
-#suffix = """<|audio_start|>
-#the<|t_0.08|><|code_start|><|257|><|740|><|636|><|913|><|788|><|1703|><|code_end|>
-#overall<|t_0.36|><|code_start|><|127|><|201|><|191|><|774|><|700|><|532|><|1056|><|557|><|798|><|298|><|1741|><|747|><|1662|><|1617|><|1702|><|1527|><|368|><|1588|><|1049|><|1008|><|1625|><|747|><|1576|><|728|><|1019|><|1696|><|1765|><|code_end|>
-#package<|t_0.56|><|code_start|><|935|><|584|><|1319|><|627|><|1016|><|1491|><|1344|><|1117|><|1526|><|1040|><|239|><|1435|><|951|><|498|><|723|><|1180|><|535|><|789|><|1649|><|1637|><|78|><|465|><|1668|><|901|><|595|><|1675|><|117|><|1009|><|1667|><|320|><|840|><|79|><|507|><|1762|><|1508|><|1228|><|1768|><|802|><|1450|><|1457|><|232|><|639|><|code_end|>
-#from<|t_0.19|><|code_start|><|604|><|782|><|1682|><|872|><|1532|><|1600|><|1036|><|1761|><|647|><|1554|><|1371|><|653|><|1595|><|950|><|code_end|>
-#just<|t_0.25|><|code_start|><|1782|><|1670|><|317|><|786|><|1748|><|631|><|599|><|1155|><|1364|><|1524|><|36|><|1591|><|889|><|1535|><|541|><|440|><|1532|><|50|><|870|><|code_end|>
-#two<|t_0.24|><|code_start|><|1681|><|1510|><|673|><|799|><|805|><|1342|><|330|><|519|><|62|><|640|><|1138|><|565|><|1552|><|1497|><|1552|><|572|><|1715|><|1732|><|code_end|>
-#people<|t_0.39|><|code_start|><|593|><|274|><|136|><|740|><|691|><|633|><|1484|><|1061|><|1138|><|1485|><|344|><|428|><|397|><|1562|><|645|><|917|><|1035|><|1449|><|1669|><|487|><|442|><|1484|><|1329|><|1832|><|1704|><|600|><|761|><|653|><|269|><|code_end|>
-#is<|t_0.16|><|code_start|><|566|><|583|><|1755|><|646|><|1337|><|709|><|802|><|1008|><|485|><|1583|><|652|><|10|><|code_end|>
-#pretty<|t_0.32|><|code_start|><|1818|><|1747|><|692|><|733|><|1010|><|534|><|406|><|1697|><|1053|><|1521|><|1355|><|1274|><|816|><|1398|><|211|><|1218|><|817|><|1472|><|1703|><|686|><|13|><|822|><|445|><|1068|><|code_end|>
-#remarkable<|t_0.68|><|code_start|><|230|><|1048|><|1705|><|355|><|706|><|1149|><|1535|><|1787|><|1356|><|1396|><|835|><|1583|><|486|><|1249|><|286|><|937|><|1076|><|1150|><|614|><|42|><|1058|><|705|><|681|><|798|><|934|><|490|><|514|><|1399|><|572|><|1446|><|1703|><|1346|><|1040|><|1426|><|1304|><|664|><|171|><|1530|><|625|><|64|><|1708|><|1830|><|1030|><|443|><|1509|><|1063|><|1605|><|1785|><|721|><|1440|><|923|><|code_end|>
-#sure<|t_0.36|><|code_start|><|792|><|1780|><|923|><|1640|><|265|><|261|><|1525|><|567|><|1491|><|1250|><|1730|><|362|><|919|><|1766|><|543|><|1|><|333|><|113|><|970|><|252|><|1606|><|133|><|302|><|1810|><|1046|><|1190|><|1675|><|code_end|>
-#i<|t_0.08|><|code_start|><|123|><|439|><|1074|><|705|><|1799|><|637|><|code_end|>
-#have<|t_0.16|><|code_start|><|1509|><|599|><|518|><|1170|><|552|><|1029|><|1267|><|864|><|419|><|143|><|1061|><|0|><|code_end|>
-#some<|t_0.16|><|code_start|><|619|><|400|><|1270|><|62|><|1370|><|1832|><|917|><|1661|><|167|><|269|><|1366|><|1508|><|code_end|>
-#critiques<|t_0.60|><|code_start|><|559|><|584|><|1163|><|1129|><|1313|><|1728|><|721|><|1146|><|1093|><|577|><|928|><|27|><|630|><|1080|><|1346|><|1337|><|320|><|1382|><|1175|><|1682|><|1556|><|990|><|1683|><|860|><|1721|><|110|><|786|><|376|><|1085|><|756|><|1523|><|234|><|1334|><|1506|><|1578|><|659|><|612|><|1108|><|1466|><|1647|><|308|><|1470|><|746|><|556|><|1061|><|code_end|>
-#about<|t_0.29|><|code_start|><|26|><|1649|><|545|><|1367|><|1263|><|1728|><|450|><|859|><|1434|><|497|><|1220|><|1285|><|179|><|755|><|1154|><|779|><|179|><|1229|><|1213|><|922|><|1774|><|1408|><|code_end|>
-#some<|t_0.23|><|code_start|><|986|><|28|><|1649|><|778|><|858|><|1519|><|1|><|18|><|26|><|1042|><|1174|><|1309|><|1499|><|1712|><|1692|><|1516|><|1574|><|code_end|>
-#of<|t_0.07|><|code_start|><|197|><|716|><|1039|><|1662|><|64|><|code_end|>
-#the<|t_0.08|><|code_start|><|1811|><|1568|><|569|><|886|><|1025|><|1374|><|code_end|>
-#gameplay<|t_0.48|><|code_start|><|1269|><|1092|><|933|><|1362|><|1762|><|1700|><|1675|><|215|><|781|><|1086|><|461|><|838|><|1022|><|759|><|649|><|1416|><|1004|><|551|><|909|><|787|><|343|><|830|><|1391|><|1040|><|1622|><|1779|><|1360|><|1231|><|1187|><|1317|><|76|><|997|><|989|><|978|><|737|><|189|><|code_end|>
-#aspects<|t_0.56|><|code_start|><|1423|><|797|><|1316|><|1222|><|147|><|719|><|1347|><|386|><|1390|><|1558|><|154|><|440|><|634|><|592|><|1097|><|1718|><|712|><|763|><|1118|><|1721|><|1311|><|868|><|580|><|362|><|1435|><|868|><|247|><|221|><|886|><|1145|><|1274|><|1284|><|457|><|1043|><|1459|><|1818|><|62|><|599|><|1035|><|62|><|1649|><|778|><|code_end|>
-#but<|t_0.20|><|code_start|><|780|><|1825|><|1681|><|1007|><|861|><|710|><|702|><|939|><|1669|><|1491|><|613|><|1739|><|823|><|1469|><|648|><|code_end|>
-#its<|t_0.09|><|code_start|><|92|><|688|><|1623|><|962|><|1670|><|527|><|599|><|code_end|>
-#still<|t_0.27|><|code_start|><|636|><|10|><|1217|><|344|><|713|><|957|><|823|><|154|><|1649|><|1286|><|508|><|214|><|1760|><|1250|><|456|><|1352|><|1368|><|921|><|615|><|5|><|code_end|>
-#really<|t_0.36|><|code_start|><|55|><|420|><|1008|><|1659|><|27|><|644|><|1266|><|617|><|761|><|1712|><|109|><|1465|><|1587|><|503|><|1541|><|619|><|197|><|1019|><|817|><|269|><|377|><|362|><|1381|><|507|><|1488|><|4|><|1695|><|code_end|>
-#enjoyable<|t_0.49|><|code_start|><|678|><|501|><|864|><|319|><|288|><|1472|><|1341|><|686|><|562|><|1463|><|619|><|1563|><|471|><|911|><|730|><|1811|><|1006|><|520|><|861|><|1274|><|125|><|1431|><|638|><|621|><|153|><|876|><|1770|><|437|><|987|><|1653|><|1109|><|898|><|1285|><|80|><|593|><|1709|><|843|><|code_end|>
-#and<|t_0.15|><|code_start|><|1285|><|987|><|303|><|1037|><|730|><|1164|><|502|><|120|><|1737|><|1655|><|1318|><|code_end|>
-#it<|t_0.09|><|code_start|><|848|><|1366|><|395|><|1601|><|1513|><|593|><|1302|><|code_end|>
-#looks<|t_0.27|><|code_start|><|1281|><|1266|><|1755|><|572|><|248|><|1751|><|1257|><|695|><|1380|><|457|><|659|><|585|><|1315|><|1105|><|1776|><|736|><|24|><|736|><|654|><|1027|><|code_end|>
-#lovely<|t_0.56|><|code_start|><|634|><|596|><|1766|><|1556|><|1306|><|1285|><|1481|><|1721|><|1123|><|438|><|1246|><|1251|><|795|><|659|><|1381|><|1658|><|217|><|1772|><|562|><|952|><|107|><|1129|><|1112|><|467|><|550|><|1079|><|840|><|1615|><|1469|><|1380|><|168|><|917|><|836|><|1827|><|437|><|583|><|67|><|595|><|1087|><|1646|><|1493|><|1677|><|code_end|>"""
-
-# TODO: tokenization is slow for some reason - here is pre-tokenized input
-suffix = [ 151667, 198, 1782, 155780, 151669, 151929, 152412, 152308, 152585, 152460, 153375, 151670, 198, 74455,
-          155808, 151669, 151799, 151873, 151863, 152446, 152372, 152204, 152728, 152229, 152470, 151970, 153413,
-          152419, 153334, 153289, 153374, 153199, 152040, 153260, 152721, 152680, 153297, 152419, 153248, 152400,
-          152691, 153368, 153437, 151670, 198, 1722, 155828, 151669, 152607, 152256, 152991, 152299, 152688, 153163,
-          153016, 152789, 153198, 152712, 151911, 153107, 152623, 152170, 152395, 152852, 152207, 152461, 153321,
-          153309, 151750, 152137, 153340, 152573, 152267, 153347, 151789, 152681, 153339, 151992, 152512, 151751,
-          152179, 153434, 153180, 152900, 153440, 152474, 153122, 153129, 151904, 152311, 151670, 198, 1499, 155791,
-          151669, 152276, 152454, 153354, 152544, 153204, 153272, 152708, 153433, 152319, 153226, 153043, 152325,
-          153267, 152622, 151670, 198, 4250, 155797, 151669, 153454, 153342, 151989, 152458, 153420, 152303, 152271,
-          152827, 153036, 153196, 151708, 153263, 152561, 153207, 152213, 152112, 153204, 151722, 152542, 151670, 198,
-          19789, 155796, 151669, 153353, 153182, 152345, 152471, 152477, 153014, 152002, 152191, 151734, 152312, 152810,
-          152237, 153224, 153169, 153224, 152244, 153387, 153404, 151670, 198, 16069, 155811, 151669, 152265, 151946,
-          151808, 152412, 152363, 152305, 153156, 152733, 152810, 153157, 152016, 152100, 152069, 153234, 152317,
-          152589, 152707, 153121, 153341, 152159, 152114, 153156, 153001, 153504, 153376, 152272, 152433, 152325,
-          151941, 151670, 198, 285, 155788, 151669, 152238, 152255, 153427, 152318, 153009, 152381, 152474, 152680,
-          152157, 153255, 152324, 151682, 151670, 198, 32955, 155804, 151669, 153490, 153419, 152364, 152405, 152682,
-          152206, 152078, 153369, 152725, 153193, 153027, 152946, 152488, 153070, 151883, 152890, 152489, 153144,
-          153375, 152358, 151685, 152494, 152117, 152740, 151670, 198, 37448, 480, 155840, 151669, 151902, 152720,
-          153377, 152027, 152378, 152821, 153207, 153459, 153028, 153068, 152507, 153255, 152158, 152921, 151958,
-          152609, 152748, 152822, 152286, 151714, 152730, 152377, 152353, 152470, 152606, 152162, 152186, 153071,
-          152244, 153118, 153375, 153018, 152712, 153098, 152976, 152336, 151843, 153202, 152297, 151736, 153380,
-          153502, 152702, 152115, 153181, 152735, 153277, 153457, 152393, 153112, 152595, 151670, 198, 19098, 155808,
-          151669, 152464, 153452, 152595, 153312, 151937, 151933, 153197, 152239, 153163, 152922, 153402, 152034,
-          152591, 153438, 152215, 151673, 152005, 151785, 152642, 151924, 153278, 151805, 151974, 153482, 152718,
-          152862, 153347, 151670, 198, 72, 155780, 151669, 151795, 152111, 152746, 152377, 153471, 152309, 151670, 198,
-          19016, 155788, 151669, 153181, 152271, 152190, 152842, 152224, 152701, 152939, 152536, 152091, 151815, 152733,
-          151672, 151670, 198, 14689, 155788, 151669, 152291, 152072, 152942, 151734, 153042, 153504, 152589, 153333,
-          151839, 151941, 153038, 153180, 151670, 198, 36996, 8303, 155832, 151669, 152231, 152256, 152835, 152801,
-          152985, 153400, 152393, 152818, 152765, 152249, 152600, 151699, 152302, 152752, 153018, 153009, 151992,
-          153054, 152847, 153354, 153228, 152662, 153355, 152532, 153393, 151782, 152458, 152048, 152757, 152428,
-          153195, 151906, 153006, 153178, 153250, 152331, 152284, 152780, 153138, 153319, 151980, 153142, 152418,
-          152228, 152733, 151670, 198, 9096, 155801, 151669, 151698, 153321, 152217, 153039, 152935, 153400, 152122,
-          152531, 153106, 152169, 152892, 152957, 151851, 152427, 152826, 152451, 151851, 152901, 152885, 152594,
-          153446, 153080, 151670, 198, 14689, 155795, 151669, 152658, 151700, 153321, 152450, 152530, 153191, 151673,
-          151690, 151698, 152714, 152846, 152981, 153171, 153384, 153364, 153188, 153246, 151670, 198, 1055, 155779,
-          151669, 151869, 152388, 152711, 153334, 151736, 151670, 198, 1782, 155780, 151669, 153483, 153240, 152241,
-          152558, 152697, 153046, 151670, 198, 5804, 1363, 155820, 151669, 152941, 152764, 152605, 153034, 153434,
-          153372, 153347, 151887, 152453, 152758, 152133, 152510, 152694, 152431, 152321, 153088, 152676, 152223,
-          152581, 152459, 152015, 152502, 153063, 152712, 153294, 153451, 153032, 152903, 152859, 152989, 151748,
-          152669, 152661, 152650, 152409, 151861, 151670, 198, 300, 7973, 155828, 151669, 153095, 152469, 152988,
-          152894, 151819, 152391, 153019, 152058, 153062, 153230, 151826, 152112, 152306, 152264, 152769, 153390,
-          152384, 152435, 152790, 153393, 152983, 152540, 152252, 152034, 153107, 152540, 151919, 151893, 152558,
-          152817, 152946, 152956, 152129, 152715, 153131, 153490, 151734, 152271, 152707, 151734, 153321, 152450,
-          151670, 198, 8088, 155792, 151669, 152452, 153497, 153353, 152679, 152533, 152382, 152374, 152611, 153341,
-          153163, 152285, 153411, 152495, 153141, 152320, 151670, 198, 1199, 155781, 151669, 151764, 152360, 153295,
-          152634, 153342, 152199, 152271, 151670, 198, 43366, 155799, 151669, 152308, 151682, 152889, 152016, 152385,
-          152629, 152495, 151826, 153321, 152958, 152180, 151886, 153432, 152922, 152128, 153024, 153040, 152593,
-          152287, 151677, 151670, 198, 53660, 155808, 151669, 151727, 152092, 152680, 153331, 151699, 152316, 152938,
-          152289, 152433, 153384, 151781, 153137, 153259, 152175, 153213, 152291, 151869, 152691, 152489, 151941,
-          152049, 152034, 153053, 152179, 153160, 151676, 153367, 151670, 198, 268, 4123, 480, 155821, 151669, 152350,
-          152173, 152536, 151991, 151960, 153144, 153013, 152358, 152234, 153135, 152291, 153235, 152143, 152583,
-          152402, 153483, 152678, 152192, 152533, 152946, 151797, 153103, 152310, 152293, 151825, 152548, 153442,
-          152109, 152659, 153325, 152781, 152570, 152957, 151752, 152265, 153381, 152515, 151670, 198, 437, 155787,
-          151669, 152957, 152659, 151975, 152709, 152402, 152836, 152174, 151792, 153409, 153327, 152990, 151670, 198,
-          275, 155781, 151669, 152520, 153038, 152067, 153273, 153185, 152265, 152974, 151670, 198, 94273, 155799,
-          151669, 152953, 152938, 153427, 152244, 151920, 153423, 152929, 152367, 153052, 152129, 152331, 152257,
-          152987, 152777, 153448, 152408, 151696, 152408, 152326, 152699, 151670, 198, 385, 16239, 155828, 151669,
-          152306, 152268, 153438, 153228, 152978, 152957, 153153, 153393, 152795, 152110, 152918, 152923, 152467,
-          152331, 153053, 153330, 151889, 153444, 152234, 152624, 151779, 152801, 152784, 152139, 152222, 152751,
-          152512, 153287, 153141, 153052, 151840, 152589, 152508, 153499, 152109, 152255, 151739, 152267, 152759,
-          153318, 153165, 153349, 151670, ]
-
-response = requests.post(
-    host_llm + "/completion",
-    json={
-        "prompt": [prefix + words, *suffix],
-        "n_predict": 1024,
-        "cache_prompt": True,
-        "return_tokens": True,
-        "samplers": ["top_k"],
-        "top_k": 16,
-        "seed": 1003,
-    }
-)
-
-response_json = response.json()
-
-#print(json.dumps(response_json, indent=4))
-#print(json.dumps(response_json["prompt"], indent=4).replace("\\n", "\n"))
-#print(json.dumps(response_json["timings"], indent=4))
-#print(json.dumps(response_json["tokens"], indent=4))
-
-codes = response_json["tokens"]
-
-codes = [t - 151672 for t in codes if t >= 151672 and t <= 155772]
-
-response = requests.post(
-    host_dec + "/embeddings",
-    json={
-        "input": [*codes],
-    }
-)
-
-response_json = response.json()
-
-#print(json.dumps(response_json, indent=4))
-
-# spectrogram
-embd = response_json[0]["embedding"]
-
-n_codes = len(embd)
-n_embd = len(embd[0])
-
-print('spectrogram generated: n_codes: %d, n_embd: %d' % (n_codes, n_embd))
-
-# post-process the spectrogram to convert to audio
-print('converting to audio ...')
-audio = embd_to_audio(embd, n_codes, n_embd)
-print('audio generated: %d samples' % len(audio))
-
-filename = "output.wav"
-sample_rate = 24000 # sampling rate
-
-# zero out first 0.25 seconds
-audio[:24000 // 4] = 0.0
-
-save_wav(filename, audio, sample_rate)
-print('audio written to file "%s"' % filename)
index 2a1bdccc9192bb8aa6c74db9577db8a0045feb93..b68edcaf5758136d882d5e5e94d2a938b4487dcb 100644 (file)
-#define _USE_MATH_DEFINES // For M_PI on MSVC
-
 #include "arg.h"
 #include "common.h"
 #include "sampling.h"
 #include "log.h"
 #include "llama.h"
+#include "mtmd.h"
+#include "mtmd-helper.h"
 
-#define JSON_ASSERT GGML_ASSERT
-#include <nlohmann/json.hpp>
-
-#include <algorithm>
-#include <clocale>
-#include <cmath>
 #include <cstdio>
-#include <fstream>
-#include <map>
-#include <regex>
+#include <cstring>
 #include <string>
-#include <thread>
-#include <vector>
-
-using json = nlohmann::ordered_json;
-
-enum outetts_version {
-    OUTETTS_V0_2,
-    OUTETTS_V0_3,
-};
-
-//
-// Terminal utils
-//
-
-#define SQR(X)    ((X) * (X))
-#define UNCUBE(x) x < 48 ? 0 : x < 115 ? 1 : (x - 35) / 40
 
 /**
- * Quantizes 24-bit RGB to xterm256 code range [16,256).
+ * Please note that this is NOT a production-ready binary.
+ * It is a playground for trying TTS support in llama.cpp.
+ * For contributors: please keep this code simple and easy to understand. Do not add unnecessary complexity. The goal is to have a simple CLI for testing TTS support.
  */
-static int rgb2xterm256(int r, int g, int b) {
-    unsigned char cube[] = {0, 0137, 0207, 0257, 0327, 0377};
-    int av, ir, ig, ib, il, qr, qg, qb, ql;
-    av = r * .299 + g * .587 + b * .114 + .5;
-    ql = (il = av > 238 ? 23 : (av - 3) / 10) * 10 + 8;
-    qr = cube[(ir = UNCUBE(r))];
-    qg = cube[(ig = UNCUBE(g))];
-    qb = cube[(ib = UNCUBE(b))];
-    if (SQR(qr - r) + SQR(qg - g) + SQR(qb - b) <=
-        SQR(ql - r) + SQR(ql - g) + SQR(ql - b))
-        return ir * 36 + ig * 6 + ib + 020;
-    return il + 0350;
-}
-
-static std::string set_xterm256_foreground(int r, int g, int b) {
-    int x = rgb2xterm256(r, g, b);
-    std::ostringstream oss;
-    oss << "\033[38;5;" << x << "m";
-    return oss.str();
-}
-
-const std::vector<std::string> k_colors = {
-    set_xterm256_foreground(220,   5,  12),
-    set_xterm256_foreground(232,  96,  28),
-    set_xterm256_foreground(241, 147,  45),
-    set_xterm256_foreground(246, 193,  65),
-    set_xterm256_foreground(247, 240,  86),
-    set_xterm256_foreground(144, 201, 135),
-    set_xterm256_foreground( 78, 178, 101),
-};
-
-static void print_usage(int, char ** argv) {
-    LOG("\nexample usage:\n");
-    LOG("\n    %s -m model.gguf -p \"Hello!\"\n", argv[0]);
-    LOG("\n");
-}
-
-struct wav_header {
-    char riff[4] = {'R', 'I', 'F', 'F'};
-    uint32_t chunk_size;
-    char wave[4] = {'W', 'A', 'V', 'E'};
-    char fmt[4] = {'f', 'm', 't', ' '};
-    uint32_t fmt_chunk_size = 16;
-    uint16_t audio_format = 1; // PCM
-    uint16_t num_channels = 1; // Mono
-    uint32_t sample_rate;
-    uint32_t byte_rate;
-    uint16_t block_align;
-    uint16_t bits_per_sample = 16;
-    char data[4] = {'d', 'a', 't', 'a'};
-    uint32_t data_size;
-};
-
-static bool save_wav16(const std::string & fname, const std::vector<float> & data, int sample_rate) {
-    std::ofstream file(fname, std::ios::binary);
-    if (!file) {
-        LOG_ERR("%s: Failed to open file '%s' for writing.\n", __func__, fname.c_str());
-        return false;
-    }
-
-    wav_header header;
-    header.sample_rate = sample_rate;
-    header.byte_rate = header.sample_rate * header.num_channels * (header.bits_per_sample / 8);
-    header.block_align = header.num_channels * (header.bits_per_sample / 8);
-    header.data_size = data.size() * (header.bits_per_sample / 8);
-    header.chunk_size = 36 + header.data_size;
-
-    file.write(reinterpret_cast<const char*>(&header), sizeof(header));
-
-    for (const auto & sample : data) {
-        int16_t pcm_sample = static_cast<int16_t>(std::clamp(sample * 32767.0, -32768.0, 32767.0));
-        file.write(reinterpret_cast<const char*>(&pcm_sample), sizeof(pcm_sample));
-    }
-
-    return file.good();
-}
 
-static void fill_hann_window(int length, bool periodic, float * output) {
-    int offset = -1;
-    if (periodic) {
-        offset = 0;
-    }
-    for (int i = 0; i < length; i++) {
-        output[i] = 0.5 * (1.0 - cosf((2.0 * M_PI * i) / (length + offset)));
-    }
-}
-
-// very poor-man fft
-static void twiddle(float * real, float * imag, int k, int N) {
-    float angle = 2 * M_PI * k / N;
-    *real = cos(angle);
-    *imag = sin(angle);
-}
-
-static void irfft(int n, const float * inp_cplx, float * out_real) {
-    int N = n / 2 + 1;
-
-    std::vector<float> real_input(N);
-    std::vector<float> imag_input(N);
-    for (int i = 0; i < N; ++i) {
-        real_input[i] = inp_cplx[2 * i];
-        imag_input[i] = inp_cplx[2 * i + 1];
-    }
+struct tts_timings {
+    int64_t t_start_us = ggml_time_us();
+    int64_t t_last_us  = t_start_us;
 
-    std::vector<float> real_output(n);
-    std::vector<float> imag_output(n);
-
-    for (int k = 0; k < n; ++k) {
-        real_output[k] = 0.0f;
-        imag_output[k] = 0.0f;
-        for (int m = 0; m < N; ++m) {
-            float twiddle_real;
-            float twiddle_imag;
-
-            twiddle(&twiddle_real, &twiddle_imag, k * m, n);
-
-            real_output[k] += real_input[m] * twiddle_real - imag_input[m] * twiddle_imag;
-            imag_output[k] += real_input[m] * twiddle_imag + imag_input[m] * twiddle_real;
+    void report(int n_frames) {
+        const int64_t t_now_us = ggml_time_us();
+        if (t_now_us - t_last_us < 2000000) {
+            return;
         }
+        t_last_us = t_now_us;
+        const double t_elapsed_s = (t_now_us - t_start_us) / 1e6;
+        const double fps = t_elapsed_s > 0 ? n_frames / t_elapsed_s : 0.0;
+        LOG_INF("frames generated: %d, speed: %.2f frames/s\n", n_frames, fps);
     }
-
-    for (int i = 0; i < n; ++i) {
-        out_real[i] = real_output[i] / N;
-    }
-}
-
-//
-//  y = torch.nn.functional.fold(
-//       data, output_size=(1, output_size), kernel_size=(1, self.win_length), stride=(1, self.hop_length),
-//  )[:, 0, 0, pad:-pad]
-//
-// data.shape =  torch.Size([1, 1280, 261])
-// output_size =  84480
-// win_length =  1280
-// hop_length =  320
-// pad =  480
-//
-static void fold(const std::vector<float> & data, int64_t n_out, int64_t n_win, int64_t n_hop, int64_t n_pad, std::vector<float> & output) {
-    int64_t output_height = n_out;
-    int64_t kernel_w = n_win;
-    int64_t stride_w = n_hop;
-    int64_t width    = n_out;
-
-    output.resize(width, 0.0f);
-
-    int64_t col_idx = 0;
-    for (int64_t w_col = 0; w_col < width; ++w_col) {
-        int64_t start = w_col * stride_w - n_pad;
-        int64_t end   = start + kernel_w;
-
-        for (int64_t w_im = start; w_im < end; ++w_im) {
-            if (w_im >= 0 && w_im < output_height && col_idx < (int64_t) data.size()) {
-                output[w_im] += data[col_idx];
-            }
-            col_idx++;
-        }
-    }
-
-    output.resize(n_out - 2 * n_pad);
-}
-
-// TODO: not optimized at all
-static std::vector<float> embd_to_audio(
-        const float * embd,
-        const int n_codes,
-        const int n_embd,
-        const int n_thread) {
-    const int n_fft = 1280;
-    const int n_hop = 320;
-    const int n_win = 1280;
-    const int n_pad = (n_win - n_hop)/2;
-    const int n_out = (n_codes - 1)*n_hop + n_win;
-
-    std::vector<float> hann(n_fft);
-
-    fill_hann_window(hann.size(), true, hann.data());
-
-    int n_spec = n_embd*n_codes;
-
-    std::vector<float> E (n_spec);
-    std::vector<float> S (n_spec);
-    std::vector<float> ST(n_spec);
-
-    for (int l = 0; l < n_codes; ++l) {
-        for (int k = 0; k < n_embd; ++k) {
-            E[k*n_codes + l] = embd[l*n_embd + k];
-        }
-    }
-
-    for (int k = 0; k < n_embd/2; ++k) {
-        for (int l = 0; l < n_codes; ++l) {
-            float mag = E[(k           )*n_codes + l];
-            float phi = E[(k + n_embd/2)*n_codes + l];
-
-            mag = exp(mag);
-
-            if (mag > 1e2) {
-                mag = 1e2;
-            }
-            S[2*(k*n_codes + l) + 0] = mag*cosf(phi);
-            S[2*(k*n_codes + l) + 1] = mag*sinf(phi);
-        }
-    }
-
-    for (int l = 0; l < n_codes; ++l) {
-        for (int k = 0; k < n_embd/2; ++k) {
-            ST[l*n_embd + 2*k + 0] = S[2*(k*n_codes + l) + 0];
-            ST[l*n_embd + 2*k + 1] = S[2*(k*n_codes + l) + 1];
-        }
-    }
-
-    std::vector<float> res  (n_codes*n_fft);
-    std::vector<float> hann2(n_codes*n_fft);
-
-    std::vector<std::thread> workers(n_thread);
-    for (int i = 0; i < n_thread; ++i) {
-        workers[i] = std::thread([&, i]() {
-            for (int l = i; l < n_codes; l += n_thread) {
-                irfft(n_fft, ST.data() + l*n_embd, res.data() + l*n_fft);
-                for (int j = 0; j < n_fft; ++j) {
-                    res  [l*n_fft + j] *= hann[j];
-                    hann2[l*n_fft + j]  = hann[j] * hann[j];
-                }
-            }
-        });
-    }
-    for (int i = 0; i < n_thread; ++i) {
-        workers[i].join();
-    }
-
-    std::vector<float> audio;
-    std::vector<float> env;
-
-    fold(res,   n_out, n_win, n_hop, n_pad, audio);
-    fold(hann2, n_out, n_win, n_hop, n_pad, env); // TODO: can be done once
-
-    for (size_t i = 0; i < audio.size(); ++i) {
-        audio[i] /= env[i];
-    }
-
-    return audio;
-}
-
-static const std::map<int, std::string> ones = {
-    {0, "zero"}, {1, "one"}, {2, "two"}, {3, "three"}, {4, "four"},
-    {5, "five"}, {6, "six"}, {7, "seven"}, {8, "eight"}, {9, "nine"},
-    {10, "ten"}, {11, "eleven"}, {12, "twelve"}, {13, "thirteen"}, {14, "fourteen"},
-    {15, "fifteen"}, {16, "sixteen"}, {17, "seventeen"}, {18, "eighteen"}, {19, "nineteen"}
 };
 
-static const std::map<int, std::string> tens = {
-    {2, "twenty"}, {3, "thirty"}, {4, "forty"}, {5, "fifty"},
-    {6, "sixty"}, {7, "seventy"}, {8, "eighty"}, {9, "ninety"}
-};
-
-// Convert a number less than 1000 to words
-static std::string convert_less_than_thousand(int num) {
-    std::string result;
-
-    if (num >= 100) {
-        result += ones.at(num / 100) + " hundred ";
-        num %= 100;
-    }
-
-    if (num >= 20) {
-        result += tens.at(num / 10);
-        if (num % 10 > 0) {
-            result += "-" + ones.at(num % 10);
-        }
-    } else if (num > 0) {
-        result += ones.at(num);
-    }
-
-    return result;
-}
-
-static std::string number_to_words(const std::string & number_str) {
-    try {
-        size_t decimal_pos = number_str.find('.');
-        std::string integer_part = number_str.substr(0, decimal_pos);
-
-        int int_number = std::stoi(integer_part);
-        std::string result;
-
-        if (int_number == 0) {
-            result = "zero";
-        } else {
-            if (int_number >= 1000000000) {
-                int billions = int_number / 1000000000;
-                result += convert_less_than_thousand(billions) + " billion ";
-                int_number %= 1000000000;
-            }
-
-            if (int_number >= 1000000) {
-                int millions = int_number / 1000000;
-                result += convert_less_than_thousand(millions) + " million ";
-                int_number %= 1000000;
-            }
-
-            if (int_number >= 1000) {
-                int thousands = int_number / 1000;
-                result += convert_less_than_thousand(thousands) + " thousand ";
-                int_number %= 1000;
-            }
-
-            if (int_number > 0) {
-                result += convert_less_than_thousand(int_number);
-            }
-        }
-
-        // Handle decimal part
-        if (decimal_pos != std::string::npos) {
-            result += " point";
-            std::string decimal_part = number_str.substr(decimal_pos + 1);
-            for (char digit : decimal_part) {
-                result += " " + ones.at(digit - '0');
-            }
-        }
-
-        return result;
-    } catch (const std::exception& e) {
-        // Skip if fails
-        return " ";
-    }
-}
-
-static std::string replace_numbers_with_words(const std::string & input_text) {
-    std::regex number_pattern(R"(\d+(\.\d+)?)");
-    std::string result;
-    auto it = std::sregex_iterator(input_text.begin(), input_text.end(), number_pattern);
-    auto end = std::sregex_iterator();
-
-    size_t last_pos = 0;
-    for (std::sregex_iterator i = it; i != end; ++i) {
-        const std::smatch& match = *i;
-        result.append(input_text, last_pos, match.position() - last_pos);
-        result.append(number_to_words(match.str()));
-        last_pos = match.position() + match.length();
-    }
-    result.append(input_text, last_pos);
-
-    return result;
-}
-
-// Based on: https://github.com/edwko/OuteTTS/blob/a613e79c489d8256dd657ea9168d78de75895d82/outetts/version/v1/prompt_processor.py#L39
-static std::string process_text(const std::string & text, const outetts_version tts_version = OUTETTS_V0_2) {
-
-    // For now I skipped text romanization as I am unsure how to handle
-    // uroman and MeCab implementations in C++
-    // maybe something like https://github.com/anyascii/anyascii/ could work.
-    // currently only English would be supported in this function
-
-    std::string processed_text = replace_numbers_with_words(text);
-
-    std::transform(processed_text.begin(), processed_text.end(),
-                  processed_text.begin(), ::tolower);
-
-    std::regex special_chars(R"([-_/,\.\\])");
-    processed_text = std::regex_replace(processed_text, special_chars, " ");
-
-    std::regex non_alpha(R"([^a-z\s])");
-    processed_text = std::regex_replace(processed_text, non_alpha, "");
-
-    std::regex multiple_spaces(R"(\s+)");
-    processed_text = std::regex_replace(processed_text, multiple_spaces, " ");
-
-    processed_text = std::regex_replace(processed_text, std::regex(R"(^\s+|\s+$)"), "");
-
-    /*
-        Replace spaces with the separator token same as in line 365
-
-        for (auto & c : prompt_user) {
-        if (c == ' ') {
-            prompt_clean += "<|text_sep|>";
-    */
-    std::string separator = (tts_version == OUTETTS_V0_3) ? "<|space|>" : "<|text_sep|>";
-    processed_text = std::regex_replace(processed_text, std::regex(R"(\s)"), separator);
-
-    return processed_text;
-}
-
-static void prompt_add(llama_tokens & prompt, llama_token token) {
-    prompt.push_back(token);
-}
-
-static void prompt_add(llama_tokens & prompt, const llama_tokens & tokens) {
-    prompt.insert(prompt.end(), tokens.begin(), tokens.end());
-}
-
-static void prompt_add(llama_tokens & prompt, const llama_vocab * vocab, const std::string & txt, bool add_special, bool parse_special) {
-    auto tmp = common_tokenize(vocab, txt, add_special, parse_special);
-    prompt_add(prompt, tmp);
-}
-
-static void prompt_init(llama_tokens & prompt, const llama_vocab * vocab) {
-    prompt.clear();
-
-    prompt_add(prompt, vocab, "<|im_start|>\n", true, true);
-}
-
-static std::vector<llama_token> prepare_guide_tokens(const llama_vocab * vocab, const std::string & str, const outetts_version tts_version = OUTETTS_V0_2) {
-    const std::string& delimiter = (tts_version == OUTETTS_V0_3 ? "<|space|>" : "<|text_sep|>");
-
-    std::vector<llama_token> result;
-    size_t start = 0;
-    size_t end = str.find(delimiter);
-
-    //first token is always a newline, as it was not previously added
-    result.push_back(common_tokenize(vocab, "\n", false, true)[0]);
-
-    while (end != std::string::npos) {
-        std::string current_word = str.substr(start, end - start);
-        auto tmp = common_tokenize(vocab, current_word, false, true);
-        result.push_back(tmp[0]);
-        start = end + delimiter.length();
-        end = str.find(delimiter, start);
-    }
-
-    // Add the last part
-    std::string current_word = str.substr(start);
-    auto tmp = common_tokenize(vocab, current_word, false, true);
-    if (tmp.size() > 0) {
-        result.push_back(tmp[0]);
-    }
-    return result;
-}
-
-static json speaker_from_file(const std::string & speaker_file) {
-    std::ifstream file(speaker_file);
-    if (!file) {
-        LOG_ERR("%s: Failed to open file '%s' for reading\n", __func__, speaker_file.c_str());
-        return json();
-    }
-
-    json speaker = json::parse(file);
-    return speaker;
-}
-
-static outetts_version get_tts_version(llama_model *model, json speaker = json::object()) {
-    if (speaker.contains("version")) {
-        std::string version = speaker["version"].get<std::string>();
-        if (version == "0.2") {
-            return OUTETTS_V0_2;
-        } else if (version == "0.3") {
-            return OUTETTS_V0_3;
-        } else {
-            LOG_ERR("%s: Unsupported speaker version '%s'\n", __func__, version.c_str());
-        }
-    }
-
-    // Also could get version from model itself
-    const char *chat_template = llama_model_chat_template(model, nullptr);
-    if (chat_template && std::string(chat_template) == "outetts-0.3") {
-        return OUTETTS_V0_3;
-    }
-
-    // Use 0.2 as the default version
-    return OUTETTS_V0_2;
-}
-
-static std::string audio_text_from_speaker(json speaker, const outetts_version tts_version = OUTETTS_V0_2) {
-    std::string audio_text = "<|text_start|>";
-
-    if (tts_version == OUTETTS_V0_2 || tts_version == OUTETTS_V0_3) {
-        std::string separator = (tts_version == OUTETTS_V0_3) ? "<|space|>" : "<|text_sep|>";
-        for (const auto &word : speaker["words"]) {
-            audio_text += word["word"].get<std::string>() + separator;
-        }
-    }
-
-    return audio_text;
-}
-
-static std::string audio_data_from_speaker(json speaker, const outetts_version tts_version = OUTETTS_V0_2) {
-    std::string audio_data = "<|audio_start|>\n";
-
-    if (tts_version == OUTETTS_V0_2 || tts_version == OUTETTS_V0_3) {
-        std::string code_start = (tts_version == OUTETTS_V0_3) ? "" : "<|code_start|>";
-        std::string code_end = (tts_version == OUTETTS_V0_3) ? "<|space|>" : "<|code_end|>";
-        for (const auto &word : speaker["words"]) {
-            std::string word_text = word["word"].get<std::string>();
-            double duration = word["duration"].get<double>();
-            std::vector<int> codes = word["codes"].get<std::vector<int>>();
-
-            // Create the audio output entry
-            std::ostringstream word_entry;
-            word_entry << word_text << "<|t_" << std::fixed << std::setprecision(2)
-                       << duration << "|>" + code_start;
-            for (const auto &Code : codes) {
-                word_entry << "<|" << Code << "|>";
-            }
-            word_entry << code_end << "\n";
-            audio_data += word_entry.str();
-        }
-    }
-
-    return audio_data;
+static void print_usage(int, char ** argv) {
+    LOG("\nexample usage:\n");
+    LOG("\n    %s -m backbone.gguf -mm mmproj.gguf -p \"text to speak\" -o output.wav", argv[0]);
+    LOG("\n    %s -hf user/model -p \"text to speak\" -o output.wav\n", argv[0]);
+    LOG("\nnote: --tts-lang and --tts-speaker-file may not be supported in all models");
+    LOG("\n      use -n to limit the output length");
+    LOG("\n      see tts/README.md for per-model usage notes");
+    LOG("\n\n");
 }
 
 int main(int argc, char ** argv) {
-    std::setlocale(LC_NUMERIC, "C");
-
     common_params params;
 
-    params.out_file = "output.wav";
-    params.prompt = "";
-
-    params.n_predict = 4096;
-    params.n_batch   = 8192;
-    params.n_ctx     = 8192;
-
-    params.sampling.top_k = 4;
-    params.sampling.samplers = { COMMON_SAMPLER_TYPE_TOP_K, };
-
     common_init();
 
     if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_TTS, print_usage)) {
         return 1;
     }
 
-    const int n_parallel = params.n_parallel;
-    const int n_predict  = params.n_predict;
-
-    // init LLM
-
-    llama_backend_init();
-    llama_numa_init(params.numa);
-
-    llama_model * model_ttc = NULL; // text-to-codes
-    llama_model * model_cts = NULL; // codes-to-speech
-
-    llama_context * ctx_ttc = NULL;
-    llama_context * ctx_cts = NULL;
-
-    auto llama_init_ttc = common_init_from_params(params);
-
-    model_ttc = llama_init_ttc->model();
-    ctx_ttc   = llama_init_ttc->context();
+    mtmd_helper_log_set(common_log_default_callback, nullptr);
 
-    if (model_ttc == nullptr || ctx_ttc == nullptr) {
-        return ENOENT;
+    if (params.prompt.empty()) {
+        LOG_ERR("no prompt provided, use -p \"text\"\n");
+        return 1;
+    }
+    if (params.mmproj.path.empty()) {
+        LOG_ERR("no mmproj provided, use --mmproj\n");
+        return 1;
     }
 
-    const llama_vocab * vocab = llama_model_get_vocab(model_ttc);
+    // important: keep this file as generic as possible
+    //            model-specific logic should be in mtmd-helper-gen or mtmd API
 
-    params.model = params.vocoder.model;
+    // always enable embd, so that we can pass hidden states to the audio generation helper
     params.embedding = true;
-    params.n_ubatch = params.n_batch;
 
-    auto llama_init_cts = common_init_from_params(params);
+    llama_backend_init();
+    llama_numa_init(params.numa);
 
-    model_cts = llama_init_cts->model();
-    ctx_cts   = llama_init_cts->context();
+    //
+    // load backbone model and mmproj
+    //
 
-    if (model_cts == nullptr || ctx_cts == nullptr) {
-        return ENOENT;
+    auto llama_init = common_init_from_params(params);
+    llama_model    * model = llama_init->model();
+    llama_context  * lctx  = llama_init->context();
+    common_sampler * smpl  = llama_init->sampler(0);
+    if (!model || !lctx) {
+        LOG_ERR("failed to init model/context\n");
+        return 1;
     }
 
-    std::vector<common_sampler *> smpl(n_parallel);
-    for (int i = 0; i < n_parallel; ++i) {
-        params.sampling.no_perf = (i != 0);
-        params.sampling.seed = params.sampling.seed + 1;
-
-        smpl[i] = common_sampler_init(model_ttc, params.sampling);
+    mtmd_context_params mtmd_params = mtmd_context_params_default();
+    mtmd_params.use_gpu = params.mmproj_use_gpu;
+    mtmd::context_ptr mctx(mtmd_init_from_file(params.mmproj.path.c_str(), model, mtmd_params));
+    if (!mctx) {
+        LOG_ERR("failed to load mmproj %s\n", params.mmproj.path.c_str());
+        return 1;
     }
-
-    LOG_INF("sampler seed: %u\n",     common_sampler_get_seed(smpl[0]));
-    LOG_INF("sampler chain: %s\n",    common_sampler_print(smpl[0]).c_str());
-    LOG_INF("sampler params: \n%s\n", params.sampling.print().c_str());
-
-    LOG_INF("%s: loading done\n", __func__);
-
-    const auto t_main_start = ggml_time_us();
-
-    std::vector<llama_token> codes;
-    std::vector<llama_token> guide_tokens;
-
-    // the default speaker profile is from: https://github.com/edwko/OuteTTS/blob/main/outetts/version/v1/default_speakers/en_male_1.json
-    std::string audio_text = "<|text_start|>the<|text_sep|>overall<|text_sep|>package<|text_sep|>from<|text_sep|>just<|text_sep|>two<|text_sep|>people<|text_sep|>is<|text_sep|>pretty<|text_sep|>remarkable<|text_sep|>sure<|text_sep|>i<|text_sep|>have<|text_sep|>some<|text_sep|>critiques<|text_sep|>about<|text_sep|>some<|text_sep|>of<|text_sep|>the<|text_sep|>gameplay<|text_sep|>aspects<|text_sep|>but<|text_sep|>its<|text_sep|>still<|text_sep|>really<|text_sep|>enjoyable<|text_sep|>and<|text_sep|>it<|text_sep|>looks<|text_sep|>lovely<|text_sep|>";
-    std::string audio_data = R"(<|audio_start|>
-the<|t_0.08|><|code_start|><|257|><|740|><|636|><|913|><|788|><|1703|><|code_end|>
-overall<|t_0.36|><|code_start|><|127|><|201|><|191|><|774|><|700|><|532|><|1056|><|557|><|798|><|298|><|1741|><|747|><|1662|><|1617|><|1702|><|1527|><|368|><|1588|><|1049|><|1008|><|1625|><|747|><|1576|><|728|><|1019|><|1696|><|1765|><|code_end|>
-package<|t_0.56|><|code_start|><|935|><|584|><|1319|><|627|><|1016|><|1491|><|1344|><|1117|><|1526|><|1040|><|239|><|1435|><|951|><|498|><|723|><|1180|><|535|><|789|><|1649|><|1637|><|78|><|465|><|1668|><|901|><|595|><|1675|><|117|><|1009|><|1667|><|320|><|840|><|79|><|507|><|1762|><|1508|><|1228|><|1768|><|802|><|1450|><|1457|><|232|><|639|><|code_end|>
-from<|t_0.19|><|code_start|><|604|><|782|><|1682|><|872|><|1532|><|1600|><|1036|><|1761|><|647|><|1554|><|1371|><|653|><|1595|><|950|><|code_end|>
-just<|t_0.25|><|code_start|><|1782|><|1670|><|317|><|786|><|1748|><|631|><|599|><|1155|><|1364|><|1524|><|36|><|1591|><|889|><|1535|><|541|><|440|><|1532|><|50|><|870|><|code_end|>
-two<|t_0.24|><|code_start|><|1681|><|1510|><|673|><|799|><|805|><|1342|><|330|><|519|><|62|><|640|><|1138|><|565|><|1552|><|1497|><|1552|><|572|><|1715|><|1732|><|code_end|>
-people<|t_0.39|><|code_start|><|593|><|274|><|136|><|740|><|691|><|633|><|1484|><|1061|><|1138|><|1485|><|344|><|428|><|397|><|1562|><|645|><|917|><|1035|><|1449|><|1669|><|487|><|442|><|1484|><|1329|><|1832|><|1704|><|600|><|761|><|653|><|269|><|code_end|>
-is<|t_0.16|><|code_start|><|566|><|583|><|1755|><|646|><|1337|><|709|><|802|><|1008|><|485|><|1583|><|652|><|10|><|code_end|>
-pretty<|t_0.32|><|code_start|><|1818|><|1747|><|692|><|733|><|1010|><|534|><|406|><|1697|><|1053|><|1521|><|1355|><|1274|><|816|><|1398|><|211|><|1218|><|817|><|1472|><|1703|><|686|><|13|><|822|><|445|><|1068|><|code_end|>
-remarkable<|t_0.68|><|code_start|><|230|><|1048|><|1705|><|355|><|706|><|1149|><|1535|><|1787|><|1356|><|1396|><|835|><|1583|><|486|><|1249|><|286|><|937|><|1076|><|1150|><|614|><|42|><|1058|><|705|><|681|><|798|><|934|><|490|><|514|><|1399|><|572|><|1446|><|1703|><|1346|><|1040|><|1426|><|1304|><|664|><|171|><|1530|><|625|><|64|><|1708|><|1830|><|1030|><|443|><|1509|><|1063|><|1605|><|1785|><|721|><|1440|><|923|><|code_end|>
-sure<|t_0.36|><|code_start|><|792|><|1780|><|923|><|1640|><|265|><|261|><|1525|><|567|><|1491|><|1250|><|1730|><|362|><|919|><|1766|><|543|><|1|><|333|><|113|><|970|><|252|><|1606|><|133|><|302|><|1810|><|1046|><|1190|><|1675|><|code_end|>
-i<|t_0.08|><|code_start|><|123|><|439|><|1074|><|705|><|1799|><|637|><|code_end|>
-have<|t_0.16|><|code_start|><|1509|><|599|><|518|><|1170|><|552|><|1029|><|1267|><|864|><|419|><|143|><|1061|><|0|><|code_end|>
-some<|t_0.16|><|code_start|><|619|><|400|><|1270|><|62|><|1370|><|1832|><|917|><|1661|><|167|><|269|><|1366|><|1508|><|code_end|>
-critiques<|t_0.60|><|code_start|><|559|><|584|><|1163|><|1129|><|1313|><|1728|><|721|><|1146|><|1093|><|577|><|928|><|27|><|630|><|1080|><|1346|><|1337|><|320|><|1382|><|1175|><|1682|><|1556|><|990|><|1683|><|860|><|1721|><|110|><|786|><|376|><|1085|><|756|><|1523|><|234|><|1334|><|1506|><|1578|><|659|><|612|><|1108|><|1466|><|1647|><|308|><|1470|><|746|><|556|><|1061|><|code_end|>
-about<|t_0.29|><|code_start|><|26|><|1649|><|545|><|1367|><|1263|><|1728|><|450|><|859|><|1434|><|497|><|1220|><|1285|><|179|><|755|><|1154|><|779|><|179|><|1229|><|1213|><|922|><|1774|><|1408|><|code_end|>
-some<|t_0.23|><|code_start|><|986|><|28|><|1649|><|778|><|858|><|1519|><|1|><|18|><|26|><|1042|><|1174|><|1309|><|1499|><|1712|><|1692|><|1516|><|1574|><|code_end|>
-of<|t_0.07|><|code_start|><|197|><|716|><|1039|><|1662|><|64|><|code_end|>
-the<|t_0.08|><|code_start|><|1811|><|1568|><|569|><|886|><|1025|><|1374|><|code_end|>
-gameplay<|t_0.48|><|code_start|><|1269|><|1092|><|933|><|1362|><|1762|><|1700|><|1675|><|215|><|781|><|1086|><|461|><|838|><|1022|><|759|><|649|><|1416|><|1004|><|551|><|909|><|787|><|343|><|830|><|1391|><|1040|><|1622|><|1779|><|1360|><|1231|><|1187|><|1317|><|76|><|997|><|989|><|978|><|737|><|189|><|code_end|>
-aspects<|t_0.56|><|code_start|><|1423|><|797|><|1316|><|1222|><|147|><|719|><|1347|><|386|><|1390|><|1558|><|154|><|440|><|634|><|592|><|1097|><|1718|><|712|><|763|><|1118|><|1721|><|1311|><|868|><|580|><|362|><|1435|><|868|><|247|><|221|><|886|><|1145|><|1274|><|1284|><|457|><|1043|><|1459|><|1818|><|62|><|599|><|1035|><|62|><|1649|><|778|><|code_end|>
-but<|t_0.20|><|code_start|><|780|><|1825|><|1681|><|1007|><|861|><|710|><|702|><|939|><|1669|><|1491|><|613|><|1739|><|823|><|1469|><|648|><|code_end|>
-its<|t_0.09|><|code_start|><|92|><|688|><|1623|><|962|><|1670|><|527|><|599|><|code_end|>
-still<|t_0.27|><|code_start|><|636|><|10|><|1217|><|344|><|713|><|957|><|823|><|154|><|1649|><|1286|><|508|><|214|><|1760|><|1250|><|456|><|1352|><|1368|><|921|><|615|><|5|><|code_end|>
-really<|t_0.36|><|code_start|><|55|><|420|><|1008|><|1659|><|27|><|644|><|1266|><|617|><|761|><|1712|><|109|><|1465|><|1587|><|503|><|1541|><|619|><|197|><|1019|><|817|><|269|><|377|><|362|><|1381|><|507|><|1488|><|4|><|1695|><|code_end|>
-enjoyable<|t_0.49|><|code_start|><|678|><|501|><|864|><|319|><|288|><|1472|><|1341|><|686|><|562|><|1463|><|619|><|1563|><|471|><|911|><|730|><|1811|><|1006|><|520|><|861|><|1274|><|125|><|1431|><|638|><|621|><|153|><|876|><|1770|><|437|><|987|><|1653|><|1109|><|898|><|1285|><|80|><|593|><|1709|><|843|><|code_end|>
-and<|t_0.15|><|code_start|><|1285|><|987|><|303|><|1037|><|730|><|1164|><|502|><|120|><|1737|><|1655|><|1318|><|code_end|>
-it<|t_0.09|><|code_start|><|848|><|1366|><|395|><|1601|><|1513|><|593|><|1302|><|code_end|>
-looks<|t_0.27|><|code_start|><|1281|><|1266|><|1755|><|572|><|248|><|1751|><|1257|><|695|><|1380|><|457|><|659|><|585|><|1315|><|1105|><|1776|><|736|><|24|><|736|><|654|><|1027|><|code_end|>
-lovely<|t_0.56|><|code_start|><|634|><|596|><|1766|><|1556|><|1306|><|1285|><|1481|><|1721|><|1123|><|438|><|1246|><|1251|><|795|><|659|><|1381|><|1658|><|217|><|1772|><|562|><|952|><|107|><|1129|><|1112|><|467|><|550|><|1079|><|840|><|1615|><|1469|><|1380|><|168|><|917|><|836|><|1827|><|437|><|583|><|67|><|595|><|1087|><|1646|><|1493|><|1677|><|code_end|>)";
-
-    // audio data for 0.3 version
-    outetts_version tts_version = get_tts_version(model_ttc);
-    if (tts_version == OUTETTS_V0_3) {
-        audio_text = std::regex_replace(audio_text, std::regex(R"(<\|text_sep\|>)"), "<|space|>");
-        audio_data = std::regex_replace(audio_data, std::regex(R"(<\|code_start\|>)"), "");
-        audio_data = std::regex_replace(audio_data, std::regex(R"(<\|code_end\|>)"), "<|space|>");
+    if (mtmd_gen_audio_get_info(mctx.get()).type == MTMD_GEN_AUDIO_TYPE_NONE) {
+        LOG_ERR("mmproj does not support audio generation\n");
+        return 1;
     }
 
-    // load speaker if given
-    if (!params.vocoder.speaker_file.empty()) {
-        LOG_INF("%s: loading speaker ..\n", __func__);
-        json speaker = speaker_from_file(params.vocoder.speaker_file);
-        if (speaker.empty()) {
-            LOG_ERR("%s: Failed to load speaker file '%s'\n", __func__, params.vocoder.speaker_file.c_str());
+    //
+    // stage 0: process speaker reference file, if any
+    //
+
+    mtmd::bitmap_ptr speaker_bitmap;
+    if (!params.tts_speaker_file.empty()) {
+        auto wrapper = mtmd_helper_bitmap_init_from_file(mctx.get(), params.tts_speaker_file.c_str(), false);
+        if (!wrapper.bitmap) {
+            LOG_ERR("failed to load speaker file %s\n", params.tts_speaker_file.c_str());
             return 1;
         }
-        audio_text = audio_text_from_speaker(speaker, tts_version);
-        audio_data = audio_data_from_speaker(speaker, tts_version);
+        speaker_bitmap.reset(wrapper.bitmap);
     }
 
-    // process prompt and generate voice codes
-    {
-        LOG_INF("%s: constructing prompt ..\n", __func__);
-
-        std::vector<llama_token> prompt_inp;
-
-        prompt_init(prompt_inp, vocab);
-
-        prompt_add(prompt_inp, vocab, audio_text, false, true);
-
-        // convert the input text into the necessary format expected by OuteTTS
-        {
-            std::string prompt_clean = process_text(params.prompt, tts_version);
-            if (params.vocoder.use_guide_tokens) {
-                guide_tokens = prepare_guide_tokens(vocab, prompt_clean, tts_version);
-            }
-
-            LOG_INF("%s: prompt: '%s'\n", __func__, prompt_clean.c_str());
-
-            prompt_add(prompt_inp, vocab, prompt_clean, false, true);
-        }
-
-        prompt_add(prompt_inp, vocab, "<|text_end|>\n", false, true);
-
-        if (!params.vocoder.speaker_file.empty()) {
-            prompt_add(prompt_inp, vocab, audio_data, false, true);
-        } else {
-            // disabled to save time on tokenizing each time
-#if 1
-            const std::string voice_data = audio_data;
-
-            auto tmp = common_tokenize(vocab, voice_data, false, true);
-
-            std::ostringstream tokens_oss;
-            for (size_t i = 0; i < tmp.size(); ++i) {
-                tokens_oss << tmp[i] << ", ";
-            }
-            LOG_INF("\n\n%s: llama tokens: %s\n\n", __func__, tokens_oss.str().c_str());
-
-            prompt_add(prompt_inp, tmp);
-#else
-            prompt_add(prompt_inp, llama_tokens {
-                151667, 198, 1782, 155780, 151669, 151929, 152412, 152308, 152585,
-                152460, 153375, 151670, 198, 74455, 155808, 151669, 151799,
-                151873, 151863, 152446, 152372, 152204, 152728, 152229, 152470,
-                151970, 153413, 152419, 153334, 153289, 153374, 153199, 152040,
-                153260, 152721, 152680, 153297, 152419, 153248, 152400, 152691,
-                153368, 153437, 151670, 198, 1722, 155828, 151669, 152607,
-                152256, 152991, 152299, 152688, 153163, 153016, 152789, 153198,
-                152712, 151911, 153107, 152623, 152170, 152395, 152852, 152207,
-                152461, 153321, 153309, 151750, 152137, 153340, 152573, 152267,
-                153347, 151789, 152681, 153339, 151992, 152512, 151751, 152179,
-                153434, 153180, 152900, 153440, 152474, 153122, 153129, 151904,
-                152311, 151670, 198, 1499, 155791, 151669, 152276, 152454,
-                153354, 152544, 153204, 153272, 152708, 153433, 152319, 153226,
-                153043, 152325, 153267, 152622, 151670, 198, 4250, 155797,
-                151669, 153454, 153342, 151989, 152458, 153420, 152303, 152271,
-                152827, 153036, 153196, 151708, 153263, 152561, 153207, 152213,
-                152112, 153204, 151722, 152542, 151670, 198, 19789, 155796,
-                151669, 153353, 153182, 152345, 152471, 152477, 153014, 152002,
-                152191, 151734, 152312, 152810, 152237, 153224, 153169, 153224,
-                152244, 153387, 153404, 151670, 198, 16069, 155811, 151669,
-                152265, 151946, 151808, 152412, 152363, 152305, 153156, 152733,
-                152810, 153157, 152016, 152100, 152069, 153234, 152317, 152589,
-                152707, 153121, 153341, 152159, 152114, 153156, 153001, 153504,
-                153376, 152272, 152433, 152325, 151941, 151670, 198, 285,
-                155788, 151669, 152238, 152255, 153427, 152318, 153009, 152381,
-                152474, 152680, 152157, 153255, 152324, 151682, 151670, 198,
-                32955, 155804, 151669, 153490, 153419, 152364, 152405, 152682,
-                152206, 152078, 153369, 152725, 153193, 153027, 152946, 152488,
-                153070, 151883, 152890, 152489, 153144, 153375, 152358, 151685,
-                152494, 152117, 152740, 151670, 198, 37448, 480, 155840, 151669,
-                151902, 152720, 153377, 152027, 152378, 152821, 153207, 153459,
-                153028, 153068, 152507, 153255, 152158, 152921, 151958, 152609,
-                152748, 152822, 152286, 151714, 152730, 152377, 152353, 152470,
-                152606, 152162, 152186, 153071, 152244, 153118, 153375, 153018,
-                152712, 153098, 152976, 152336, 151843, 153202, 152297, 151736,
-                153380, 153502, 152702, 152115, 153181, 152735, 153277, 153457,
-                152393, 153112, 152595, 151670, 198, 19098, 155808, 151669,
-                152464, 153452, 152595, 153312, 151937, 151933, 153197, 152239,
-                153163, 152922, 153402, 152034, 152591, 153438, 152215, 151673,
-                152005, 151785, 152642, 151924, 153278, 151805, 151974, 153482,
-                152718, 152862, 153347, 151670, 198, 72, 155780, 151669, 151795,
-                152111, 152746, 152377, 153471, 152309, 151670, 198, 19016,
-                155788, 151669, 153181, 152271, 152190, 152842, 152224, 152701,
-                152939, 152536, 152091, 151815, 152733, 151672, 151670, 198,
-                14689, 155788, 151669, 152291, 152072, 152942, 151734, 153042,
-                153504, 152589, 153333, 151839, 151941, 153038, 153180, 151670,
-                198, 36996, 8303, 155832, 151669, 152231, 152256, 152835,
-                152801, 152985, 153400, 152393, 152818, 152765, 152249, 152600,
-                151699, 152302, 152752, 153018, 153009, 151992, 153054, 152847,
-                153354, 153228, 152662, 153355, 152532, 153393, 151782, 152458,
-                152048, 152757, 152428, 153195, 151906, 153006, 153178, 153250,
-                152331, 152284, 152780, 153138, 153319, 151980, 153142, 152418,
-                152228, 152733, 151670, 198, 9096, 155801, 151669, 151698,
-                153321, 152217, 153039, 152935, 153400, 152122, 152531, 153106,
-                152169, 152892, 152957, 151851, 152427, 152826, 152451, 151851,
-                152901, 152885, 152594, 153446, 153080, 151670, 198, 14689,
-                155795, 151669, 152658, 151700, 153321, 152450, 152530, 153191,
-                151673, 151690, 151698, 152714, 152846, 152981, 153171, 153384,
-                153364, 153188, 153246, 151670, 198, 1055, 155779, 151669,
-                151869, 152388, 152711, 153334, 151736, 151670, 198, 1782,
-                155780, 151669, 153483, 153240, 152241, 152558, 152697, 153046,
-                151670, 198, 5804, 1363, 155820, 151669, 152941, 152764, 152605,
-                153034, 153434, 153372, 153347, 151887, 152453, 152758, 152133,
-                152510, 152694, 152431, 152321, 153088, 152676, 152223, 152581,
-                152459, 152015, 152502, 153063, 152712, 153294, 153451, 153032,
-                152903, 152859, 152989, 151748, 152669, 152661, 152650, 152409,
-                151861, 151670, 198, 300, 7973, 155828, 151669, 153095, 152469,
-                152988, 152894, 151819, 152391, 153019, 152058, 153062, 153230,
-                151826, 152112, 152306, 152264, 152769, 153390, 152384, 152435,
-                152790, 153393, 152983, 152540, 152252, 152034, 153107, 152540,
-                151919, 151893, 152558, 152817, 152946, 152956, 152129, 152715,
-                153131, 153490, 151734, 152271, 152707, 151734, 153321, 152450,
-                151670, 198, 8088, 155792, 151669, 152452, 153497, 153353,
-                152679, 152533, 152382, 152374, 152611, 153341, 153163, 152285,
-                153411, 152495, 153141, 152320, 151670, 198, 1199, 155781,
-                151669, 151764, 152360, 153295, 152634, 153342, 152199, 152271,
-                151670, 198, 43366, 155799, 151669, 152308, 151682, 152889,
-                152016, 152385, 152629, 152495, 151826, 153321, 152958, 152180,
-                151886, 153432, 152922, 152128, 153024, 153040, 152593, 152287,
-                151677, 151670, 198, 53660, 155808, 151669, 151727, 152092,
-                152680, 153331, 151699, 152316, 152938, 152289, 152433, 153384,
-                151781, 153137, 153259, 152175, 153213, 152291, 151869, 152691,
-                152489, 151941, 152049, 152034, 153053, 152179, 153160, 151676,
-                153367, 151670, 198, 268, 4123, 480, 155821, 151669, 152350,
-                152173, 152536, 151991, 151960, 153144, 153013, 152358, 152234,
-                153135, 152291, 153235, 152143, 152583, 152402, 153483, 152678,
-                152192, 152533, 152946, 151797, 153103, 152310, 152293, 151825,
-                152548, 153442, 152109, 152659, 153325, 152781, 152570, 152957,
-                151752, 152265, 153381, 152515, 151670, 198, 437, 155787,
-                151669, 152957, 152659, 151975, 152709, 152402, 152836, 152174,
-                151792, 153409, 153327, 152990, 151670, 198, 275, 155781,
-                151669, 152520, 153038, 152067, 153273, 153185, 152265, 152974,
-                151670, 198, 94273, 155799, 151669, 152953, 152938, 153427,
-                152244, 151920, 153423, 152929, 152367, 153052, 152129, 152331,
-                152257, 152987, 152777, 153448, 152408, 151696, 152408, 152326,
-                152699, 151670, 198, 385, 16239, 155828, 151669, 152306, 152268,
-                153438, 153228, 152978, 152957, 153153, 153393, 152795, 152110,
-                152918, 152923, 152467, 152331, 153053, 153330, 151889, 153444,
-                152234, 152624, 151779, 152801, 152784, 152139, 152222, 152751,
-                152512, 153287, 153141, 153052, 151840, 152589, 152508, 153499,
-                152109, 152255, 151739, 152267, 152759, 153318, 153165, 153349,
-                151670,});
-#endif
-        }
-
-        // print the prompt token-by-token
-
-        LOG("\n");
-
-        for (auto id : prompt_inp) {
-            LOG("%s", common_token_to_piece(ctx_ttc, id).c_str());
-        }
+    mtmd_helper::gen_audio gen(lctx, mctx.get());
+    mtmd_helper_gen_audio_inp inp{};
+    inp.seq_id      = 0;
+    inp.prompt      = params.prompt.c_str();
+    inp.prompt_len  = params.prompt.size();
+    inp.speaker_ref = speaker_bitmap.get();
+    inp.lang        = params.tts_lang.c_str();
+    inp.top_k       = params.sampling.top_k;
+    inp.top_p       = params.sampling.top_p;
+    inp.out_type    = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
 
-        LOG_INF("%s: prompt size: %d\n", __func__, (int) prompt_inp.size());
+    //
+    // stage 1: process prompt via backbone model, generate semantic representation
+    //
 
-        LOG("\n");
-
-        // create a llama_batch
-        // we use this object to submit token data for decoding
-        llama_batch batch = llama_batch_init(std::max(prompt_inp.size(), (size_t) n_parallel), 0, n_parallel);
-
-        std::vector<llama_seq_id> seq_ids(n_parallel, 0);
-        for (int32_t i = 0; i < n_parallel; ++i) {
-            seq_ids[i] = i;
-        }
-
-        // evaluate the initial prompt
-        for (size_t i = 0; i < prompt_inp.size(); ++i) {
-            common_batch_add(batch, prompt_inp[i], i, seq_ids, false);
-        }
-        GGML_ASSERT(batch.n_tokens == (int) prompt_inp.size());
+    if (gen.set_input(&inp) != 0) {
+        LOG_ERR("set_input failed\n");
+        return 1;
+    }
 
-        // llama_decode will output logits only for the last token of the prompt
-        batch.logits[batch.n_tokens - 1] = true;
+    const int64_t t_prompt_start_us = ggml_time_us();
 
-        if (llama_decode(ctx_ttc, batch) != 0) {
-            LOG_ERR("%s: llama_decode() failed\n", __func__);
+    for (;;) {
+        int32_t ret = gen.step_prompt(params.n_batch);
+        if (ret < 0) {
+            LOG_ERR("prompt processing failed\n");
             return 1;
         }
-
-        if (n_parallel > 1) {
-            LOG_INF("\n\n%s: generating %d sequences ...\n", __func__, n_parallel);
-        }
-
-        llama_synchronize(ctx_ttc);
-
-        LOG_INF("%s: time for prompt: %.3f ms\n\n", __func__, (ggml_time_us() - t_main_start) / 1000.0f);
-
-        const auto t_dec_start = ggml_time_us();
-
-        // main loop
-
-        // remember the batch index of the last token for each parallel sequence
-        // we need this to determine which logits to sample from
-        std::vector<int32_t> i_batch(n_parallel, batch.n_tokens - 1);
-
-        int n_past   = batch.n_tokens;
-        int n_decode = 0;
-
-        bool next_token_uses_guide_token = true;
-
-        while (n_decode <= n_predict) {
-            // prepare the next batch
-            common_batch_clear(batch);
-
-            // sample the next token for each parallel sequence / stream
-            for (int32_t i = 0; i < n_parallel; ++i) {
-                if (i_batch[i] < 0) {
-                    // the stream has already finished
-                    continue;
-                }
-
-                llama_token new_token_id = common_sampler_sample(smpl[i], ctx_ttc, i_batch[i]);
-
-                //guide tokens help prevent hallucinations by forcing the TTS to use the correct word
-                if (!guide_tokens.empty() && next_token_uses_guide_token && !llama_vocab_is_control(vocab, new_token_id) && !llama_vocab_is_eog(vocab, new_token_id)) {
-                    llama_token guide_token = guide_tokens[0];
-                    guide_tokens.erase(guide_tokens.begin());
-                    new_token_id = guide_token; //ensure correct word fragment is used
-                }
-
-                //this is the token id that always precedes a new word
-                next_token_uses_guide_token = (new_token_id == 198);
-
-                common_sampler_accept(smpl[i], new_token_id, true);
-
-                codes.push_back(new_token_id);
-
-                const auto * cands = common_sampler_get_candidates(smpl[i], false);
-
-                // is it an end of generation? -> mark the stream as finished
-                if (llama_vocab_is_eog(vocab, new_token_id) || n_decode == n_predict) {
-                    std::string reason;
-                    if (llama_vocab_is_eog(vocab, new_token_id)) {
-                        reason = "eos";
-                    } else {
-                        reason = "n_predict";
-                    }
-
-                    i_batch[i] = -1;
-
-                    LOG("\n");
-                    if (n_parallel > 1) {
-                        LOG_CNT("\n");
-                        LOG_INF("%s: stream %d finished at n_past = %d, reason = '%s'\n", __func__, i, n_past, reason.c_str());
-                    }
-
-                    continue;
-                }
-
-                {
-                    const float p = cands->data[cands->selected].p;
-
-                    const int col = std::max(0, std::min((int) k_colors.size() - 1, (int) ((3*p)*float(k_colors.size()))));
-
-                    LOG_CNT("%s%d%s", k_colors[col].c_str(), i, "\033[0m");
-                    //LOG_CNT("%d", i);
-                }
-
-                i_batch[i] = batch.n_tokens;
-
-                // push this new token for next evaluation
-                common_batch_add(batch, new_token_id, n_past, { i }, true);
-            }
-
-            // all streams are finished
-            if (batch.n_tokens == 0) {
-                break;
-            }
-
-            n_decode += 1;
-            n_past += 1;
-
-            // evaluate the current batch with the transformer model
-            if (llama_decode(ctx_ttc, batch)) {
-                LOG_ERR("%s : failed to eval, return code %d\n", __func__, 1);
-                return 1;
-            }
+        if (ret == 0) {
+            break;
         }
-
-        llama_batch_free(batch);
-
-        LOG("\n");
-        LOG_INF("%s: time for decoder:       %.3f ms\n", __func__, (ggml_time_us() - t_dec_start) / 1000.0f);
     }
 
-    common_perf_print(ctx_ttc, smpl[0]);
+    const llama_vocab * vocab = llama_model_get_vocab(model);
 
-    //std::vector<llama_token> codes = {198, 88225, 155856, 151669, 152205,
-    //    153064, 152537, 153421, 153209, 152524, 151689, 152993, 152438, 152695,
-    //    153091, 152945, 152829, 152534, 152934, 153020, 151997, 152263, 153010,
-    //    153146, 152399, 153208, 152496, 151793, 152848, 152263, 152571, 153286,
-    //    152227, 153300, 152934, 152263, 153208, 152263, 152965, 152430, 152296,
-    //    153146, 152920, 152376, 152556, 153363, 151775, 152044, 152972, 152690,
-    //    153379, 152368, 152233, 153422, 152490, 151996, 152022, 151694, 152061,
-    //    153238, 152539, 153356, 152640, 153021, 153123, 151962, 153094, 151670,
-    //    198, 20339, 13189, 155824, 151669, 152070, 152007, 152910, 151683,
-    //    152000, 152373, 152760, 152046, 151735, 152334, 152394, 153073, 152908,
-    //    151856, 151953, 153247, 153293, 151903, 153480, 153168, 152478, 153359,
-    //    153429, 151905, 151678, 152567, 152411, 152165, 152556, 153075, 153424,
-    //    151993, 152999, 153078, 152151, 152088, 153389, 152484, 151874, 151670,
-    //    198, 285, 155784, 151669, 152226, 152126, 152638, 153215, 151729,
-    //    152959, 153479, 153059, 151838, 151670, 198, 1782, 155783, 151669,
-    //    153288, 153055, 153314, 152497, 152962, 152741, 152076, 153253, 151670,
-    //    198, 471, 16488, 155825, 151669, 152060, 152916, 151893, 153469, 152501,
-    //    152080, 152743, 151932, 153161, 152096, 152761, 152698, 153401, 153242,
-    //    153336, 152441, 152838, 153467, 152706, 153496, 153310, 152422, 153360,
-    //    153115, 152763, 151998, 152373, 153450, 152554, 151968, 153323, 152055,
-    //    152468, 153111, 153358, 152813, 152010, 151770, 152823, 152960, 151670,
-    //    198, 22627, 155823, 151669, 152814, 152366, 153484, 152931, 153441,
-    //    152164, 152877, 152915, 153463, 151692, 152911, 152747, 152776, 151831,
-    //    153449, 151882, 152975, 152031, 152513, 153150, 152448, 152667, 153133,
-    //    153189, 152619, 153466, 152054, 152106, 153119, 152277, 152439, 153109,
-    //    152997, 152141, 153154, 153256, 153311, 151922, 151670, 198, 1055,
-    //    155781, 151669, 152633, 151850, 153060, 153270, 152560, 153348, 152729,
-    //    151670, 198, 25312, 155803, 151669, 152521, 153403, 152561, 153337,
-    //    153383, 152199, 153493, 153326, 151830, 152254, 152248, 152349, 152153,
-    //    153007, 151823, 153037, 152575, 152457, 152406, 152592, 153116, 153365,
-    //    153456, 151670, 198, 88225, 155817, 151669, 153271, 151925, 152218,
-    //    152418, 152253, 153140, 151903, 153151, 152626, 152338, 152647, 153464,
-    //    152785, 152768, 151711, 152037, 152033, 151804, 152216, 151701, 151855,
-    //    152348, 152995, 152955, 152905, 152342, 152340, 153391, 153453, 152418,
-    //    153415, 151990, 153083, 152884, 151670, 198, 151668, 198, 151645};
+    auto sample_semantic_code = [&]() -> llama_token {
+        llama_token t = common_sampler_sample(smpl, lctx, -1);
+        common_sampler_accept(smpl, t, true);
+        return t;
+    };
 
-    {
-        const std::string inp_txt = common_detokenize(ctx_ttc, codes, true);
+    const int max_new = params.n_predict > 0 ? params.n_predict : 512;
+    int n_frames = 0;
+    llama_token sampled = sample_semantic_code();
+    const float * h_state = llama_get_embeddings_ith(lctx, -1);
 
-        LOG("\n");
-        LOG_INF("codes: '%s'\n", inp_txt.c_str());
-        LOG_INF("%s: codes size: %d\n", __func__, (int) codes.size());
-    }
-
-    // remove all non-audio tokens (i.e. < 151672 || > 155772)
-    codes.erase(std::remove_if(codes.begin(), codes.end(), [](llama_token t) { return t < 151672 || t > 155772; }), codes.end());
+    tts_timings timings;
+    const int64_t t_gen_start_us = ggml_time_us();
 
-    {
-        const std::string inp_txt = common_detokenize(ctx_ttc, codes, true);
-        LOG_INF("codes audio: '%s'\n", inp_txt.c_str());
-        LOG_INF("%s: codes audio size: %d\n", __func__, (int) codes.size());
-    }
+    for (; n_frames < max_new && !llama_vocab_is_eog(vocab, sampled); n_frames++) {
+        const float * h_next = nullptr;
 
-    for (auto & token : codes) {
-        token -= 151672;
-    }
-
-    const auto t_voc_start = ggml_time_us();
-
-    const int n_codes = codes.size();
-
-    llama_batch batch = llama_batch_init(n_codes, 0, 1);
-
-    for (size_t i = 0; i < codes.size(); ++i) {
-        common_batch_add(batch, codes[i], i, { 0 }, true); // TODO: all logits?
-    }
-    GGML_ASSERT(batch.n_tokens == n_codes);
-
-    if (llama_encode(ctx_cts, batch) != 0) {
-        LOG_ERR("%s: llama_encode() failed\n", __func__);
-        return 1;
-    }
-
-    llama_synchronize(ctx_cts);
-
-    LOG_INF("%s: time for vocoder:      %.3f ms\n", __func__, (ggml_time_us() - t_voc_start) / 1000.0f);
-
-    const auto t_spec_start = ggml_time_us();
-
-#if 1
-    // spectral operations
-    const int n_embd = llama_model_n_embd_out(model_cts);
-    const float * embd = llama_get_embeddings(ctx_cts);
-
-    auto audio = embd_to_audio(embd, n_codes, n_embd, params.cpuparams.n_threads);
-
-#else
-    // read the spectrogram from a file for debugging purposes
-    std::vector<float> audio;
-    {
-        std::ifstream fin("out.bin", std::ios::binary);
-        if (!fin) {
-            LOG_ERR("%s: failed to open file '%s'\n", __func__, "out.bin");
+        // stage 2+3: semantic --> acoustic details --> audio waveform
+        //            step_gen() runs both stages and returns new h_state for next step
+        if (gen.step_gen(sampled, h_state, &h_next) != 0) {
+            LOG_ERR("step_gen failed at frame %d\n", n_frames);
             return 1;
         }
 
-        std::vector<float> embd;
-
-        int n_codes;
-        int n_embd;
-
-        fin.read(reinterpret_cast<char *>(&n_codes), sizeof(int));
-        fin.read(reinterpret_cast<char *>(&n_embd), sizeof(int));
-
-        embd.resize(n_codes * n_embd);
-        fin.read(reinterpret_cast<char *>(embd.data()), n_codes * n_embd * sizeof(float));
-        fin.close();
-
-        LOG_INF("%s: n_codes: %d, n_embd: %d\n", __func__, n_codes, n_embd);
-
-        audio = embd_to_audio(embd.data(), n_codes, n_embd, params.cpuparams.n_threads);
+        h_state = h_next;
+        sampled = sample_semantic_code();
+        timings.report(n_frames + 1);
     }
-#endif
-
-    const int n_sr = 24000; // sampling rate
+    const double t_gen_s = (ggml_time_us() - t_gen_start_us) / 1e6;
 
-    // zero out first 0.25 seconds
-    for (int i = 0; i < 24000/4; ++i) {
-        audio[i] = 0.0f;
+    int32_t      sample_rate = 0;
+    const char * data        = nullptr;
+    size_t       data_len    = 0;
+    int64_t      n_samples   = 0;
+    if (gen.get_output(&sample_rate, &data, &data_len, &n_samples) != 0) {
+        LOG_ERR("get_output failed\n");
+        return 1;
     }
 
-    LOG_INF("%s: time for spectral ops: %.3f ms\n", __func__, (ggml_time_us() - t_spec_start) / 1000.0f);
-    LOG_INF("%s: total time:            %.3f ms\n", __func__, (ggml_time_us() - t_main_start) / 1000.0f);
-
-    int retval = 0;
+    LOG_INF("generated %d frames, %zu bytes of WAV audio (%d Hz)\n", n_frames, data_len, sample_rate);
 
-    if (save_wav16(params.out_file, audio, n_sr)) {
-        LOG_INF("%s: audio written to file '%s'\n", __func__, params.out_file.c_str());
-    } else {
-        retval = ENOENT;
+    const double t_prompt_s = (t_gen_start_us - t_prompt_start_us) / 1e6;
+    const double t_total_s  = t_prompt_s + t_gen_s;
+    const double audio_s    = sample_rate > 0 ? (double) n_samples / sample_rate : 0.0;
+    LOG_INF("timings: prompt eval %.2fs + generation %.2fs = total %.2fs\n", t_prompt_s, t_gen_s, t_total_s);
+    LOG_INF("         output audio = %.2fs (audio time = %.2fx process time)\n", audio_s, t_total_s > 0 ? audio_s / t_total_s : 0.0);
+    FILE * f = fopen(params.out_file.c_str(), "wb");
+    if (!f) {
+        LOG_ERR("failed to open %s\n", params.out_file.c_str());
+        return 1;
     }
+    fwrite(data, 1, data_len, f);
+    fclose(f);
+    LOG_INF("wrote %s\n", params.out_file.c_str());
 
     llama_backend_free();
-
-    return retval;
+    return 0;
 }