diff --git a/.github/workflows/README.md b/.github/workflows/README.md new file mode 100644 index 0000000..c0c52e3 --- /dev/null +++ b/.github/workflows/README.md @@ -0,0 +1,8 @@ +# GitHub Actions Workflows + +| Workflow | Trigger | Description | +|----------|---------|-------------| +| [CI](ci.yml) | Push to `main`, PRs, manual | Runs `cargo fmt`, `clippy`, tests, and doc build | +| [Release-plz](release-plz.yml) | Push to `main`, manual | Automates crate releases to crates.io and opens release PRs | +| [Update benchmark table](update-bench.yml) | Push to `main` (when `bench/results/**.csv` changes), manual | Regenerates the benchmark table in README.md | +| [Export ONNX](export-onnx.yml) | Manual | Exports Qwen3-TTS to ONNX (FP32 + INT4), validates, and uploads to HuggingFace. Supports `voicedesign` (1.7B) and `clone` (0.6B Base) variants via input selector. | diff --git a/.github/workflows/export-onnx.yml b/.github/workflows/export-onnx.yml index 3bb00be..3c7f55f 100644 --- a/.github/workflows/export-onnx.yml +++ b/.github/workflows/export-onnx.yml @@ -3,13 +3,20 @@ name: Export ONNX & Publish to HuggingFace on: workflow_dispatch: inputs: + variant: + description: "Model variant to export" + default: "voicedesign" + type: choice + options: + - voicedesign + - clone model_id: - description: "HuggingFace source model" - default: "Qwen/Qwen3-TTS-12Hz-1.7B-VoiceDesign" + description: "HuggingFace source model (leave default for selected variant)" + default: "" type: string hf_repo: - description: "Target HuggingFace repo (org/name)" - default: "wavekat/Qwen3-TTS-1.7B-VoiceDesign-ONNX" + description: "Target HuggingFace repo (leave default for selected variant)" + default: "" type: string runner: description: "GitHub Actions runner label" @@ -20,7 +27,7 @@ on: default: false type: boolean revision: - description: "HuggingFace revision/branch to push to (e.g. v1, 2026-04-06). Defaults to main." + description: "HuggingFace revision/branch to push to. Defaults to main." default: "" type: string @@ -29,6 +36,11 @@ jobs: runs-on: ${{ inputs.runner }} timeout-minutes: 120 + env: + MODEL_ID: ${{ inputs.model_id != '' && inputs.model_id || (inputs.variant == 'clone' && 'Qwen/Qwen3-TTS-12Hz-0.6B-Base' || 'Qwen/Qwen3-TTS-12Hz-1.7B-VoiceDesign') }} + HF_REPO: ${{ inputs.hf_repo != '' && inputs.hf_repo || (inputs.variant == 'clone' && 'wavekat/Qwen3-TTS-0.6B-Base-ONNX' || 'wavekat/Qwen3-TTS-1.7B-VoiceDesign-ONNX') }} + OUTPUT_DIR: ${{ inputs.variant == 'clone' && './output/qwen3-tts-0.6b-base' || './output/qwen3-tts-1.7b-voicedesign' }} + steps: - name: Free disk space run: | @@ -61,12 +73,23 @@ jobs: working-directory: tools/qwen3-tts-onnx run: make venv - - name: Export FP32 ONNX models + - name: Export FP32 ONNX models (VoiceDesign) + if: inputs.variant == 'voicedesign' working-directory: tools/qwen3-tts-onnx env: HF_TOKEN: ${{ secrets.HF_TOKEN }} run: | - make export MODEL_ID="${{ inputs.model_id }}" + make export MODEL_ID="${{ env.MODEL_ID }}" + echo "=== Disk after export ===" + df -h / + + - name: Export FP32 ONNX models (Clone — 6 components) + if: inputs.variant == 'clone' + working-directory: tools/qwen3-tts-onnx + env: + HF_TOKEN: ${{ secrets.HF_TOKEN }} + run: | + make clone-export CLONE_MODEL_ID="${{ env.MODEL_ID }}" echo "=== Disk after export ===" df -h / @@ -75,7 +98,7 @@ jobs: working-directory: tools/qwen3-tts-onnx env: HF_TOKEN: ${{ secrets.HF_TOKEN }} - run: make validate MODEL_ID="${{ inputs.model_id }}" + run: make validate MODEL_ID="${{ env.MODEL_ID }}" OUTPUT_DIR="${{ env.OUTPUT_DIR }}" - name: Free PyTorch model cache run: | @@ -83,17 +106,32 @@ jobs: echo "=== Disk after cache cleanup ===" df -h / - - name: Quantize to INT4 + - name: Quantize to INT4 (VoiceDesign) + if: inputs.variant == 'voicedesign' working-directory: tools/qwen3-tts-onnx run: | make quantize echo "=== Disk after quantize ===" df -h / - - name: Package for HuggingFace + - name: Quantize to INT4 (Clone — encoders stay FP32) + if: inputs.variant == 'clone' + working-directory: tools/qwen3-tts-onnx + run: | + make clone-base-preset CLONE_MODEL_ID="${{ env.MODEL_ID }}" + echo "=== Disk after quantize ===" + df -h / + + - name: Package for HuggingFace (VoiceDesign) + if: inputs.variant == 'voicedesign' working-directory: tools/qwen3-tts-onnx run: make hf + - name: Package for HuggingFace (Clone) + if: inputs.variant == 'clone' + working-directory: tools/qwen3-tts-onnx + run: make clone-hf + - name: Upload to HuggingFace working-directory: tools/qwen3-tts-onnx env: @@ -102,10 +140,10 @@ jobs: .venv/bin/pip install -q huggingface_hub REVISION="${{ inputs.revision }}" .venv/bin/huggingface-cli upload \ - "${{ inputs.hf_repo }}" \ - ./output/qwen3-tts-1.7b-voicedesign/ \ + "${{ env.HF_REPO }}" \ + ${{ env.OUTPUT_DIR }}/ \ --repo-type model \ - --commit-message "export: ${{ inputs.model_id }} (run #${{ github.run_number }})" \ + --commit-message "export: ${{ env.MODEL_ID }} (run #${{ github.run_number }})" \ ${REVISION:+--revision "$REVISION"} - name: Job summary @@ -113,14 +151,15 @@ jobs: run: | REVISION="${{ inputs.revision }}" REVISION_DISPLAY="${REVISION:-main}" - HF_URL="https://huggingface.co/${{ inputs.hf_repo }}/tree/${REVISION_DISPLAY}" + HF_URL="https://huggingface.co/${{ env.HF_REPO }}/tree/${REVISION_DISPLAY}" cat >> "$GITHUB_STEP_SUMMARY" < **Requires a reference WAV file** (`ref.wav`) — a short mono clip of the voice +> you want to clone, plus a transcript of what is spoken in the clip. -println!("{}s at {} Hz", audio.duration_secs(), audio.sample_rate()); +```rust +use wavekat_tts::AudioFrame; +use wavekat_tts::backends::qwen3_tts::{Qwen3TtsClone, CloneRequest}; +// use wavekat_tts::backends::qwen3_tts::{ModelConfig, ModelPrecision}; + +fn main() { + let ref_audio = AudioFrame::from_wav("ref.wav").unwrap(); // 24 kHz mono WAV + let tts = Qwen3TtsClone::new().unwrap(); // auto-downloads 0.6B Base INT4 model + + // For FP32 precision: + // let config = ModelConfig::default().with_precision(ModelPrecision::Fp32); + // let tts = Qwen3TtsClone::from_config(config).unwrap(); + + let req = CloneRequest::new( + "Text to say in the cloned voice", + ref_audio.samples(), + 24000, + "Transcript of the reference clip.", + ).with_language("en"); + let audio = tts.synthesize_clone(&req).unwrap(); + audio.write_wav("clone_output.wav").unwrap(); + println!("Wrote clone_output.wav ({:.2}s)", audio.duration_secs()); +} ``` Model files are cached by the HF Hub client at `$HF_HOME/hub/` (default `~/.cache/huggingface/hub/`). @@ -92,10 +121,17 @@ Two trait families: Generate a WAV file from text (model files are auto-downloaded on first run): ```sh +# VoiceDesign (1.7B) cargo run --example synthesize --features qwen3-tts -- "Hello, world\!" cargo run --example synthesize --features qwen3-tts -- --instruction "Speak in a warm, friendly tone." "Give every small business the voice of a big one." -cargo run --example synthesize --features qwen3-tts -- --precision fp32 "Hello" -cargo run --example synthesize --features qwen3-tts -- --model-dir /path/to/model --output hello.wav "Hello" +# cargo run --example synthesize --features qwen3-tts -- --precision fp32 "Hello, world\!" + +# Voice Clone (0.6B) +cargo run --example synthesize_clone --features qwen3-tts -- \ + --ref-audio ref.wav --ref-text "Transcript of the reference clip." \ + "Text to synthesize in the cloned voice." +# cargo run --example synthesize_clone --features qwen3-tts -- --precision fp32 \ +# --ref-audio ref.wav --ref-text "Transcript." "Text to synthesize." ``` ## Performance diff --git a/crates/wavekat-tts/Cargo.toml b/crates/wavekat-tts/Cargo.toml index fb5028c..cabe924 100644 --- a/crates/wavekat-tts/Cargo.toml +++ b/crates/wavekat-tts/Cargo.toml @@ -13,7 +13,7 @@ categories = ["multimedia::audio"] default = [] # Local inference backends (all ONNX-based) -qwen3-tts = ["dep:ort", "dep:ndarray", "dep:tokenizers", "dep:npyz", "dep:rand", "dep:hf-hub"] +qwen3-tts = ["dep:ort", "dep:ndarray", "dep:tokenizers", "dep:npyz", "dep:rand", "dep:hf-hub", "dep:realfft"] cosyvoice = ["dep:ort", "dep:ndarray"] # Execution providers — composable with any ONNX backend feature @@ -22,7 +22,7 @@ cuda = ["ort?/cuda"] tensorrt = ["ort?/tensorrt"] [dependencies] -wavekat-core = { version = "0.0.5", features = ["wav"] } +wavekat-core = { version = "0.0.7", features = ["wav"] } thiserror = "2" serde = { version = "1", features = ["derive"] } serde_json = "1" @@ -34,11 +34,16 @@ tokenizers = { version = "0.21", optional = true, default-features = false, feat npyz = { version = "0.8", optional = true } rand = { version = "0.9", optional = true } hf-hub = { version = "0.5", optional = true, default-features = false, features = ["ureq"] } +realfft = { version = "3", optional = true } [[example]] name = "synthesize" required-features = ["qwen3-tts"] +[[example]] +name = "synthesize_clone" +required-features = ["qwen3-tts"] + [[example]] name = "bench_rtf" required-features = ["qwen3-tts"] diff --git a/crates/wavekat-tts/examples/synthesize.rs b/crates/wavekat-tts/examples/synthesize.rs index 768efbe..ce06cd4 100644 --- a/crates/wavekat-tts/examples/synthesize.rs +++ b/crates/wavekat-tts/examples/synthesize.rs @@ -259,4 +259,5 @@ fn synthesize_one( audio.write_wav(output).expect("failed to write WAV"); eprintln!("Wrote {}", output.display()); + eprintln!("Done."); } diff --git a/crates/wavekat-tts/examples/synthesize_clone.rs b/crates/wavekat-tts/examples/synthesize_clone.rs new file mode 100644 index 0000000..39ac729 --- /dev/null +++ b/crates/wavekat-tts/examples/synthesize_clone.rs @@ -0,0 +1,165 @@ +//! Synthesize text in a cloned voice using Qwen3-TTS 0.6B Base. +//! +//! Usage: +//! cargo run --example synthesize_clone --features qwen3-tts -- [OPTIONS] +//! +//! Options: +//! --ref-audio Reference audio WAV (24 kHz mono) +//! --ref-text Transcript of the reference audio +//! --text Text to synthesize in the cloned voice +//! --language Language code (default: en) +//! --model-dir Model directory (default: auto-download) +//! --precision Model precision: int4 (default) or fp32 +//! --provider Execution provider: cpu (default), cuda, tensorrt, coreml +//! --output Output WAV path (default: clone_output.wav) +//! +//! Example: +//! cargo run --example synthesize_clone --features qwen3-tts -- \ +//! --ref-audio ref.wav \ +//! --ref-text "Give every small business the voice of a big one." \ +//! --text "Your customers deserve a voice they can trust." + +use std::path::PathBuf; + +use wavekat_tts::backends::qwen3_tts::{ + CloneRequest, ExecutionProvider, ModelConfig, ModelPrecision, Qwen3TtsClone, +}; +use wavekat_tts::AudioFrame; + +fn main() { + let args: Vec = std::env::args().skip(1).collect(); + + let mut model_dir: Option = None; + let mut precision = ModelPrecision::Int4; + let mut provider = ExecutionProvider::Cpu; + let mut language = "en".to_string(); + let mut output = PathBuf::from("clone_output.wav"); + let mut ref_audio_path: Option = None; + let mut ref_text: Option = None; + let mut text: Option = None; + + let mut i = 0; + while i < args.len() { + match args[i].as_str() { + "--model-dir" => { + i += 1; + model_dir = Some(PathBuf::from(&args[i])); + } + "--precision" => { + i += 1; + precision = match args[i].as_str() { + "int4" => ModelPrecision::Int4, + "fp32" => ModelPrecision::Fp32, + other => { + eprintln!("error: unknown precision \"{other}\""); + std::process::exit(1); + } + }; + } + "--provider" => { + i += 1; + provider = match args[i].as_str() { + "cpu" => ExecutionProvider::Cpu, + "cuda" => ExecutionProvider::Cuda, + "tensorrt" => ExecutionProvider::TensorRt, + "coreml" => ExecutionProvider::CoreMl, + other => { + eprintln!("error: unknown provider \"{other}\""); + std::process::exit(1); + } + }; + } + "--language" => { + i += 1; + language = args[i].clone(); + } + "--output" => { + i += 1; + output = PathBuf::from(&args[i]); + } + "--ref-audio" => { + i += 1; + ref_audio_path = Some(PathBuf::from(&args[i])); + } + "--ref-text" => { + i += 1; + ref_text = Some(args[i].clone()); + } + "--text" => { + i += 1; + text = Some(args[i].clone()); + } + other => { + eprintln!("error: unknown argument \"{other}\""); + std::process::exit(1); + } + } + i += 1; + } + + let ref_audio_path = ref_audio_path.unwrap_or_else(|| { + eprintln!("error: --ref-audio is required"); + std::process::exit(1); + }); + let ref_text = ref_text.unwrap_or_else(|| { + eprintln!("error: --ref-text is required"); + std::process::exit(1); + }); + let text = text.unwrap_or_else(|| { + eprintln!("error: --text is required"); + std::process::exit(1); + }); + + // Read reference audio WAV + eprintln!("Reading reference audio: {} ...", ref_audio_path.display()); + let ref_audio = AudioFrame::from_wav(&ref_audio_path).expect("failed to read reference WAV"); + if ref_audio.sample_rate() != 24000 { + eprintln!( + "error: reference audio must be 24 kHz, got {} Hz", + ref_audio.sample_rate() + ); + eprintln!("hint: resample with: ffmpeg -i input.wav -ar 24000 -ac 1 ref_24k.wav"); + std::process::exit(1); + } + eprintln!( + " {:.1}s, {} Hz, {} samples", + ref_audio.duration_secs(), + ref_audio.sample_rate(), + ref_audio.len(), + ); + + // Load model + eprintln!("Loading clone model ..."); + let mut config = ModelConfig::default() + .with_precision(precision) + .with_execution_provider(provider); + if let Some(dir) = model_dir { + config = config.with_dir(dir); + } + let tts = Qwen3TtsClone::from_config(config).expect("failed to load clone model"); + + // Synthesize + let request = + CloneRequest::new(&text, ref_audio.samples(), 24000, &ref_text).with_language(&language); + + eprintln!("Synthesizing: \"{text}\" (language={language})"); + let start = std::time::Instant::now(); + let audio = tts.synthesize_clone(&request).expect("synthesis failed"); + let elapsed = start.elapsed(); + + let duration = audio.duration_secs(); + let rtf = elapsed.as_secs_f64() / duration; + + eprintln!( + "Generated {} samples at {} Hz ({:.2}s) in {:.2}s (RTF: {:.2})", + audio.len(), + audio.sample_rate(), + duration, + elapsed.as_secs_f64(), + rtf, + ); + + audio.write_wav(&output).expect("failed to write WAV"); + eprintln!("Wrote {}", output.display()); + eprintln!("Done."); +} diff --git a/crates/wavekat-tts/src/backends/qwen3_tts/clone_model.rs b/crates/wavekat-tts/src/backends/qwen3_tts/clone_model.rs new file mode 100644 index 0000000..5e9be14 --- /dev/null +++ b/crates/wavekat-tts/src/backends/qwen3_tts/clone_model.rs @@ -0,0 +1,768 @@ +//! Voice-clone ONNX pipeline for Qwen3-TTS 0.6B Base. +//! +//! Chains 6 ONNX models: tokenizer_encoder → speaker_encoder → talker (prefill +//! + decode loop) → code_predictor → vocoder. +//! +//! Reference audio codes are prepended to the generated codes before vocoding, +//! then the leading reference portion is trimmed proportionally. + +use std::path::Path; +use std::sync::Mutex; + +use ndarray::{concatenate, s, Array1, Array2, Array3, Array5, Axis}; +use ort::session::Session; +use ort::value::TensorRef; +use wavekat_core::AudioFrame; + +use crate::TtsError; + +use super::mel::MelSpectrogram; +use super::model::{ + apply_execution_provider, load_npy1, load_npy2, prepare_onnx_dir, text_project, +}; +use super::sampler::{self, SamplerConfig}; +use super::tokenizer::{self, ASSISTANT, IM_START, NEWLINE, TTS_BOS, TTS_EOS, TTS_PAD}; + +// Codec control token IDs (shared across all Qwen3-TTS variants) +const CODEC_PAD: i64 = 2148; +const CODEC_BOS: i64 = 2149; +const CODEC_EOS: i64 = 2150; +const CODEC_THINK: i64 = 2154; +const CODEC_THINK_BOS: i64 = 2156; +const CODEC_THINK_EOS: i64 = 2157; + +// Model dimensions — Qwen3-TTS-12Hz-0.6B-Base +const HIDDEN_DIM: usize = 1024; +const NUM_LAYERS: usize = 28; +const NUM_KV_HEADS: usize = 8; +const HEAD_DIM: usize = 128; +const TALKER_VOCAB_SIZE: usize = 3072; +const CP_NUM_LAYERS: usize = 5; +const CP_NUM_KV_HEADS: usize = 8; +const NUM_CP_GROUPS: usize = 15; // codebook groups 1-15 +const SAMPLE_RATE: u32 = 24000; +const MAX_NEW_TOKENS: usize = 8192; + +// Tokenizer encoder constants +const TOKENIZER_CANONICAL_SAMPLES: usize = 240_000; // 10s @ 24kHz +const TOKENIZER_DOWNSAMPLE: usize = 1920; // 24000 Hz / 12.5 Hz + +/// Sampling defaults from config.json. +const TALKER_SAMPLER: SamplerConfig = SamplerConfig { + temperature: 0.9, + top_k: 50, + repetition_penalty: 1.05, +}; + +const CP_SAMPLER: SamplerConfig = SamplerConfig { + temperature: 0.9, + top_k: 50, + repetition_penalty: 1.0, +}; + +/// Talker output: (logits, hidden_state, kv_keys, kv_values). +type TalkerOutput = (Vec, Array3, Array5, Array5); + +/// All ONNX sessions and embedding tables for the 0.6B Base voice clone pipeline. +/// +/// Chains six ONNX models: tokenizer encoder → speaker encoder → talker +/// (prefill + decode loop) → code predictor → vocoder. +pub struct CloneModel { + talker_prefill: Mutex, + talker_decode: Mutex, + code_predictor: Mutex, + vocoder: Mutex, + speaker_encoder: Mutex, + tokenizer_encoder: Mutex, + + // Embedding tables (immutable after construction) + text_embedding: Array2, + text_proj_fc1_weight: Array2, + text_proj_fc1_bias: Array1, + text_proj_fc2_weight: Array2, + text_proj_fc2_bias: Array1, + talker_codec_embedding: Array2, + cp_codec_embeddings: Vec>, + + // Precomputed + tts_pad_embed: Array1, + mel: MelSpectrogram, +} + +impl CloneModel { + /// Load all 6 ONNX sessions and embedding tables from `model_dir`. + pub fn load(model_dir: &Path, config: &super::ModelConfig) -> Result { + let onnx_dir = prepare_onnx_dir(&model_dir.join(config.precision.subdir()))?; + + let load_session = |name: &str, dir: &Path| -> Result { + let path = dir.join(name); + let builder = Session::builder() + .map_err(|e| TtsError::Model(format!("session builder error: {e}")))?; + apply_execution_provider(builder, config.execution_provider)? + .commit_from_file(&path) + .map_err(|e| TtsError::Model(format!("failed to load {name}: {e}"))) + }; + + eprint!("Loading talker prefill ... "); + let talker_prefill = load_session("talker_prefill.onnx", &onnx_dir)?; + eprintln!("done"); + + eprint!("Loading talker decode ... "); + let talker_decode = load_session("talker_decode.onnx", &onnx_dir)?; + eprintln!("done"); + + eprint!("Loading code predictor ... "); + let code_predictor = load_session("code_predictor.onnx", &onnx_dir)?; + eprintln!("done"); + + eprint!("Loading vocoder ... "); + let vocoder = load_session("vocoder.onnx", &onnx_dir)?; + eprintln!("done"); + + // Speaker encoder and tokenizer encoder are FP32-only, stored at model root. + // They have external .data files, so we need prepare_onnx_dir to resolve + // HF Hub symlinks (same issue as the talker models). + let root_dir = prepare_onnx_dir(model_dir)?; + + eprint!("Loading speaker encoder ... "); + let speaker_encoder = load_session("speaker_encoder.onnx", &root_dir)?; + eprintln!("done"); + + eprint!("Loading tokenizer encoder ... "); + let tokenizer_encoder = load_session("tokenizer_encoder.onnx", &root_dir)?; + eprintln!("done"); + + eprint!("Loading embeddings ... "); + let text_embedding = load_npy2(model_dir, "embeddings/text_embedding.npy")?; + let text_proj_fc1_weight = + load_npy2(model_dir, "embeddings/text_projection_fc1_weight.npy")?; + let text_proj_fc1_bias = load_npy1(model_dir, "embeddings/text_projection_fc1_bias.npy")?; + let text_proj_fc2_weight = + load_npy2(model_dir, "embeddings/text_projection_fc2_weight.npy")?; + let text_proj_fc2_bias = load_npy1(model_dir, "embeddings/text_projection_fc2_bias.npy")?; + let talker_codec_embedding = load_npy2(model_dir, "embeddings/talker_codec_embedding.npy")?; + + let mut cp_codec_embeddings = Vec::with_capacity(NUM_CP_GROUPS); + for i in 0..NUM_CP_GROUPS { + cp_codec_embeddings.push(load_npy2( + model_dir, + &format!("embeddings/cp_codec_embedding_{i}.npy"), + )?); + } + eprintln!("done"); + + let tts_pad_raw = text_embedding.row(TTS_PAD as usize).to_owned(); + let tts_pad_embed = text_project( + &tts_pad_raw, + &text_proj_fc1_weight, + &text_proj_fc1_bias, + &text_proj_fc2_weight, + &text_proj_fc2_bias, + ); + + eprintln!("Clone model ready."); + + Ok(Self { + talker_prefill: Mutex::new(talker_prefill), + talker_decode: Mutex::new(talker_decode), + code_predictor: Mutex::new(code_predictor), + vocoder: Mutex::new(vocoder), + speaker_encoder: Mutex::new(speaker_encoder), + tokenizer_encoder: Mutex::new(tokenizer_encoder), + text_embedding, + text_proj_fc1_weight, + text_proj_fc1_bias, + text_proj_fc2_weight, + text_proj_fc2_bias, + talker_codec_embedding, + cp_codec_embeddings, + tts_pad_embed, + mel: MelSpectrogram::new(), + }) + } + + /// Run the full voice-clone pipeline. + /// + /// `pcm_24k` — reference audio resampled to 24 kHz mono + /// `ref_tokens` — tokenized reference transcript + /// `text_tokens` — tokenized target text + /// `language` — language code (e.g. "en") + pub fn synthesize( + &self, + pcm_24k: &[f32], + ref_tokens: &[u32], + text_tokens: &[u32], + language: &str, + ) -> Result, TtsError> { + let lang_id = tokenizer::language_id(language) + .ok_or_else(|| TtsError::UnsupportedLanguage(language.to_string()))?; + + // 1. Speaker embedding from mel spectrogram + let speaker_embed = self.encode_speaker(pcm_24k)?; + + // 2. Reference codes from tokenizer encoder + let ref_codes = self.encode_ref_codes(pcm_24k)?; + let ref_frames = ref_codes.nrows(); + + // 3. Build ICL prefill embeddings + let prefill_embeds = + self.build_icl_prefill(ref_tokens, text_tokens, lang_id, &speaker_embed, &ref_codes)?; + let prefill_len = prefill_embeds.shape()[1]; + + // 4. Run talker prefill + let (logits, hidden_states, mut past_keys, mut past_values) = + self.run_talker_prefill(&prefill_embeds, prefill_len)?; + + // 5. Autoregressive decode loop + let mut all_codes: Vec<[i64; 16]> = Vec::new(); + let mut talker_past_tokens: Vec = Vec::new(); + let mut current_logits = logits; + + let mut current_hidden = hidden_states + .slice(s![0, prefill_len - 1.., ..]) + .to_owned() + .into_shape_with_order((1, 1, HIDDEN_DIM)) + .map_err(|e| TtsError::Synthesis(format!("reshape hidden: {e}")))?; + + for step in 0..MAX_NEW_TOKENS { + let group0 = sampler::sample( + ¤t_logits, + &TALKER_SAMPLER, + &talker_past_tokens, + |tok| sampler::talker_mask(tok) || (step < 2 && tok == CODEC_EOS as usize), + ) as i64; + + if group0 == CODEC_EOS { + break; + } + talker_past_tokens.push(group0); + + // Code predictor for groups 1-15 + let mut codes = [0i64; 16]; + codes[0] = group0; + self.run_code_predictor(¤t_hidden, &mut codes)?; + all_codes.push(codes); + + // Next talker input: sum all 16 group embeddings + tts_pad + let mut next_embed = self.talker_codec_embedding.row(group0 as usize).to_owned(); + for g in 0..NUM_CP_GROUPS { + let cp_embed = self.cp_codec_embeddings[g].row(codes[g + 1] as usize); + next_embed += &cp_embed; + } + next_embed += &self.tts_pad_embed; + + let next_embed = next_embed + .into_shape_with_order((1, 1, HIDDEN_DIM)) + .map_err(|e| TtsError::Synthesis(format!("reshape next_embed: {e}")))?; + + let total_seq = prefill_len + step + 1; + let position = (prefill_len + step) as i64; + + let (new_logits, new_hidden, new_keys, new_values) = + self.run_talker_decode(&next_embed, total_seq, position, &past_keys, &past_values)?; + + current_logits = new_logits; + current_hidden = new_hidden; + past_keys = new_keys; + past_values = new_values; + } + + if all_codes.is_empty() { + return Err(TtsError::Synthesis("model produced no audio tokens".into())); + } + + // 6. Vocoder: prepend ref codes, decode, trim reference portion + self.run_vocoder_clone(&all_codes, &ref_codes, ref_frames) + } + + // ------------------------------------------------------------------ + // Reference-audio preprocessing + // ------------------------------------------------------------------ + + /// Compute mel spectrogram → run speaker_encoder.onnx → (1024,) speaker embed. + fn encode_speaker(&self, pcm_24k: &[f32]) -> Result, TtsError> { + let mel = self.mel.compute(pcm_24k); // (T_mel, 128) + let n_frames = mel.nrows(); + let mel_3d = mel + .into_shape_with_order((1, n_frames, 128)) + .map_err(|e| TtsError::Synthesis(format!("reshape mel: {e}")))?; + + let t_mel = TensorRef::from_array_view(&mel_3d) + .map_err(|e| TtsError::Synthesis(format!("tensor mel: {e}")))?; + + let mut session = self.speaker_encoder.lock().unwrap(); + let outputs = session + .run(ort::inputs!["mels" => t_mel]) + .map_err(|e| TtsError::Synthesis(format!("speaker encoder failed: {e}")))?; + + let (_, data) = outputs[0] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract speaker embed: {e}")))?; + + Ok(Array1::from(data.to_vec())) + } + + /// Zero-pad audio to canonical 10 s → run tokenizer_encoder.onnx → (T_ref, 16) codes. + fn encode_ref_codes(&self, pcm_24k: &[f32]) -> Result, TtsError> { + let n = pcm_24k.len().min(TOKENIZER_CANONICAL_SAMPLES); + let mut padded = vec![0.0f32; TOKENIZER_CANONICAL_SAMPLES]; + padded[..n].copy_from_slice(&pcm_24k[..n]); + + let waveform = Array2::from_shape_vec((1, TOKENIZER_CANONICAL_SAMPLES), padded) + .map_err(|e| TtsError::Synthesis(format!("reshape waveform: {e}")))?; + + let t_wav = TensorRef::from_array_view(&waveform) + .map_err(|e| TtsError::Synthesis(format!("tensor waveform: {e}")))?; + + let mut session = self.tokenizer_encoder.lock().unwrap(); + let outputs = session + .run(ort::inputs!["waveform" => t_wav]) + .map_err(|e| TtsError::Synthesis(format!("tokenizer encoder failed: {e}")))?; + + // Output shape: (1, 16, 125) + let (_, codes_data) = outputs[0] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract ref codes: {e}")))?; + + let actual_frames = (n as f64 / TOKENIZER_DOWNSAMPLE as f64).ceil() as usize; + + // Reshape (1, 16, 125) → take only actual_frames, transpose to (T_ref, 16) + let full = Array3::from_shape_vec((1, 16, 125), codes_data.to_vec()) + .map_err(|e| TtsError::Synthesis(format!("reshape codes: {e}")))?; + + let trimmed = full.slice(s![0, .., ..actual_frames]).t().to_owned(); // (actual_frames, 16) + + Ok(trimmed) + } + + // ------------------------------------------------------------------ + // ICL prefill construction + // ------------------------------------------------------------------ + + /// Project a text token through the embedding table + SiLU MLP. + fn text_project_token(&self, token: u32) -> Array1 { + let raw = self.text_embedding.row(token as usize).to_owned(); + text_project( + &raw, + &self.text_proj_fc1_weight, + &self.text_proj_fc1_bias, + &self.text_proj_fc2_weight, + &self.text_proj_fc2_bias, + ) + } + + /// Build ICL prefill embeddings for the voice-clone pipeline. + /// + /// Layout: + /// ```text + /// [im_start, assistant, \n] — role prefix (text_proj only) + /// [tts_pad + codec(think)] + /// [tts_pad + codec(think_bos)] — codec think prefix + /// [tts_pad + codec(lang_id)] + /// [tts_pad + codec(think_eos)] + /// [tts_pad + speaker_embed] — speaker slot + /// [tts_bos + codec_pad] — transition + /// ICL block: + /// text side: [text_proj(ref_tokens ++ text_tokens), tts_eos] + codec_pad (T1) + /// codec side: [codec_bos, Σ_g codec_embed_g[ref_code]] + tts_pad (T2) + /// ``` + fn build_icl_prefill( + &self, + ref_tokens: &[u32], + text_tokens: &[u32], + lang_id: i64, + speaker_embed: &Array1, + ref_codes: &Array2, // (T_ref, 16) + ) -> Result, TtsError> { + let codec_pad_embed = self + .talker_codec_embedding + .row(CODEC_PAD as usize) + .to_owned(); + let codec_bos_embed = self + .talker_codec_embedding + .row(CODEC_BOS as usize) + .to_owned(); + let tts_bos_embed = self.text_project_token(TTS_BOS); + let tts_eos_embed = self.text_project_token(TTS_EOS); + + let codec_prefix = [CODEC_THINK, CODEC_THINK_BOS, lang_id, CODEC_THINK_EOS]; + let ref_frames = ref_codes.nrows(); + + // ICL text side: text_proj(ref_tokens ++ text_tokens) | tts_eos + let combined_text_len = ref_tokens.len() + text_tokens.len(); + let t1 = combined_text_len + 1; // +1 for tts_eos + + // ICL codec side: codec_bos | Σ_g codec_embed per ref frame + let t2 = 1 + ref_frames; + + // Total: 3 role + 4 codec_prefix + 1 speaker + 1 transition + T1 + T2 + let seq_len = 3 + codec_prefix.len() + 1 + 1 + t1 + t2; + let mut embeds = Array3::::zeros((1, seq_len, HIDDEN_DIM)); + let mut pos = 0; + + // Part A: Role prefix (3 tokens, text_proj only) + for &tok in &[IM_START, ASSISTANT, NEWLINE] { + let embed = self.text_project_token(tok); + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + + // Part B: Codec think prefix — tts_pad + codec_embed(token) + for &codec_tok in &codec_prefix { + let mut embed = self.tts_pad_embed.clone(); + embed += &self.talker_codec_embedding.row(codec_tok as usize); + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + + // Part C: Speaker slot — tts_pad + speaker_embed + { + let embed = &self.tts_pad_embed + speaker_embed; + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + + // Part D: Transition — tts_bos + codec_pad + { + let embed = &tts_bos_embed + &codec_pad_embed; + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + + // Part E: ICL block + // Text side: text_proj(ref_tokens ++ text_tokens) | tts_eos, all + codec_pad + for &tok in ref_tokens.iter().chain(text_tokens.iter()) { + let embed = self.text_project_token(tok) + &codec_pad_embed; + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + { + let embed = &tts_eos_embed + &codec_pad_embed; + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + + // Codec side: codec_bos + tts_pad, then Σ codec_embed_g[ref_code] + tts_pad per frame + { + let embed = &codec_bos_embed + &self.tts_pad_embed; + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + for f in 0..ref_frames { + // Group 0: talker codec embedding + let mut embed = self + .talker_codec_embedding + .row(ref_codes[[f, 0]] as usize) + .to_owned(); + // Groups 1-15: CP codec embeddings + for g in 0..NUM_CP_GROUPS { + embed += &self.cp_codec_embeddings[g].row(ref_codes[[f, g + 1]] as usize); + } + embed += &self.tts_pad_embed; + embeds.slice_mut(s![0, pos, ..]).assign(&embed); + pos += 1; + } + + debug_assert_eq!(pos, seq_len); + Ok(embeds) + } + + // ------------------------------------------------------------------ + // Talker prefill / decode (same structure as model.rs, 1024-dim) + // ------------------------------------------------------------------ + + fn run_talker_prefill( + &self, + inputs_embeds: &Array3, + seq_len: usize, + ) -> Result { + let attention_mask = Array2::::ones((1, seq_len)); + + let positions: Vec = (0..seq_len as i64).collect(); + let pos_2d = Array1::from(positions) + .into_shape_with_order((1, seq_len)) + .map_err(|e| TtsError::Synthesis(format!("reshape pos: {e}")))?; + let position_ids = ndarray::stack(Axis(0), &[pos_2d.view(), pos_2d.view(), pos_2d.view()]) + .map_err(|e| TtsError::Synthesis(format!("stack pos: {e}")))?; + + let t_embeds = TensorRef::from_array_view(inputs_embeds) + .map_err(|e| TtsError::Synthesis(format!("tensor inputs_embeds: {e}")))?; + let t_mask = TensorRef::from_array_view(&attention_mask) + .map_err(|e| TtsError::Synthesis(format!("tensor mask: {e}")))?; + let t_pos = TensorRef::from_array_view(&position_ids) + .map_err(|e| TtsError::Synthesis(format!("tensor pos: {e}")))?; + + let mut session = self.talker_prefill.lock().unwrap(); + let outputs = session + .run(ort::inputs![ + "inputs_embeds" => t_embeds, + "attention_mask" => t_mask, + "position_ids" => t_pos, + ]) + .map_err(|e| TtsError::Synthesis(format!("talker prefill failed: {e}")))?; + + let (_, logits_data) = outputs[0] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract logits: {e}")))?; + let logits: Vec = logits_data[logits_data.len() - TALKER_VOCAB_SIZE..].to_vec(); + + let (_, hidden_data) = outputs[1] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract hidden: {e}")))?; + let hidden = Array3::from_shape_vec((1, seq_len, HIDDEN_DIM), hidden_data.to_vec()) + .map_err(|e| TtsError::Synthesis(format!("reshape hidden: {e}")))?; + + let mut key_layers = Vec::with_capacity(NUM_LAYERS); + let mut value_layers = Vec::with_capacity(NUM_LAYERS); + for layer in 0..NUM_LAYERS { + let key_idx = 2 + layer * 2; + let val_idx = 2 + layer * 2 + 1; + + let (_, key_data) = outputs[key_idx] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract key layer {layer}: {e}")))?; + let (_, val_data) = outputs[val_idx] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract val layer {layer}: {e}")))?; + + let key_arr = ndarray::ArrayD::from_shape_vec( + vec![1, NUM_KV_HEADS, seq_len, HEAD_DIM], + key_data.to_vec(), + ) + .map_err(|e| TtsError::Synthesis(format!("reshape key {layer}: {e}")))? + .insert_axis(Axis(0)); + let val_arr = ndarray::ArrayD::from_shape_vec( + vec![1, NUM_KV_HEADS, seq_len, HEAD_DIM], + val_data.to_vec(), + ) + .map_err(|e| TtsError::Synthesis(format!("reshape val {layer}: {e}")))? + .insert_axis(Axis(0)); + + key_layers.push(key_arr); + value_layers.push(val_arr); + } + + let past_keys = concatenate( + Axis(0), + &key_layers.iter().map(|a| a.view()).collect::>(), + ) + .map_err(|e| TtsError::Synthesis(format!("stack keys: {e}")))? + .into_shape_with_order((NUM_LAYERS, 1, NUM_KV_HEADS, seq_len, HEAD_DIM)) + .map_err(|e| TtsError::Synthesis(format!("reshape stacked keys: {e}")))?; + + let past_values = concatenate( + Axis(0), + &value_layers.iter().map(|a| a.view()).collect::>(), + ) + .map_err(|e| TtsError::Synthesis(format!("stack values: {e}")))? + .into_shape_with_order((NUM_LAYERS, 1, NUM_KV_HEADS, seq_len, HEAD_DIM)) + .map_err(|e| TtsError::Synthesis(format!("reshape stacked values: {e}")))?; + + Ok((logits, hidden, past_keys, past_values)) + } + + fn run_talker_decode( + &self, + inputs_embeds: &Array3, + total_seq: usize, + position: i64, + past_keys: &Array5, + past_values: &Array5, + ) -> Result { + let attention_mask = Array2::::ones((1, total_seq)); + let position_ids = Array3::::from_elem((3, 1, 1), position); + + let t_embeds = TensorRef::from_array_view(inputs_embeds) + .map_err(|e| TtsError::Synthesis(format!("tensor embeds: {e}")))?; + let t_mask = TensorRef::from_array_view(&attention_mask) + .map_err(|e| TtsError::Synthesis(format!("tensor mask: {e}")))?; + let t_pos = TensorRef::from_array_view(&position_ids) + .map_err(|e| TtsError::Synthesis(format!("tensor pos: {e}")))?; + let t_keys = TensorRef::from_array_view(past_keys) + .map_err(|e| TtsError::Synthesis(format!("tensor past_keys: {e}")))?; + let t_values = TensorRef::from_array_view(past_values) + .map_err(|e| TtsError::Synthesis(format!("tensor past_values: {e}")))?; + + let mut session = self.talker_decode.lock().unwrap(); + let outputs = session + .run(ort::inputs![ + "inputs_embeds" => t_embeds, + "attention_mask" => t_mask, + "position_ids" => t_pos, + "past_keys" => t_keys, + "past_values" => t_values, + ]) + .map_err(|e| TtsError::Synthesis(format!("talker decode failed: {e}")))?; + + let (_, logits_data) = outputs[0] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract decode logits: {e}")))?; + let logits = logits_data.to_vec(); + + let (_, hidden_data) = outputs[1] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract decode hidden: {e}")))?; + let hidden = Array3::from_shape_vec((1, 1, HIDDEN_DIM), hidden_data.to_vec()) + .map_err(|e| TtsError::Synthesis(format!("reshape decode hidden: {e}")))?; + + let (_, keys_data) = outputs[2] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract decode keys: {e}")))?; + let new_keys = Array5::from_shape_vec( + (NUM_LAYERS, 1, NUM_KV_HEADS, total_seq, HEAD_DIM), + keys_data.to_vec(), + ) + .map_err(|e| TtsError::Synthesis(format!("reshape decode keys: {e}")))?; + + let (_, values_data) = outputs[3] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract decode values: {e}")))?; + let new_values = Array5::from_shape_vec( + (NUM_LAYERS, 1, NUM_KV_HEADS, total_seq, HEAD_DIM), + values_data.to_vec(), + ) + .map_err(|e| TtsError::Synthesis(format!("reshape decode values: {e}")))?; + + Ok((logits, hidden, new_keys, new_values)) + } + + // ------------------------------------------------------------------ + // Code predictor (groups 1-15) + // ------------------------------------------------------------------ + + fn run_code_predictor( + &self, + hidden_state: &Array3, + codes: &mut [i64; 16], + ) -> Result<(), TtsError> { + let group0_embed = self + .talker_codec_embedding + .row(codes[0] as usize) + .to_owned() + .into_shape_with_order((1, 1, HIDDEN_DIM)) + .map_err(|e| TtsError::Synthesis(format!("reshape group0 embed: {e}")))?; + + let first_input = concatenate(Axis(1), &[hidden_state.view(), group0_embed.view()]) + .map_err(|e| TtsError::Synthesis(format!("concat cp input: {e}")))?; + + let mut cp_past_keys = + Array5::::zeros((CP_NUM_LAYERS, 1, CP_NUM_KV_HEADS, 0, HEAD_DIM)); + let mut cp_past_values = + Array5::::zeros((CP_NUM_LAYERS, 1, CP_NUM_KV_HEADS, 0, HEAD_DIM)); + let mut cp_input = first_input; + + let mut session = self.code_predictor.lock().unwrap(); + + for group_idx in 0..NUM_CP_GROUPS { + let generation_steps = Array1::::from_elem(1, group_idx as i64); + + let t_input = TensorRef::from_array_view(&cp_input) + .map_err(|e| TtsError::Synthesis(format!("tensor cp input: {e}")))?; + let t_steps = TensorRef::from_array_view(&generation_steps) + .map_err(|e| TtsError::Synthesis(format!("tensor gen steps: {e}")))?; + let t_keys = TensorRef::from_array_view(&cp_past_keys) + .map_err(|e| TtsError::Synthesis(format!("tensor cp keys: {e}")))?; + let t_values = TensorRef::from_array_view(&cp_past_values) + .map_err(|e| TtsError::Synthesis(format!("tensor cp values: {e}")))?; + + let outputs = session + .run(ort::inputs![ + "inputs_embeds" => t_input, + "generation_steps" => t_steps, + "past_keys" => t_keys, + "past_values" => t_values, + ]) + .map_err(|e| { + TtsError::Synthesis(format!("code predictor group {group_idx} failed: {e}")) + })?; + + let (_, logits_data) = outputs[0] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract cp logits: {e}")))?; + let cp_vocab_size = 2048; + let last_logits = &logits_data[logits_data.len() - cp_vocab_size..]; + + let token = sampler::sample(last_logits, &CP_SAMPLER, &[], sampler::no_mask) as i64; + codes[group_idx + 1] = token; + + let seq_so_far = if group_idx == 0 { 2 } else { group_idx + 2 }; + + let (_, keys_data) = outputs[1] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract cp keys: {e}")))?; + let (_, values_data) = outputs[2] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract cp values: {e}")))?; + + cp_past_keys = Array5::from_shape_vec( + (CP_NUM_LAYERS, 1, CP_NUM_KV_HEADS, seq_so_far, HEAD_DIM), + keys_data.to_vec(), + ) + .map_err(|e| TtsError::Synthesis(format!("reshape cp keys: {e}")))?; + cp_past_values = Array5::from_shape_vec( + (CP_NUM_LAYERS, 1, CP_NUM_KV_HEADS, seq_so_far, HEAD_DIM), + values_data.to_vec(), + ) + .map_err(|e| TtsError::Synthesis(format!("reshape cp values: {e}")))?; + + if group_idx < NUM_CP_GROUPS - 1 { + let next_embed = self.cp_codec_embeddings[group_idx] + .row(token as usize) + .to_owned() + .into_shape_with_order((1, 1, HIDDEN_DIM)) + .map_err(|e| TtsError::Synthesis(format!("reshape cp embed: {e}")))?; + cp_input = next_embed; + } + } + + Ok(()) + } + + // ------------------------------------------------------------------ + // Vocoder with reference-code prepend + proportional trim + // ------------------------------------------------------------------ + + fn run_vocoder_clone( + &self, + gen_codes: &[[i64; 16]], + ref_codes: &Array2, // (T_ref, 16) + ref_frames: usize, + ) -> Result, TtsError> { + let gen_frames = gen_codes.len(); + let total_frames = ref_frames + gen_frames; + + // Build (1, 16, total_frames) code tensor: ref_codes | gen_codes + let mut codes = Array3::::zeros((1, 16, total_frames)); + + // Fill ref codes (stored as T_ref × 16) + for f in 0..ref_frames { + for g in 0..16 { + codes[[0, g, f]] = ref_codes[[f, g]]; + } + } + // Fill generated codes + for (t, frame_codes) in gen_codes.iter().enumerate() { + for (g, &code) in frame_codes.iter().enumerate() { + codes[[0, g, ref_frames + t]] = code; + } + } + + let t_codes = TensorRef::from_array_view(&codes) + .map_err(|e| TtsError::Synthesis(format!("tensor codes: {e}")))?; + + let mut session = self.vocoder.lock().unwrap(); + let outputs = session + .run(ort::inputs!["codes" => t_codes]) + .map_err(|e| TtsError::Synthesis(format!("vocoder failed: {e}")))?; + + let (_, waveform) = outputs[0] + .try_extract_tensor::() + .map_err(|e| TtsError::Synthesis(format!("extract waveform: {e}")))?; + + // Trim leading reference portion proportionally + let cut = ref_frames as f64 / total_frames.max(1) as f64 * waveform.len() as f64; + let trimmed = waveform[cut as usize..].to_vec(); + + Ok(AudioFrame::from_vec(trimmed, SAMPLE_RATE)) + } +} diff --git a/crates/wavekat-tts/src/backends/qwen3_tts/download.rs b/crates/wavekat-tts/src/backends/qwen3_tts/download.rs index c426678..28f15dc 100644 --- a/crates/wavekat-tts/src/backends/qwen3_tts/download.rs +++ b/crates/wavekat-tts/src/backends/qwen3_tts/download.rs @@ -10,6 +10,9 @@ use crate::TtsError; const REPO_ID: &str = "wavekat/Qwen3-TTS-1.7B-VoiceDesign-ONNX"; const REVISION: &str = "2026-04-06"; +const CLONE_REPO_ID: &str = "wavekat/Qwen3-TTS-0.6B-Base-ONNX"; +const CLONE_REVISION: &str = "main"; + /// ONNX model files for INT4 precision. const ONNX_FILES_INT4: &[&str] = &[ "int4/talker_prefill.onnx", @@ -158,3 +161,141 @@ pub fn resolve_model_dir(config: &super::ModelConfig) -> Result Result { + if let Some(dir) = &config.model_dir { + return Ok(dir.clone()); + } + + let cache_dir_override = match std::env::var("WAVEKAT_CLONE_MODEL_DIR") { + Ok(dir) => { + let path = PathBuf::from(&dir); + if path.join("config.json").exists() { + return Ok(path); + } + Some(path) + } + Err(_) => None, + }; + + let precision = config.precision; + + let mut builder = ApiBuilder::from_env(); + if let Some(ref dir) = cache_dir_override { + builder = builder.with_cache_dir(dir.clone()); + } + if let Ok(token) = std::env::var("HF_TOKEN") { + if !token.is_empty() { + builder = builder.with_token(Some(token)); + } + } + let api = builder + .build() + .map_err(|e| TtsError::Model(format!("failed to initialize HF Hub client: {e}")))?; + + let repo = api.repo(Repo::with_revision( + CLONE_REPO_ID.to_string(), + RepoType::Model, + CLONE_REVISION.to_string(), + )); + + let onnx_files = match precision { + super::ModelPrecision::Int4 => CLONE_ONNX_FILES_INT4, + super::ModelPrecision::Fp32 => CLONE_ONNX_FILES_FP32, + }; + let total = 1 + onnx_files.len() + CLONE_SHARED_FILES[1..].len(); + + eprintln!( + "Ensuring Qwen3-TTS 0.6B Clone ({}) model ({total} files from {CLONE_REPO_ID})...", + precision.subdir() + ); + + eprintln!("[1/{total}] {}", CLONE_SHARED_FILES[0]); + let config_path = repo.get(CLONE_SHARED_FILES[0]).map_err(|e| { + TtsError::Model(format!("failed to download {}: {e}", CLONE_SHARED_FILES[0])) + })?; + + let model_dir = config_path + .parent() + .ok_or_else(|| TtsError::Model("unexpected cache path for config.json".into()))? + .to_path_buf(); + + for (i, filename) in onnx_files + .iter() + .chain(CLONE_SHARED_FILES[1..].iter()) + .enumerate() + { + eprintln!("[{}/{total}] {filename}", i + 2); + repo.get(filename) + .map_err(|e| TtsError::Model(format!("failed to download {filename}: {e}")))?; + } + + eprintln!("Files ready. Loading clone model ..."); + Ok(model_dir) +} diff --git a/crates/wavekat-tts/src/backends/qwen3_tts/mel.rs b/crates/wavekat-tts/src/backends/qwen3_tts/mel.rs new file mode 100644 index 0000000..33f1c47 --- /dev/null +++ b/crates/wavekat-tts/src/backends/qwen3_tts/mel.rs @@ -0,0 +1,193 @@ +//! Mel-spectrogram computation matching the Qwen3-TTS reference (librosa). +//! +//! Parameters: sr=24000, n_fft=1024, hop=256, win=1024, n_mels=128, +//! fmin=0, fmax=12000, center=False, power=2.0, log on top. + +use ndarray::Array2; +use realfft::num_complex::Complex; +use realfft::RealFftPlanner; + +const SR: f32 = 24000.0; +const N_FFT: usize = 1024; +const HOP: usize = 256; +const WIN: usize = 1024; +const N_MELS: usize = 128; +const FMIN: f32 = 0.0; +const FMAX: f32 = 12000.0; +const N_BINS: usize = N_FFT / 2 + 1; // 513 + +/// Pre-computed mel filterbank and Hann window for repeated use. +pub struct MelSpectrogram { + window: Vec, + filterbank: Array2, // (N_MELS, N_BINS) +} + +impl MelSpectrogram { + pub fn new() -> Self { + Self { + window: hann_window(WIN), + filterbank: mel_filterbank(N_MELS, N_FFT, SR, FMIN, FMAX), + } + } + + /// Compute log-mel spectrogram. Returns `(T_mel, 128)` f32. + pub fn compute(&self, audio: &[f32]) -> Array2 { + let n_frames = if audio.len() >= WIN { + 1 + (audio.len() - WIN) / HOP + } else { + 0 + }; + + let mut planner = RealFftPlanner::::new(); + let fft = planner.plan_fft_forward(N_FFT); + + let mut mel = Array2::::zeros((n_frames, N_MELS)); + let mut frame_buf = vec![0.0f32; N_FFT]; + let mut spectrum = vec![Complex::new(0.0f32, 0.0f32); N_BINS]; + + for i in 0..n_frames { + let start = i * HOP; + + // Window the frame + for j in 0..WIN { + frame_buf[j] = audio[start + j] * self.window[j]; + } + // FFT + fft.process(&mut frame_buf, &mut spectrum).unwrap(); + + // Power spectrum → mel filterbank → log + for m in 0..N_MELS { + let mut sum = 0.0f32; + for (k, s) in spectrum.iter().enumerate() { + let power = s.re * s.re + s.im * s.im; + sum += self.filterbank[[m, k]] * power; + } + mel[[i, m]] = sum.max(1e-5).ln(); + } + } + + mel + } +} + +/// Periodic Hann window (matches librosa's STFT default). +fn hann_window(length: usize) -> Vec { + let n = length as f32; + (0..length) + .map(|i| 0.5 * (1.0 - (2.0 * std::f32::consts::PI * i as f32 / n).cos())) + .collect() +} + +// --------------------------------------------------------------------------- +// Slaney mel scale (matches librosa default, htk=False) +// --------------------------------------------------------------------------- + +const F_SP: f32 = 200.0 / 3.0; // 66.667 Hz per mel (linear region) +const MIN_LOG_HZ: f32 = 1000.0; +const MIN_LOG_MEL: f32 = MIN_LOG_HZ / F_SP; // 15.0 + +/// log(6.4) / 27 ≈ 0.06875 +fn logstep() -> f32 { + 6.4f32.ln() / 27.0 +} + +fn hz_to_mel(freq: f32) -> f32 { + if freq < MIN_LOG_HZ { + freq / F_SP + } else { + MIN_LOG_MEL + (freq / MIN_LOG_HZ).ln() / logstep() + } +} + +fn mel_to_hz(mel: f32) -> f32 { + if mel < MIN_LOG_MEL { + mel * F_SP + } else { + MIN_LOG_HZ * (logstep() * (mel - MIN_LOG_MEL)).exp() + } +} + +/// Build a `(n_mels, n_bins)` triangular mel filterbank (no normalization). +fn mel_filterbank(n_mels: usize, n_fft: usize, sr: f32, fmin: f32, fmax: f32) -> Array2 { + let n_bins = n_fft / 2 + 1; + let min_mel = hz_to_mel(fmin); + let max_mel = hz_to_mel(fmax); + + // n_mels + 2 mel-spaced center frequencies + let mel_points: Vec = (0..n_mels + 2) + .map(|i| mel_to_hz(min_mel + (max_mel - min_mel) * i as f32 / (n_mels + 1) as f32)) + .collect(); + + // FFT bin center frequencies + let fft_freqs: Vec = (0..n_bins).map(|k| k as f32 * sr / n_fft as f32).collect(); + + let mut fb = Array2::::zeros((n_mels, n_bins)); + + for m in 0..n_mels { + let f_left = mel_points[m]; + let f_center = mel_points[m + 1]; + let f_right = mel_points[m + 2]; + + let d_left = f_center - f_left; + let d_right = f_right - f_center; + + for k in 0..n_bins { + let f = fft_freqs[k]; + if f >= f_left && f <= f_center && d_left > 0.0 { + fb[[m, k]] = (f - f_left) / d_left; + } else if f > f_center && f <= f_right && d_right > 0.0 { + fb[[m, k]] = (f_right - f) / d_right; + } + } + } + + fb +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn mel_scale_roundtrip() { + for &freq in &[0.0, 500.0, 1000.0, 4000.0, 12000.0] { + let m = hz_to_mel(freq); + let f = mel_to_hz(m); + assert!( + (f - freq).abs() < 0.01, + "roundtrip failed for {freq}: got {f}" + ); + } + } + + #[test] + fn filterbank_shape() { + let fb = mel_filterbank(128, 1024, 24000.0, 0.0, 12000.0); + assert_eq!(fb.shape(), &[128, 513]); + } + + #[test] + fn filterbank_non_negative() { + let fb = mel_filterbank(128, 1024, 24000.0, 0.0, 12000.0); + assert!(fb.iter().all(|&v| v >= 0.0)); + } + + #[test] + fn mel_output_shape() { + let mel = MelSpectrogram::new(); + // 1 second of audio at 24kHz + let audio = vec![0.0f32; 24000]; + let result = mel.compute(&audio); + // (24000 - 1024) / 256 + 1 = 89 + 1 = 90 frames + let expected_frames = 1 + (24000 - WIN) / HOP; + assert_eq!(result.shape(), &[expected_frames, 128]); + } + + #[test] + fn hann_window_properties() { + let w = hann_window(1024); + assert_eq!(w.len(), 1024); + assert!((w[0] - 0.0).abs() < 1e-6); // starts near zero + assert!(w[512] > 0.99); // peak near middle + } +} diff --git a/crates/wavekat-tts/src/backends/qwen3_tts/mod.rs b/crates/wavekat-tts/src/backends/qwen3_tts/mod.rs index 0f60c65..df1d262 100644 --- a/crates/wavekat-tts/src/backends/qwen3_tts/mod.rs +++ b/crates/wavekat-tts/src/backends/qwen3_tts/mod.rs @@ -1,27 +1,30 @@ -//! Qwen3-TTS backend (ONNX, 1.7B VoiceDesign). +//! Qwen3-TTS backends (ONNX). //! -//! Runs the Qwen3-TTS-12Hz-1.7B-VoiceDesign model via ONNX Runtime. -//! Supports INT4 (weight-only quantized, default) and FP32 precision. +//! Two sibling structs: +//! +//! - [`Qwen3Tts`] — 1.7B VoiceDesign (prompt-based voice styling) +//! - [`Qwen3TtsClone`] — 0.6B Base (reference-audio voice cloning) +//! +//! # VoiceDesign (1.7B) //! //! ```ignore //! use wavekat_tts::{TtsBackend, SynthesizeRequest}; //! use wavekat_tts::backends::qwen3_tts::{Qwen3Tts, ModelConfig, ModelPrecision, ExecutionProvider}; //! -//! // Auto-download INT4 model files via HF Hub, run on CPU (default): //! let tts = Qwen3Tts::new()?; +//! let request = SynthesizeRequest::new("Hello, world"); +//! let audio = tts.synthesize(&request)?; +//! ``` //! -//! // Auto-download FP32, run on CPU: -//! let tts = Qwen3Tts::from_config(ModelConfig::default().with_precision(ModelPrecision::Fp32))?; +//! # Voice Clone (0.6B) //! -//! // INT4 from a local directory, run on CUDA: -//! let tts = Qwen3Tts::from_config( -//! ModelConfig::default() -//! .with_dir("models/qwen3-tts-1.7b") -//! .with_execution_provider(ExecutionProvider::Cuda), -//! )?; +//! ```ignore +//! use wavekat_tts::backends::qwen3_tts::{Qwen3TtsClone, CloneRequest, ModelConfig}; //! -//! let request = SynthesizeRequest::new("Hello, world"); -//! let audio = tts.synthesize(&request)?; +//! let tts = Qwen3TtsClone::new()?; +//! let ref_audio: Vec = todo!("24 kHz mono PCM"); +//! let req = CloneRequest::new("Text to say", &ref_audio, 24000, "Transcript of ref audio"); +//! let audio = tts.synthesize_clone(&req)?; //! ``` use std::path::PathBuf; @@ -38,7 +41,9 @@ use tokenizer::{IM_END, IM_START, NEWLINE}; static WARNED_NO_INSTRUCTION: Once = Once::new(); +mod clone_model; mod download; +mod mel; mod model; mod sampler; mod tokenizer; @@ -87,7 +92,7 @@ pub enum ExecutionProvider { CoreMl, } -/// Model loading configuration for [`Qwen3Tts`]. +/// Model loading configuration for [`Qwen3Tts`] and [`Qwen3TtsClone`]. /// /// All fields default to sensible values: INT4 quantization, CPU inference, /// and auto-download from HF Hub. @@ -142,7 +147,28 @@ impl ModelConfig { } } -/// Qwen3-TTS backend using ONNX Runtime. +/// Qwen3-TTS 1.7B VoiceDesign backend using ONNX Runtime. +/// +/// Generates speech from text using a style instruction to control voice +/// characteristics (tone, pace, emotion). Implements [`TtsBackend`]. +/// +/// # Examples +/// +/// ```rust,no_run +/// use wavekat_tts::{TtsBackend, SynthesizeRequest}; +/// use wavekat_tts::backends::qwen3_tts::Qwen3Tts; +/// +/// let tts = Qwen3Tts::new()?; +/// let audio = tts.synthesize( +/// &SynthesizeRequest::new("Hello, world") +/// .with_instruction("Speak naturally and clearly."), +/// )?; +/// audio.write_wav("output.wav")?; +/// # Ok::<(), wavekat_tts::TtsError>(()) +/// ``` +/// +/// Use [`Qwen3Tts::from_config`] with [`ModelConfig`] to select FP32 +/// precision or a GPU execution provider. pub struct Qwen3Tts { model: model::Model, tokenizer: tokenizer::Tokenizer, @@ -172,6 +198,135 @@ impl Qwen3Tts { } } +// --------------------------------------------------------------------------- +// Voice Clone (0.6B Base) +// --------------------------------------------------------------------------- + +/// A voice-clone synthesis request. +/// +/// Requires a reference audio clip (3–10 s, mono) and its transcript. The +/// model produces speech in the cloned voice speaking `text`. +/// +/// Reference audio **must be 24 kHz mono float32 PCM**. If your audio is at +/// a different sample rate, resample before passing it in. +/// +/// # Examples +/// +/// ```rust,no_run +/// use wavekat_tts::backends::qwen3_tts::CloneRequest; +/// +/// let ref_samples: Vec = vec![]; // 24 kHz mono float32 +/// let req = CloneRequest::new("Hello", &ref_samples, 24000, "ref transcript") +/// .with_language("en"); +/// ``` +#[derive(Debug, Clone)] +pub struct CloneRequest<'a> { + /// Text to synthesize in the cloned voice. + pub text: &'a str, + /// Reference audio samples (24 kHz mono float32). + pub ref_samples: &'a [f32], + /// Sample rate of `ref_samples` (must be 24000). + pub ref_sample_rate: u32, + /// Transcript of the reference audio (required for ICL mode). + pub ref_text: &'a str, + /// Language code (e.g. `"en"`, `"zh"`). `None` defaults to `"en"`. + pub language: Option<&'a str>, +} + +impl<'a> CloneRequest<'a> { + /// Create a clone request with all required fields. + pub fn new( + text: &'a str, + ref_samples: &'a [f32], + ref_sample_rate: u32, + ref_text: &'a str, + ) -> Self { + Self { + text, + ref_samples, + ref_sample_rate, + ref_text, + language: None, + } + } + + /// Set the language code. + pub fn with_language(mut self, language: &'a str) -> Self { + self.language = Some(language); + self + } +} + +/// Qwen3-TTS 0.6B Base voice-clone backend using ONNX Runtime. +/// +/// Clones a speaker's voice from a short reference clip (3–10 s) and its +/// transcript, then synthesizes new text in that voice. +/// +/// # Examples +/// +/// ```rust,no_run +/// use wavekat_tts::AudioFrame; +/// use wavekat_tts::backends::qwen3_tts::{Qwen3TtsClone, CloneRequest}; +/// +/// let ref_audio = AudioFrame::from_wav("ref.wav")?; +/// let tts = Qwen3TtsClone::new()?; +/// let req = CloneRequest::new( +/// "Text to say in the cloned voice", +/// ref_audio.samples(), +/// 24000, +/// "Transcript of the reference clip.", +/// ); +/// let audio = tts.synthesize_clone(&req)?; +/// audio.write_wav("clone_output.wav")?; +/// # Ok::<(), wavekat_tts::TtsError>(()) +/// ``` +/// +/// Use [`Qwen3TtsClone::from_config`] with [`ModelConfig`] to select FP32 +/// precision or a GPU execution provider. +pub struct Qwen3TtsClone { + model: clone_model::CloneModel, + tokenizer: tokenizer::Tokenizer, +} + +impl Qwen3TtsClone { + /// Create a new clone backend with default config (INT4, CPU, auto-download). + pub fn new() -> Result { + Self::from_config(ModelConfig::default()) + } + + /// Create a new clone backend with the given [`ModelConfig`]. + pub fn from_config(config: ModelConfig) -> Result { + let model_dir = download::resolve_clone_model_dir(&config)?; + let model = clone_model::CloneModel::load(model_dir.as_ref(), &config)?; + let tokenizer = tokenizer::Tokenizer::new(&model_dir)?; + Ok(Self { model, tokenizer }) + } + + /// Synthesize text in a cloned voice. + pub fn synthesize_clone( + &self, + request: &CloneRequest, + ) -> Result, TtsError> { + if request.ref_sample_rate != 24000 { + return Err(TtsError::Synthesis(format!( + "reference audio must be 24 kHz, got {} Hz", + request.ref_sample_rate, + ))); + } + + let language = request.language.unwrap_or("en"); + let ref_tokens = self.tokenizer.encode(request.ref_text)?; + let text_tokens = self.tokenizer.encode(request.text)?; + + self.model + .synthesize(request.ref_samples, &ref_tokens, &text_tokens, language) + } +} + +// --------------------------------------------------------------------------- +// TtsBackend for Qwen3Tts (1.7B VoiceDesign) +// --------------------------------------------------------------------------- + impl TtsBackend for Qwen3Tts { fn synthesize(&self, request: &SynthesizeRequest) -> Result, TtsError> { let tokens = self.tokenizer.encode(request.text)?; diff --git a/crates/wavekat-tts/src/backends/qwen3_tts/model.rs b/crates/wavekat-tts/src/backends/qwen3_tts/model.rs index 41d064c..77d532a 100644 --- a/crates/wavekat-tts/src/backends/qwen3_tts/model.rs +++ b/crates/wavekat-tts/src/backends/qwen3_tts/model.rs @@ -648,7 +648,7 @@ impl Model { /// a full copy only on cross-device mounts. /// /// Returns `onnx_dir` unchanged if no symlinks are present. -fn prepare_onnx_dir(onnx_dir: &Path) -> Result { +pub(super) fn prepare_onnx_dir(onnx_dir: &Path) -> Result { let entries: Vec<_> = std::fs::read_dir(onnx_dir) .map_err(|e| TtsError::Model(format!("cannot read {}: {e}", onnx_dir.display())))? .filter_map(|e| e.ok()) @@ -665,6 +665,10 @@ fn prepare_onnx_dir(onnx_dir: &Path) -> Result { for entry in &entries { let src = entry.path(); + // Skip directories — we only need to resolve file symlinks. + if src.is_dir() { + continue; + } let dst = resolved.join(entry.file_name()); if dst.exists() { continue; @@ -691,7 +695,7 @@ fn prepare_onnx_dir(onnx_dir: &Path) -> Result { /// /// CPU is the ORT default — no registration needed. CUDA and CoreML require /// an ORT build that includes those providers; otherwise ORT will return an error. -fn apply_execution_provider( +pub(super) fn apply_execution_provider( builder: ort::session::builder::SessionBuilder, ep: super::ExecutionProvider, ) -> Result { @@ -716,8 +720,8 @@ fn apply_execution_provider( } } -/// SiLU-gated MLP text projection: 2048 → 2048. -fn text_project( +/// SiLU-gated MLP text projection. +pub(super) fn text_project( input: &Array1, fc1_weight: &Array2, fc1_bias: &Array1, @@ -730,7 +734,7 @@ fn text_project( } /// Load a 2D .npy file into Array2. -fn load_npy2(dir: &Path, name: &str) -> Result, TtsError> { +pub(super) fn load_npy2(dir: &Path, name: &str) -> Result, TtsError> { let path = dir.join(name); let bytes = std::fs::read(&path) .map_err(|e| TtsError::Model(format!("failed to read {}: {e}", path.display())))?; @@ -745,7 +749,7 @@ fn load_npy2(dir: &Path, name: &str) -> Result, TtsError> { } /// Load a 1D .npy file into Array1. -fn load_npy1(dir: &Path, name: &str) -> Result, TtsError> { +pub(super) fn load_npy1(dir: &Path, name: &str) -> Result, TtsError> { let path = dir.join(name); let bytes = std::fs::read(&path) .map_err(|e| TtsError::Model(format!("failed to read {}: {e}", path.display())))?; diff --git a/crates/wavekat-tts/src/error.rs b/crates/wavekat-tts/src/error.rs index 5392b0c..0997f10 100644 --- a/crates/wavekat-tts/src/error.rs +++ b/crates/wavekat-tts/src/error.rs @@ -25,3 +25,12 @@ pub enum TtsError { #[error("io error: {0}")] Io(#[from] std::io::Error), } + +impl From for TtsError { + fn from(err: wavekat_core::CoreError) -> Self { + match err { + wavekat_core::CoreError::Io(io) => Self::Io(io), + wavekat_core::CoreError::Audio(msg) => Self::Synthesis(msg), + } + } +} diff --git a/crates/wavekat-tts/src/types.rs b/crates/wavekat-tts/src/types.rs index 847d16f..3ed1322 100644 --- a/crates/wavekat-tts/src/types.rs +++ b/crates/wavekat-tts/src/types.rs @@ -100,7 +100,10 @@ pub struct VoiceInfo { /// Voice gender hint. #[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub enum Gender { + /// Male voice. Male, + /// Female voice. Female, + /// Gender-neutral or unspecified voice. Neutral, } diff --git a/docs/08-qwen3-tts-0.6b-voice-clone.md b/docs/08-qwen3-tts-0.6b-voice-clone.md new file mode 100644 index 0000000..a0dd09f --- /dev/null +++ b/docs/08-qwen3-tts-0.6b-voice-clone.md @@ -0,0 +1,325 @@ +# Qwen3-TTS 0.6B Voice Clone + +> **Status: Implemented** — all planned work is complete. This doc describes the +> design and what was built. + +## Goal + +A second variant of the Qwen3-TTS backend that performs **voice cloning** from a +reference audio clip (3–10 s) plus its transcript, using the +[`Qwen/Qwen3-TTS-12Hz-0.6B-Base`][hf-base] checkpoint. + +The existing 1.7B VoiceDesign backend is untouched. The two coexist as sibling +structs: `Qwen3Tts` (1.7B VoiceDesign) and `Qwen3TtsClone` (0.6B Base). + +The user provides: + +- `text` — what to say +- `ref_samples` — short waveform (24 kHz mono float32 PCM) +- `ref_text` — transcript of the reference audio (required for ICL mode) +- `language` — target language (optional, defaults to `"en"`) + +The model returns audio in the cloned voice. This first pass covers the **0.6B** +size; the 1.7B Base variant uses the same architecture and can be added later by +swapping model files. + +## Why a separate variant + +Voice clone is **not** just "VoiceDesign with a different prompt". The Base +checkpoint differs from VoiceDesign in three ways that matter for inference: + +| | VoiceDesign (1.7B, current) | Base (0.6B, this work) | +|---|---|---| +| Talker hidden | 2048 | 1024 | +| Talker layers | 28 | 28 | +| Talker KV heads | 8 | 8 | +| Speaker encoder | none | ECAPA-TDNN, mel→1024-dim embed | +| Speech tokenizer encoder | unused at inference | **used** to encode ref_audio | +| Prefill prefix | text-only instruction tokens | speaker_embed + ref_text + ref_codes interleaving (ICL) | + +So we need (a) one or more new ONNX assets, (b) new prefill construction logic +in Rust, and (c) a new request entry point (`synthesize_clone`) that takes +reference audio. + +## Architecture: how voice clone works + +Source of truth: [`generate_voice_clone`][gen_clone] and +[`generate_icl_prompt`][gen_icl] in `qwen_tts/inference/qwen3_tts_model.py` +and `qwen_tts/core/models/modeling_qwen3_tts.py`. + +### Reference-audio preprocessing (host code, runs once per ref clip) + +``` +ref_audio (any sr) + │ + ├──► resample to 24 kHz ──► mel-spectrogram (n_fft=1024, hop=256, win=1024, + │ n_mels=128, fmin=0, fmax=12000) + │ │ + │ ▼ + │ [Speaker Encoder] ECAPA-TDNN + │ │ + │ ▼ + │ ref_spk_embedding (1024,) + │ + └──► [Speech Tokenizer Encoder] Mimi-based (12 Hz, 16 quantizer groups) + │ + ▼ + ref_code (T_ref, 16) i64 codebook indices +``` + +ICL mode uses **both** outputs. `x_vector_only_mode` uses only the speaker +embedding (lower quality, no transcript needed). We implement ICL first. + +### Prefill embedding construction (Base + ICL) + +``` +talker_input_embed = concat over seq dim: + + [ im_start, assistant, \n ] text_proj only (3) + [ tts_pad + codec_embed[think] ] + [ tts_pad + codec_embed[think_bos] ] + [ tts_pad + codec_embed[lang_id] ] codec think prefix (4) + [ tts_pad + codec_embed[think_eos] ] + [ speaker_embed + codec_embed[codec_pad] ] speaker slot (1) + [ tts_bos + codec_embed[codec_pad] ] transition (1) + + ICL block (non_streaming, len = max(text_lens, codec_lens)): + text_part = [ text_proj(ref_id ++ text_id), tts_eos ] (T1) + codec_part = [ codec_embed[codec_bos], Σ_g codec_embed_g[ref_code[:,g]] ] (T2 = 1 + T_ref) + + text_part += codec_embed[codec_pad] (broadcast over T1) + codec_part += tts_pad (broadcast over T2) + + icl_input = concat([text_part, codec_part], dim=1) +``` + +Two crucial differences from VoiceDesign: + +1. The codec prefix carries an **extra speaker slot** (`speaker_embed + + codec_pad`) inserted after `think_eos` and before the transition. This is + how the model receives the cloned-voice condition. +2. The text path (`text_proj(ref_id ++ text_id) + tts_eos`) is **summed + element-wise** with the codec path (`codec_bos + Σ ref_code_embeds`) along + the sequence dim, using `codec_pad` / `tts_pad` to pad whichever is shorter. + This is the in-context-learning mechanism — the ref text and ref codes are + delivered together as a single positional stream the model has been trained + to "continue". + +### Decode loop + +After prefill, the autoregressive decode loop is **identical** to VoiceDesign: +sample group-0 from talker logits, run code predictor for groups 1–15, sum 16 +codec embeddings + `trailing_text_hidden` (= `tts_pad` in non-streaming) for the +next step, stop on `codec_eos`. + +### Vocoder / output trim + +The vocoder is run on `concat([ref_code, generated_codes])` so the model's +generated tail blends smoothly with the reference timbre. Then we **cut off** +the leading portion proportional to `ref_len / total_len` so the returned +waveform contains only the new content. (See lines 612–631 in +`qwen3_tts_model.py`.) + +## ONNX assets + +The 0.6B Base checkpoint provides a different talker plus extra modules (speaker +encoder, tokenizer encoder). All are exported via `tools/qwen3-tts-onnx/`. + +### Reused (re-exported with smaller dims) + +| File | Notes | +|---|---| +| `talker_prefill.onnx` | hidden=1024, vocab=3072, 28 layers, 8 KV heads | +| `talker_decode.onnx` | same | +| `code_predictor.onnx` | 1024 hidden, 5 layers | +| `vocoder.onnx` | Mimi v2 decoder | +| `embeddings/text_embedding.npy` etc. | re-extracted from Base weights | +| `embeddings/cp_codec_embedding_0..14.npy` | re-extracted | + +### New for voice clone + +| File | Source | Shape | Purpose | +|---|---|---|---| +| `speaker_encoder.onnx` | `model.speaker_encoder` (ECAPA-TDNN) | in: `(1, T_mel, 128)` mel; out: `(1, 1024)` | Encode ref audio → speaker embed | +| `tokenizer_encoder.onnx` | `model.speech_tokenizer.encode()` | in: `(1, 1, S_audio)` waveform; out: `(1, T_codes, 16)` i64 | Encode ref audio → ref codes | + +Mel computation is done in host code (pure-Rust STFT + mel filterbank in +`mel.rs`) rather than a separate ONNX. Params match the reference exactly: +`n_fft=1024, hop=256, win=1024, n_mels=128, fmin=0, fmax=12000`, `center=False`, +log on top. + +### Export tooling + +`tools/qwen3-tts-onnx/` contains: + +- `export_speaker_encoder.py` — ECAPA-TDNN, dynamic axis on mel time dim, + opset 17. Validates against PyTorch (atol=1e-4). +- `export_tokenizer_encoder.py` — Mimi-based tokenizer encoder. Uses JIT trace + with fixed size (240k samples = 10 s @ 24 kHz). Applies `mask_patch.py` to + handle Mimi's causal mask incompatibility with tracing. Validates exact i64 + code match. +- `mask_patch.py` — patches Mimi causal mask for JIT tracing compatibility. +- `generate_clone_onnx.py` — end-to-end Python ONNX voice clone reference + (547 lines). Loads all 6 ONNX sessions, builds ICL prefill, runs decode loop, + vocoders, and trims the reference portion. +- Existing `export_talker.py`, `export_code_predictor.py`, `export_vocoder.py`, + `export_embeddings.py` are reused with `MODEL_ID=Qwen/Qwen3-TTS-12Hz-0.6B-Base`. +- `quantize_int4.py` quantizes talker/CP/vocoder; speaker and tokenizer encoders + stay FP32 (small, conditioning path — no upside to quantizing). + +### Makefile targets + +``` +make clone-all # full export + quantize + HF packaging +make clone-export # orchestrate 6 component exports +make clone-base-preset # INT4 quantization (encoders stay FP32) +make clone-hf # package for HF Hub +make clone-fixture # regenerate ref WAV fixture +make clone-generate # test FP32 output +``` + +### Published HF repo + +[`wavekat/Qwen3-TTS-0.6B-Base-ONNX`](https://huggingface.co/wavekat/Qwen3-TTS-0.6B-Base-ONNX) +— separate from the existing 1.7B repo. Layout: + +``` +speaker_encoder.onnx FP32 only +tokenizer_encoder.onnx FP32 only +fp32/ talker_prefill, talker_decode, code_predictor, vocoder +int4/ same (INT4 quantized) +embeddings/ text_embedding, text_projection, codec embeddings +tokenizer/ vocab.json, merges.txt +config.json +``` + +## Rust backend + +### Module structure + +``` +crates/wavekat-tts/src/backends/qwen3_tts/ +├── mod.rs — Qwen3Tts (1.7B) + Qwen3TtsClone (0.6B) + CloneRequest +├── download.rs — resolve_model_dir + resolve_clone_model_dir +├── model.rs — existing talker/CP/vocoder pipeline (untouched) +├── clone_model.rs — CloneModel: 6 ONNX sessions + ICL prefill builder +├── mel.rs — pure-Rust STFT + mel filterbank (realfft) +├── tokenizer.rs — shared text tokenization (unchanged) +└── sampler.rs — shared sampling logic (unchanged) +``` + +### Public API + +```rust +use wavekat_tts::backends::qwen3_tts::{Qwen3TtsClone, CloneRequest, ModelConfig}; + +let tts = Qwen3TtsClone::new()?; // 0.6B Base, INT4, CPU +let req = CloneRequest::new("Text to say", &pcm_24k, 24000, "transcript of ref") + .with_language("en"); +let frame: AudioFrame = tts.synthesize_clone(&req)?; +``` + +`Qwen3TtsClone` exposes `fn synthesize_clone(&self, req: &CloneRequest) -> +Result, TtsError>`. The existing `TtsBackend` contract +doesn't fit (no place for a reference clip), so `Qwen3TtsClone` has its own +method rather than implementing `TtsBackend`. + +`CloneRequest`: +```rust +pub struct CloneRequest<'a> { + pub text: &'a str, + pub ref_samples: &'a [f32], // 24 kHz mono float32 PCM + pub ref_sample_rate: u32, // must be 24000 + pub ref_text: &'a str, // required for ICL mode + pub language: Option<&'a str>, // defaults to "en" +} +``` + +`ModelConfig` is shared between both backends (precision, EP, model dir). + +### Implementation + +`clone_model.rs` (`CloneModel`) handles the full pipeline: + +1. **`encode_speaker()`** — computes mel spectrogram via `mel.rs`, runs + `speaker_encoder.onnx` → `(1, 1024)` speaker embedding. +2. **`encode_ref_codes()`** — runs `tokenizer_encoder.onnx` on raw PCM → + `(T_ref, 16)` i64 codebook indices. +3. **`build_icl_prefill()`** — constructs the prefill embedding tensor: + role prefix → codec prefix → speaker slot → transition → ICL block + (text + codec interleaving). Produces `(1, T, 1024)`. +4. **`run_talker_prefill()` / `run_talker_decode()`** — talker pipeline + (same structure as 1.7B, adapted to 1024-dim hidden). +5. **`run_code_predictor()`** — predicts codec groups 1–15 from group 0. +6. **`run_vocoder_clone()`** — prepends ref codes before vocoding, then trims + the leading portion proportional to `ref_len / total_len`. + +### Mel-spectrogram in Rust + +`mel.rs` — pure-Rust implementation using `realfft`: +- `MelSpectrogram` struct with precomputed Hann window + Slaney mel filterbank +- `.compute(audio) → (T, 128)` log-mel frames +- Params match the reference: `n_fft=1024, hop=256, win=1024, n_mels=128, + fmin=0, fmax=12000, center=False` +- Unit tests for scale roundtrip, filterbank shape, output shape + +### Audio I/O and resampling + +`CloneRequest` takes raw `&[f32]` PCM samples. Reading WAV files is the +caller's responsibility (example uses `hound`). Sample rate must be 24 kHz — +validated at runtime with a clear error. Resampling is not done internally; +callers resample before calling. + +## Example + +`examples/synthesize_clone.rs` — full CLI example (165 lines) with `--ref-audio`, +`--ref-text`, `--text`, `--language`, `--precision`, `--provider` flags. Reads a +24 kHz mono WAV, runs `Qwen3TtsClone::synthesize_clone`, writes output, and +prints RTF. + +## CI + +`.github/workflows/export-onnx.yml` — unified workflow with a variant selector +dropdown (`voicedesign` | `clone`). The `clone` variant runs `clone-export`, +`clone-base-preset` (INT4 quantization, encoders FP32), and `clone-hf` +(HF Hub packaging). Conditional validation and cleanup between steps for +runner disk space. + +## Verification + +1. **Per-component parity** (in export scripts): + - speaker_encoder ONNX vs PyTorch on a fixed mel: max abs err < 1e-4 + - tokenizer_encoder ONNX vs PyTorch on a fixed wav: identical i64 codes +2. **End-to-end ONNX (Python)** via `generate_clone_onnx.py`: + - Loads all 6 ONNX sessions, builds ICL prefill, runs decode + vocoder + - Supports FP32 and INT4 variants, all 10 languages +3. **Rust vs Python parity**: + - Same ICL prefill layout, same vocoder trim logic +4. **Reference fixture**: `tools/qwen3-tts-onnx/fixtures/ref_clone.wav` for + integration tests and examples. + +## Decisions made + +1. **API shape**: sibling struct `Qwen3TtsClone` + `CloneRequest`. Existing + `Qwen3Tts` and `TtsBackend` unchanged. +2. **Speaker / tokenizer encoder precision**: FP32 only. They're small and + sit on the conditioning path — no upside to quantizing. +3. **HF repo**: separate `wavekat/Qwen3-TTS-0.6B-Base-ONNX` (does not reuse + the 1.7B repo). +4. **Mel in host code**: pure-Rust (`realfft`) rather than a separate ONNX + model. Deterministic, small, and avoids an extra session. +5. **No internal resampling**: callers must provide 24 kHz audio. Keeps the + backend simple and avoids pulling in `rubato`. +6. **Tokenizer encoder uses fixed-size JIT trace** (240k samples) rather than + dynamic axes, due to Mimi causal mask incompatibility with tracing. + +## Future work + +- 1.7B Base voice clone (same architecture; swap model dir) +- `x_vector_only_mode` (no `ref_text`) — speaker-embedding-only conditioning +- Streaming voice clone +- True batch inference + +[hf-base]: https://huggingface.co/Qwen/Qwen3-TTS-12Hz-0.6B-Base +[gen_clone]: https://github.com/QwenLM/Qwen3-TTS/blob/main/qwen_tts/inference/qwen3_tts_model.py +[gen_icl]: https://github.com/QwenLM/Qwen3-TTS/blob/main/qwen_tts/core/models/modeling_qwen3_tts.py diff --git a/tools/qwen3-tts-onnx/.gitignore b/tools/qwen3-tts-onnx/.gitignore index 9f07dfe..57f1000 100644 --- a/tools/qwen3-tts-onnx/.gitignore +++ b/tools/qwen3-tts-onnx/.gitignore @@ -2,4 +2,5 @@ output/ __pycache__/ .venv/ .claude/ -*.wav \ No newline at end of file +*.wav +!fixtures/*.wav \ No newline at end of file diff --git a/tools/qwen3-tts-onnx/Makefile b/tools/qwen3-tts-onnx/Makefile index 76deb30..ae87676 100644 --- a/tools/qwen3-tts-onnx/Makefile +++ b/tools/qwen3-tts-onnx/Makefile @@ -5,7 +5,7 @@ PYTHON ?= .venv/bin/python FP32_DIR = $(OUTPUT_DIR)/fp32 INT4_DIR = $(OUTPUT_DIR)/int4 -.PHONY: all venv export embeddings talker code-predictor vocoder validate validate-int4 quantize hf generate generate-int4 generate-zh generate-zh-int4 clean help +.PHONY: all venv export embeddings talker code-predictor vocoder speaker-encoder tokenizer-encoder validate validate-int4 quantize hf generate generate-int4 generate-zh generate-zh-int4 clean help clone-all clone-export clone-base-preset help: ## Show this help @echo "Qwen3-TTS ONNX Export & Quantization Toolkit" @@ -18,6 +18,8 @@ help: ## Show this help @echo "Variables:" @echo " MODEL_ID HuggingFace model ID (default: Qwen/Qwen3-TTS-12Hz-1.7B-VoiceDesign)" @echo " OUTPUT_DIR Output directory (default: ./output/qwen3-tts-1.7b-voicedesign)" + @echo " CLONE_MODEL_ID Voice clone model ID (default: Qwen/Qwen3-TTS-12Hz-0.6B-Base)" + @echo " CLONE_OUTPUT_DIR Voice clone output dir (default: ./output/qwen3-tts-0.6b-base)" @echo " PYTHON Python interpreter (default: .venv/bin/python)" venv: ## Create virtualenv and install dependencies @@ -40,6 +42,12 @@ code-predictor: ## Export code predictor ONNX model vocoder: ## Export vocoder ONNX model $(PYTHON) export_vocoder.py --model-id $(MODEL_ID) --output-dir $(OUTPUT_DIR) +speaker-encoder: ## Export ECAPA-TDNN speaker encoder ONNX (Base models only) + $(PYTHON) export_speaker_encoder.py --model-id $(MODEL_ID) --output-dir $(OUTPUT_DIR) + +tokenizer-encoder: ## Export speech-tokenizer encoder ONNX (Base models only) + $(PYTHON) export_tokenizer_encoder.py --model-id $(MODEL_ID) --output-dir $(OUTPUT_DIR) + validate: ## Validate ONNX exports against PyTorch $(PYTHON) validate.py --model-id $(MODEL_ID) --onnx-dir $(OUTPUT_DIR) @@ -85,3 +93,59 @@ generate-zh-int4: ## Generate Chinese sample audio (INT4) clean: ## Remove output directory rm -rf $(OUTPUT_DIR) + +# --------------------------------------------------------------------------- +# 0.6B Base voice clone preset +# +# Same export scripts as VoiceDesign, plus the two extra encoders. Pin to a +# separate output directory so 1.7B and 0.6B builds don't stomp each other. +# --------------------------------------------------------------------------- + +CLONE_MODEL_ID ?= Qwen/Qwen3-TTS-12Hz-0.6B-Base +CLONE_OUTPUT_DIR ?= ./output/qwen3-tts-0.6b-base + +clone-all: clone-export clone-base-preset clone-hf ## Full 0.6B Base export + INT4 + HF packaging + +clone-export: ## Export all ONNX models for the 0.6B Base voice clone variant + @echo "\n[1/6] Embeddings, config, tokenizer ..." + $(MAKE) embeddings MODEL_ID=$(CLONE_MODEL_ID) OUTPUT_DIR=$(CLONE_OUTPUT_DIR) + @echo "\n[2/6] Talker prefill + decode ..." + $(MAKE) talker MODEL_ID=$(CLONE_MODEL_ID) OUTPUT_DIR=$(CLONE_OUTPUT_DIR) + @echo "\n[3/6] Code predictor ..." + $(MAKE) code-predictor MODEL_ID=$(CLONE_MODEL_ID) OUTPUT_DIR=$(CLONE_OUTPUT_DIR) + @echo "\n[4/6] Vocoder ..." + $(MAKE) vocoder MODEL_ID=$(CLONE_MODEL_ID) OUTPUT_DIR=$(CLONE_OUTPUT_DIR) + @echo "\n[5/6] Speaker encoder (ECAPA-TDNN) ..." + $(MAKE) speaker-encoder MODEL_ID=$(CLONE_MODEL_ID) OUTPUT_DIR=$(CLONE_OUTPUT_DIR) + @echo "\n[6/6] Tokenizer encoder (Mimi) ..." + $(MAKE) tokenizer-encoder MODEL_ID=$(CLONE_MODEL_ID) OUTPUT_DIR=$(CLONE_OUTPUT_DIR) + @echo "\nAll 0.6B Base clone exports complete." + +clone-hf: ## Package 0.6B Base clone output for Hugging Face upload + @echo "Packaging 0.6B Base clone for Hugging Face..." + cp README_CLONE.md $(CLONE_OUTPUT_DIR)/README.md + sed 's|./output/qwen3-tts-0.6b-base|.|' generate_clone_onnx.py > $(CLONE_OUTPUT_DIR)/generate_clone_onnx.py + printf 'onnxruntime>=1.24.0\nnumpy>=2.0\nsoundfile>=0.13\ntransformers>=4.57.0\nlibrosa>=0.10\n' > $(CLONE_OUTPUT_DIR)/requirements.txt + printf '*.onnx filter=lfs diff=lfs merge=lfs -text\n*.onnx.data filter=lfs diff=lfs merge=lfs -text\n*.npy filter=lfs diff=lfs merge=lfs -text\n' > $(CLONE_OUTPUT_DIR)/.gitattributes + @echo "Done: $(CLONE_OUTPUT_DIR) is ready for 'huggingface-cli upload'" + +clone-base-preset: ## INT4 quantize the talker/CP for the 0.6B Base build (encoders stay FP32) + @echo "\nQuantizing talker + CP to INT4 (speaker/tokenizer encoders stay FP32) ..." + $(MAKE) quantize MODEL_ID=$(CLONE_MODEL_ID) OUTPUT_DIR=$(CLONE_OUTPUT_DIR) + +CLONE_REF_AUDIO ?= fixtures/ref_clone.wav +CLONE_REF_TEXT ?= Give every small business the voice of a big one. With WaveKat, your phone is always answered by a voice your customers trust. + +clone-fixture: ## Regenerate fixtures/ref_clone.wav via wavekat-tts (1.7B VoiceDesign) + @mkdir -p fixtures + cd ../.. && cargo run --release --example synthesize --features qwen3-tts -- \ + --precision fp32 \ + --instruction "Speak in a warm and friendly female voice with natural intonation." \ + --output tools/qwen3-tts-onnx/fixtures/ref_clone.wav \ + "$(CLONE_REF_TEXT)" + +clone-generate: ## Generate sample voice-clone audio (FP32) + $(PYTHON) generate_clone_onnx.py --model-dir $(CLONE_OUTPUT_DIR) \ + --ref-audio "$(CLONE_REF_AUDIO)" --ref-text "$(CLONE_REF_TEXT)" \ + --text "Your customers deserve a voice they can trust, every time they call." \ + -o clone_output.wav diff --git a/tools/qwen3-tts-onnx/README_CLONE.md b/tools/qwen3-tts-onnx/README_CLONE.md new file mode 100644 index 0000000..f93bcc0 --- /dev/null +++ b/tools/qwen3-tts-onnx/README_CLONE.md @@ -0,0 +1,151 @@ +--- +language: + - en + - zh + - ja + - ko + - de + - fr + - es + - it + - pt + - ru +license: apache-2.0 +tags: + - text-to-speech + - tts + - onnx + - qwen3-tts + - voice-cloning +library_name: onnxruntime +pipeline_tag: text-to-speech +base_model: Qwen/Qwen3-TTS-12Hz-0.6B-Base +--- + +

+ + WaveKat TTS + +

+ +# Qwen3-TTS 0.6B Base — Voice Clone (ONNX) + +ONNX export of [Qwen/Qwen3-TTS-12Hz-0.6B-Base](https://huggingface.co/Qwen/Qwen3-TTS-12Hz-0.6B-Base) for **voice cloning** with ONNX Runtime. No PyTorch required at inference time. + +Provide a short reference audio clip and its transcript, and the model synthesizes new text in the same voice using In-Context Learning (ICL). + +Both FP32 and INT4 (weight-only, RTN) variants are included. + +> Exported and maintained by [WaveKat](https://github.com/wavekat) as part of the [wavekat-tts](https://github.com/wavekat/wavekat-tts) voice pipeline. + +## Quick Start + +```bash +pip install -r requirements.txt + +# FP32 +python generate_clone_onnx.py \ + --ref-audio ref.wav --ref-text "Transcript of the reference audio." \ + --text "New text to synthesize in the cloned voice." \ + -o output_fp32.wav + +# INT4 (~4x smaller, faster) +python generate_clone_onnx.py --variant int4 \ + --ref-audio ref.wav --ref-text "Transcript of the reference audio." \ + --text "New text to synthesize in the cloned voice." \ + -o output_int4.wav +``` + +Reference audio should be **mono 24 kHz WAV** with a clear, single-speaker recording (3–10 seconds works well). The script will resample automatically via librosa if needed. + +## Model Architecture + +Qwen3-TTS 0.6B Base uses a 6-model voice-clone pipeline: + +``` +Ref audio --> [Speaker Encoder] ECAPA-TDNN → 1024-d speaker embedding + \-> [Tokenizer Encoder] Mimi encoder → 16-group ref codes (12 Hz) + +Text + Ref text + Speaker embed + Ref codes + | + v (ICL prefill) + [Talker LM] 28 layers, 1024 hidden + predicts codebook group 0 + | + v + [Code Predictor] 5 layers, 1024 hidden + predicts groups 1-15 + | + v + [Vocoder] single forward pass + concat(ref_codes, gen_codes) → 24 kHz waveform → trim ref portion +``` + +The pipeline is split into 6 ONNX models: + +| Model | Description | Precision | +|-------|-------------|-----------| +| `speaker_encoder.onnx` | ECAPA-TDNN: mel → 1024-d speaker embedding | FP32 only | +| `tokenizer_encoder.onnx` | Mimi encoder: audio → 16-group codec codes | FP32 only | +| `talker_prefill.onnx` | Full sequence prefill with KV cache output | FP32 / INT4 | +| `talker_decode.onnx` | Single-step decode with KV cache | FP32 / INT4 | +| `code_predictor.onnx` | Predict codebook groups 1-15 | FP32 / INT4 | +| `vocoder.onnx` | Codes to 24 kHz waveform | FP32 / INT4 | + +> Speaker encoder and tokenizer encoder are always FP32 — they run once per request and are small. + +## Repository Structure + +``` +. +├── config.json # Model config (dimensions, token IDs, language map) +├── speaker_encoder.onnx # ECAPA-TDNN speaker encoder (FP32) +├── tokenizer_encoder.onnx # Mimi speech tokenizer encoder (FP32) +├── tokenizer/ # Text tokenizer (vocab, merges) +├── embeddings/ # Pre-extracted embedding weights (.npy) +├── fp32/ # FP32 ONNX models +│ ├── talker_prefill.onnx +│ ├── talker_decode.onnx +│ ├── code_predictor.onnx +│ └── vocoder.onnx +├── int4/ # INT4 weight-only quantized models +│ ├── talker_prefill.onnx +│ ├── talker_decode.onnx +│ ├── code_predictor.onnx +│ └── vocoder.onnx +├── generate_clone_onnx.py # Reference ONNX-only voice clone script +└── requirements.txt # Inference dependencies +``` + +## How It Works + +1. **Speaker encoding** — the reference audio is converted to a log-mel spectrogram and passed through an ECAPA-TDNN encoder to produce a 1024-d speaker embedding. +2. **Reference code extraction** — the same audio is encoded by a Mimi tokenizer encoder into 16-group discrete codes at 12 Hz. +3. **ICL prefill** — the talker LM is prefilled with an interleaved sequence: text embeddings (ref transcript + target text) paired with codec embeddings (speaker embed + reference codes). +4. **Autoregressive decode** — the talker generates group-0 codec tokens, and the code predictor fills in groups 1-15 per frame. +5. **Vocoder** — reference codes are prepended to generated codes, the vocoder decodes the combined sequence, and the leading reference portion is trimmed proportionally. + +## Supported Languages + +English, Chinese, Japanese, Korean, German, French, Spanish, Italian, Portuguese, Russian. + +## Reproducing the Export + +The export scripts are in the [wavekat-tts](https://github.com/wavekat/wavekat-tts) repository: + +```bash +cd tools/qwen3-tts-onnx +pip install -r requirements.txt + +# Export FP32, quantize INT4, and package for HF +make clone-all +``` + +## About WaveKat + +[WaveKat](https://github.com/wavekat) builds open-source voice pipeline components in Rust. +This ONNX export is maintained as part of [wavekat-tts](https://github.com/wavekat/wavekat-tts), which provides unified TTS inference across multiple backends. + +## Acknowledgements + +- [Qwen3-TTS](https://huggingface.co/Qwen/Qwen3-TTS-12Hz-0.6B-Base) by the Qwen team at Alibaba Cloud diff --git a/tools/qwen3-tts-onnx/export_speaker_encoder.py b/tools/qwen3-tts-onnx/export_speaker_encoder.py new file mode 100644 index 0000000..0606c20 --- /dev/null +++ b/tools/qwen3-tts-onnx/export_speaker_encoder.py @@ -0,0 +1,178 @@ +#!/usr/bin/env python3 +"""Export Qwen3-TTS-Base speaker encoder (ECAPA-TDNN) as speaker_encoder.onnx. + +The speaker encoder is the conditioning module on the Base ("voice clone") +checkpoints. It maps an 80-bin mel-spectrogram of the reference clip to a +1024-dim speaker embedding, which is inserted into the talker prefill at the +speaker-slot position (after `codec_think_eos`, before the transition). + +Mel input: (B, T_mel, mel_dim=128), float32 + Computed in host code (not part of ONNX) with the standard HiFi-GAN + parameters used by the reference pipeline: + n_fft=1024, hop_size=256, win_size=1024, num_mels=128, fmin=0, fmax=12000, + sample_rate=24000, center=False, log on top of dynamic-range compression. + +Output: (B, enc_dim=1024), float32 + +The encoder lives on `model.speaker_encoder` and only exists when +`config.tts_model_type == "base"`. VoiceDesign / CustomVoice checkpoints don't +have it; running this script against them is a no-op + clear error. +""" + +import argparse +import os + +import numpy as np +import onnx +import torch +import torch.nn as nn + +from qwen_tts.core.models.modeling_qwen3_tts import Qwen3TTSForConditionalGeneration + + +class SpeakerEncoderWrapper(nn.Module): + """Pass-through wrapper around the ECAPA-TDNN speaker encoder. + + The reference pipeline calls `speaker_encoder(mels)[0]` to drop the batch + dim. We keep the batch dim in ONNX so callers can decide what to do. + """ + + def __init__(self, speaker_encoder): + super().__init__() + self.speaker_encoder = speaker_encoder + + def forward(self, mels): # (B, T_mel, mel_dim) + return self.speaker_encoder(mels) # (B, enc_dim) + + +def export_speaker_encoder(model_id: str, output_dir: str): + print(f"Loading model: {model_id}") + model = Qwen3TTSForConditionalGeneration.from_pretrained( + model_id, dtype=torch.float32, attn_implementation="eager" + ) + model.eval() + + if model.speaker_encoder is None: + raise SystemExit( + f"Model {model_id} (tts_model_type={model.config.tts_model_type}) " + "has no speaker encoder. Voice clone export only applies to Base " + "checkpoints (e.g. Qwen/Qwen3-TTS-12Hz-0.6B-Base)." + ) + + spk_cfg = model.config.speaker_encoder_config + print( + f" Speaker encoder: mel_dim={spk_cfg.mel_dim}, " + f"enc_dim={spk_cfg.enc_dim}, sample_rate={spk_cfg.sample_rate}" + ) + + wrapper = SpeakerEncoderWrapper(model.speaker_encoder) + wrapper.eval() + + # Trace with a representative mel length: ~3 s at 24 kHz / hop 256 + # ≈ 281 frames. Use 300 for headroom. + T_mel = 300 + dummy_mels = torch.randn(1, T_mel, spk_cfg.mel_dim, dtype=torch.float32) + + onnx_path = os.path.join(output_dir, "speaker_encoder.onnx") + os.makedirs(output_dir, exist_ok=True) + print(f"\nExporting speaker_encoder.onnx (trace, dynamic mel length, T={T_mel}) ...") + + pre_export = set(os.listdir(output_dir)) if os.path.exists(output_dir) else set() + + with torch.no_grad(): + torch.onnx.export( + wrapper, + (dummy_mels,), + onnx_path, + opset_version=17, + dynamo=False, + input_names=["mels"], + output_names=["speaker_embedding"], + dynamic_axes={ + "mels": {0: "batch", 1: "mel_frames"}, + "speaker_embedding": {0: "batch"}, + }, + ) + + _try_consolidate(onnx_path, pre_export) + print(f" Saved: {onnx_path}") + + _validate(wrapper, dummy_mels, onnx_path) + # Spot-check dynamic axis at a couple of other lengths. + for test_T in [100, 500]: + test_mels = torch.randn(1, test_T, spk_cfg.mel_dim, dtype=torch.float32) + _validate(wrapper, test_mels, onnx_path, label=f"T_mel={test_T}") + + print("\nSpeaker encoder export complete.") + + +def _try_consolidate(onnx_path: str, pre_export_files: set | None = None): + """Consolidate external data into a single .onnx.data file if needed.""" + onnx_dir = os.path.dirname(onnx_path) + data_path = onnx_path + ".data" + + try: + m = onnx.load(onnx_path) + onnx.save_model( + m, + onnx_path, + save_as_external_data=True, + all_tensors_to_one_file=True, + location=os.path.basename(data_path), + ) + except Exception as e: + print(f" Note: consolidation skipped ({e})") + return + + if pre_export_files is not None: + current_files = set(os.listdir(onnx_dir)) + scattered = current_files - pre_export_files - { + os.path.basename(onnx_path), + os.path.basename(data_path), + } + for f in scattered: + path = os.path.join(onnx_dir, f) + if os.path.isfile(path): + os.remove(path) + if scattered: + print(f" Cleaned up {len(scattered)} scattered external data files") + + +def _validate(wrapper, mels, onnx_path, label=None): + import onnxruntime as ort + + with torch.no_grad(): + pt_out = wrapper(mels) + + sess = ort.InferenceSession(onnx_path) + ort_out = sess.run(None, {"mels": mels.numpy()})[0] + + pt_arr = pt_out.numpy() + max_err = float(np.max(np.abs(pt_arr - ort_out))) + tag = f" ({label})" if label else "" + print( + f" Speaker encoder validation{tag}: max_err={max_err:.6e}, " + f"shape={ort_out.shape}" + ) + if max_err > 1e-4: + print(f" WARNING: max error {max_err:.6e} exceeds 1e-4 threshold") + + +def main(): + parser = argparse.ArgumentParser(description="Export Qwen3-TTS speaker encoder to ONNX") + parser.add_argument( + "--model-id", + default="Qwen/Qwen3-TTS-12Hz-0.6B-Base", + help="HuggingFace model ID (must be a Base checkpoint)", + ) + parser.add_argument( + "--output-dir", + default="./output/qwen3-tts-0.6b-base", + help="Output directory (speaker_encoder.onnx is written at the root, alongside fp32/, int4/, etc.)", + ) + args = parser.parse_args() + export_speaker_encoder(args.model_id, args.output_dir) + + +if __name__ == "__main__": + main() diff --git a/tools/qwen3-tts-onnx/export_tokenizer_encoder.py b/tools/qwen3-tts-onnx/export_tokenizer_encoder.py new file mode 100644 index 0000000..773168d --- /dev/null +++ b/tools/qwen3-tts-onnx/export_tokenizer_encoder.py @@ -0,0 +1,220 @@ +#!/usr/bin/env python3 +"""Export the speech-tokenizer encoder as tokenizer_encoder.onnx. + +Voice clone needs the *reference* audio turned into the same discrete codes the +talker normally produces. The Base model uses `model.speech_tokenizer.encode()` +for this — a Mimi-based residual VQ encoder that runs at 12 Hz with 16 +quantizer groups. + +Pipeline (as in qwen_tts.inference.qwen3_tts_tokenizer.Qwen3TTSTokenizer.encode): + + audio (24 kHz mono float32, 1-D) + └── unsqueeze ──> (1, T) waveform + └── inner encoder.encode(input_values=(1, 1, T)) + └── codes: (1, num_quantizers, frames) + └── slice [:, :encoder_valid_num_quantizers] + +We expose the same signature in ONNX: + + Input : waveform `(1, FIXED_SAMPLES)` float32, 24 kHz + Output: audio_codes `(1, num_valid, frames)` int64 + +NOTE: The Mimi encoder uses data-dependent conv padding (`.item()` in +`_get_extra_padding_for_conv1d`), which makes `torch.export` / dynamo fail +on dynamic shapes. We therefore use the **legacy JIT tracer** with a fixed +canonical sample length. Host code must zero-pad (or truncate) the reference +waveform to exactly `CANONICAL_SAMPLES` before feeding it to the ONNX session, +and trim trailing code frames based on the original audio length. + +The encoder only runs *once per reference clip* (not in the autoregressive +loop), so the fixed-size constraint has no performance impact. +""" + +import argparse +import os + +import numpy as np +import onnx +import torch +import torch.nn as nn + +from qwen_tts.core.models.modeling_qwen3_tts import Qwen3TTSForConditionalGeneration +from mask_patch import patch_causal_mask + +# Canonical sample count: 10 s × 24 kHz = 240 000 samples. +# Covers reference clips up to 10 seconds. Shorter clips are zero-padded; +# the host code trims the output codes based on the original sample count. +CANONICAL_SECONDS = 10 +CANONICAL_SR = 24000 +CANONICAL_SAMPLES = CANONICAL_SECONDS * CANONICAL_SR # 240_000 + + +class TokenizerEncoderWrapper(nn.Module): + """Wraps `Qwen3TTSTokenizerV2Encoder` (a `MimiModel` subclass). + + Calls the inner encoder's `.encode()` and slices to the valid quantizer + count. Input is a fixed-length waveform; output is the code matrix. + """ + + def __init__(self, encoder, num_valid_quantizers: int): + super().__init__() + self.encoder = encoder + self.num_valid = num_valid_quantizers + + def forward(self, waveform): # (1, CANONICAL_SAMPLES) float32 + # Mimi encoder expects (B, 1, T) + encoded = self.encoder.encode( + input_values=waveform.unsqueeze(1), + return_dict=True, + ) + # encoded.audio_codes: (B, num_quantizers, frames) int64 + return encoded.audio_codes[:, : self.num_valid] + + +def export_tokenizer_encoder(model_id: str, output_dir: str): + # Mimi's encoder_transformer uses create_causal_mask which calls torch.vmap, + # incompatible with JIT tracing. Patch it the same way as the talker export. + patch_causal_mask() + + print(f"Loading model: {model_id}") + model = Qwen3TTSForConditionalGeneration.from_pretrained( + model_id, dtype=torch.float32, attn_implementation="eager" + ) + model.eval() + + speech_tokenizer = model.speech_tokenizer + if speech_tokenizer is None: + raise SystemExit( + f"Model {model_id} does not expose a speech_tokenizer. " + "Voice clone export requires a Base checkpoint with the 12Hz " + "tokenizer (e.g. Qwen/Qwen3-TTS-12Hz-0.6B-Base)." + ) + + inner = speech_tokenizer.model + encoder = inner.encoder + num_valid = int(inner.encoder_valid_num_quantizers) + input_sr = int(inner.input_sample_rate) + encode_downsample = int(inner.encode_downsample_rate) + print( + f" Tokenizer encoder: input_sr={input_sr}, " + f"downsample={encode_downsample}, num_valid_quantizers={num_valid}" + ) + assert input_sr == CANONICAL_SR, ( + f"Expected input_sr={CANONICAL_SR}, got {input_sr}. " + f"Update CANONICAL_SR in this script." + ) + + wrapper = TokenizerEncoderWrapper(encoder, num_valid) + wrapper.eval() + + dummy_wav = torch.randn(1, CANONICAL_SAMPLES, dtype=torch.float32) * 0.1 + expected_frames = CANONICAL_SAMPLES // encode_downsample + print( + f" Fixed trace size: {CANONICAL_SAMPLES} samples " + f"({CANONICAL_SECONDS}s @ {CANONICAL_SR} Hz) → {expected_frames} code frames" + ) + + onnx_path = os.path.join(output_dir, "tokenizer_encoder.onnx") + os.makedirs(output_dir, exist_ok=True) + + pre_export = set(os.listdir(output_dir)) if os.path.exists(output_dir) else set() + + # Legacy JIT tracer — no dynamic shapes. The Mimi encoder's conv padding + # uses .item() which creates data-dependent guards incompatible with dynamo. + print( + f"\nExporting tokenizer_encoder.onnx " + f"(JIT trace, fixed T={CANONICAL_SAMPLES}) ..." + ) + with torch.no_grad(): + torch.onnx.export( + wrapper, + (dummy_wav,), + onnx_path, + opset_version=17, + dynamo=False, + input_names=["waveform"], + output_names=["audio_codes"], + ) + + _try_consolidate(onnx_path, pre_export) + print(f" Saved: {onnx_path}") + + _validate(wrapper, dummy_wav, onnx_path) + + print(f"\nTokenizer encoder export complete.") + print(f" IMPORTANT: host code must zero-pad waveforms to exactly {CANONICAL_SAMPLES}") + print(f" samples ({CANONICAL_SECONDS}s @ {CANONICAL_SR} Hz) before inference,") + print(f" then trim output codes to ceil(original_samples / {encode_downsample}) frames.") + + +def _try_consolidate(onnx_path: str, pre_export_files: set | None = None): + onnx_dir = os.path.dirname(onnx_path) + data_path = onnx_path + ".data" + + try: + m = onnx.load(onnx_path) + onnx.save_model( + m, + onnx_path, + save_as_external_data=True, + all_tensors_to_one_file=True, + location=os.path.basename(data_path), + ) + except Exception as e: + print(f" Note: consolidation skipped ({e})") + return + + if pre_export_files is not None: + current_files = set(os.listdir(onnx_dir)) + scattered = current_files - pre_export_files - { + os.path.basename(onnx_path), + os.path.basename(data_path), + } + for f in scattered: + path = os.path.join(onnx_dir, f) + if os.path.isfile(path): + os.remove(path) + if scattered: + print(f" Cleaned up {len(scattered)} scattered external data files") + + +def _validate(wrapper, waveform, onnx_path, label=None): + import onnxruntime as ort + + with torch.no_grad(): + pt_codes = wrapper(waveform) + + sess = ort.InferenceSession(onnx_path) + ort_codes = sess.run(None, {"waveform": waveform.numpy()})[0] + + pt_arr = pt_codes.numpy() + # Codes are integer — they should match exactly. + same = bool(np.array_equal(pt_arr, ort_codes)) + diffs = int(np.sum(pt_arr != ort_codes)) + tag = f" ({label})" if label else "" + print( + f" Tokenizer encoder validation{tag}: identical={same}, " + f"diffs={diffs}, shape={ort_codes.shape}" + ) + if not same: + print(f" WARNING: {diffs} code mismatches between PyTorch and ONNX") + + +def main(): + parser = argparse.ArgumentParser(description="Export Qwen3-TTS speech tokenizer encoder to ONNX") + parser.add_argument( + "--model-id", + default="Qwen/Qwen3-TTS-12Hz-0.6B-Base", + help="HuggingFace model ID (must include a 12Hz speech tokenizer)", + ) + parser.add_argument( + "--output-dir", + default="./output/qwen3-tts-0.6b-base", + help="Output directory (tokenizer_encoder.onnx is written at the root)", + ) + args = parser.parse_args() + export_tokenizer_encoder(args.model_id, args.output_dir) + + +if __name__ == "__main__": + main() diff --git a/tools/qwen3-tts-onnx/fixtures/README.md b/tools/qwen3-tts-onnx/fixtures/README.md new file mode 100644 index 0000000..68b7fd6 --- /dev/null +++ b/tools/qwen3-tts-onnx/fixtures/README.md @@ -0,0 +1,17 @@ +# fixtures/ + +## ref_clone.wav + +Reference audio clip used for voice-clone testing and examples. + +- **Generated by**: `make clone-fixture` (1.7B VoiceDesign, FP32) +- **Voice instruction**: "Speak in a warm and friendly female voice with natural intonation." +- **Text**: "Give every small business the voice of a big one. With WaveKat, your phone is always answered by a voice your customers trust." +- **Sample rate**: 24 kHz mono + +To regenerate: + +```bash +cd tools/qwen3-tts-onnx +make clone-fixture +``` diff --git a/tools/qwen3-tts-onnx/fixtures/ref_clone.wav b/tools/qwen3-tts-onnx/fixtures/ref_clone.wav new file mode 100644 index 0000000..fad4d51 Binary files /dev/null and b/tools/qwen3-tts-onnx/fixtures/ref_clone.wav differ diff --git a/tools/qwen3-tts-onnx/generate_clone_onnx.py b/tools/qwen3-tts-onnx/generate_clone_onnx.py new file mode 100644 index 0000000..fc513cd --- /dev/null +++ b/tools/qwen3-tts-onnx/generate_clone_onnx.py @@ -0,0 +1,546 @@ +#!/usr/bin/env python3 +"""Generate WAV with voice cloning using ONNX models (no PyTorch at inference). + +This is the end-to-end verification script for the 0.6B Base voice-clone ONNX +pipeline. It chains all 6 exported models: + + 1. tokenizer_encoder.onnx — ref audio → ref codes (16 groups, 12 Hz) + 2. speaker_encoder.onnx — ref mel → speaker embedding (1024-d) + 3. talker_prefill.onnx — ICL prefill → logits + KV cache + 4. talker_decode.onnx — single-step decode loop + 5. code_predictor.onnx — groups 1-15 per frame + 6. vocoder.onnx — codes → 24 kHz waveform + +Examples: + python generate_clone_onnx.py \\ + --ref-audio clone.wav --ref-text "Okay. Yeah. I resent you." \\ + --text "Give every small business the voice of a big one." \\ + -o cloned.wav + + python generate_clone_onnx.py --variant int4 \\ + --ref-audio clone.wav --ref-text "Okay. Yeah. I resent you." \\ + --text "让每一家小企业,都拥有大企业的声音。" --lang chinese \\ + -o cloned_zh.wav +""" + +import argparse +import json +import os +import time + +import librosa +import numpy as np +import onnxruntime as ort +import soundfile as sf +from transformers import AutoTokenizer + +# Mel-spectrogram parameters (match the reference HiFi-GAN config). +MEL_SR = 24000 +MEL_N_FFT = 1024 +MEL_HOP = 256 +MEL_WIN = 1024 +MEL_N_MELS = 128 +MEL_FMIN = 0 +MEL_FMAX = 12000 + +# Tokenizer encoder canonical input length (must match export_tokenizer_encoder.py). +TOKENIZER_CANONICAL_SAMPLES = 10 * MEL_SR # 240_000 + + +# --------------------------------------------------------------------------- +# Audio / mel helpers +# --------------------------------------------------------------------------- + +def load_ref_audio(path: str, target_sr: int = MEL_SR) -> np.ndarray: + """Load + resample reference audio to mono float32 at target_sr.""" + audio, sr = librosa.load(path, sr=target_sr, mono=True) + return audio.astype(np.float32) + + +def compute_mel(audio: np.ndarray) -> np.ndarray: + """Compute log-mel spectrogram matching the Qwen3-TTS reference. + + Returns (1, T_mel, 128) float32. + """ + mel = librosa.feature.melspectrogram( + y=audio, sr=MEL_SR, n_fft=MEL_N_FFT, hop_length=MEL_HOP, + win_length=MEL_WIN, n_mels=MEL_N_MELS, fmin=MEL_FMIN, fmax=MEL_FMAX, + center=False, + ) + mel = np.log(np.clip(mel, a_min=1e-5, a_max=None)) + # (n_mels, frames) → (1, frames, n_mels) + return mel.T[np.newaxis, :, :].astype(np.float32) + + +def pad_for_tokenizer_encoder(audio: np.ndarray) -> tuple[np.ndarray, int]: + """Zero-pad / truncate to the canonical sample count. + + Returns (padded_waveform (1, CANONICAL), original_sample_count). + """ + n = len(audio) + if n > TOKENIZER_CANONICAL_SAMPLES: + audio = audio[:TOKENIZER_CANONICAL_SAMPLES] + n = TOKENIZER_CANONICAL_SAMPLES + padded = np.zeros(TOKENIZER_CANONICAL_SAMPLES, dtype=np.float32) + padded[:n] = audio + return padded[np.newaxis, :], n + + +# --------------------------------------------------------------------------- +# Embedding / sampling helpers (copied from generate_onnx.py) +# --------------------------------------------------------------------------- + +def text_project_numpy(token_ids, text_emb, fc1_w, fc1_b, fc2_w, fc2_b): + embeds = text_emb[token_ids] + hidden = embeds @ fc1_w.T + fc1_b + activated = hidden * (1.0 / (1.0 + np.exp(-hidden))) + return activated @ fc2_w.T + fc2_b + + +def load_embeddings(onnx_dir): + edir = os.path.join(onnx_dir, "embeddings") + d = {} + for name in [ + "text_embedding", + "text_projection_fc1_weight", "text_projection_fc1_bias", + "text_projection_fc2_weight", "text_projection_fc2_bias", + "talker_codec_embedding", + ]: + d[name] = np.load(os.path.join(edir, f"{name}.npy")) + d["cp_codec_embeddings"] = [] + i = 0 + while True: + path = os.path.join(edir, f"cp_codec_embedding_{i}.npy") + if not os.path.exists(path): + break + d["cp_codec_embeddings"].append(np.load(path)) + i += 1 + return d + + +def load_config(onnx_dir): + with open(os.path.join(onnx_dir, "config.json")) as f: + return json.load(f) + + +def sample_top_k(logits, top_k, temperature): + if temperature != 1.0: + logits = logits / temperature + if top_k > 0 and top_k < len(logits): + idx = np.argpartition(logits, -top_k)[-top_k:] + mask = np.full_like(logits, -np.inf) + mask[idx] = logits[idx] + logits = mask + logits = logits - np.max(logits) + probs = np.exp(logits) + probs = probs / probs.sum() + return int(np.random.choice(len(probs), p=probs)) + + +# --------------------------------------------------------------------------- +# Main pipeline +# --------------------------------------------------------------------------- + +def generate_clone_onnx( + model_dir: str, + variant: str, + text: str, + ref_audio_path: str, + ref_text: str, + language: str, + output_path: str, + max_new_tokens: int, + temperature: float, + top_k: int, + repetition_penalty: float, + seed: int | None, +): + if seed is not None: + np.random.seed(seed) + + onnx_dir = os.path.join(model_dir, variant) + config = load_config(model_dir) + emb = load_embeddings(model_dir) + tokenizer = AutoTokenizer.from_pretrained(os.path.join(model_dir, "tokenizer")) + + # ------------------------------------------------------------------ + # 1. Load ONNX sessions + # ------------------------------------------------------------------ + print(f"Loading ONNX models ({variant}) ...") + prefill_sess = ort.InferenceSession(os.path.join(onnx_dir, "talker_prefill.onnx")) + decode_sess = ort.InferenceSession(os.path.join(onnx_dir, "talker_decode.onnx")) + cp_sess = ort.InferenceSession(os.path.join(onnx_dir, "code_predictor.onnx")) + vocoder_sess = ort.InferenceSession(os.path.join(onnx_dir, "vocoder.onnx")) + spk_sess = ort.InferenceSession(os.path.join(model_dir, "speaker_encoder.onnx")) + tok_enc_sess = ort.InferenceSession(os.path.join(model_dir, "tokenizer_encoder.onnx")) + + # ------------------------------------------------------------------ + # 2. Precompute embedding accessors + # ------------------------------------------------------------------ + text_emb = emb["text_embedding"] + fc1_w = emb["text_projection_fc1_weight"] + fc1_b = emb["text_projection_fc1_bias"] + fc2_w = emb["text_projection_fc2_weight"] + fc2_b = emb["text_projection_fc2_bias"] + codec_emb = emb["talker_codec_embedding"] + cp_codec_embs = emb["cp_codec_embeddings"] + hidden_size = config["talker_hidden_size"] + num_code_groups = config["talker_num_code_groups"] + + def text_proj(token_ids): + return text_project_numpy(token_ids, text_emb, fc1_w, fc1_b, fc2_w, fc2_b) + + # ------------------------------------------------------------------ + # 3. Encode reference audio + # ------------------------------------------------------------------ + print(f" Ref audio: {ref_audio_path}") + ref_audio = load_ref_audio(ref_audio_path) + ref_duration = len(ref_audio) / MEL_SR + print(f" Ref duration: {ref_duration:.1f}s ({len(ref_audio)} samples)") + + # Speaker embedding + mel = compute_mel(ref_audio) + print(f" Mel shape: {mel.shape}") + spk_embed = spk_sess.run(None, {"mels": mel})[0] # (1, 1024) + print(f" Speaker embedding: {spk_embed.shape}") + + # Reference codes + padded_wav, orig_samples = pad_for_tokenizer_encoder(ref_audio) + ref_codes_full = tok_enc_sess.run(None, {"waveform": padded_wav})[0] # (1, 16, 125) + # Trim to actual frames + encode_downsample = 1920 + actual_frames = int(np.ceil(orig_samples / encode_downsample)) + ref_codes = ref_codes_full[:, :, :actual_frames] # (1, 16, actual_frames) + print(f" Ref codes: {ref_codes.shape} (trimmed from {ref_codes_full.shape[2]})") + + # ------------------------------------------------------------------ + # 4. Tokenize text & ref_text + # ------------------------------------------------------------------ + chat_text = f"<|im_start|>assistant\n{text}<|im_end|>\n<|im_start|>assistant\n" + input_ids = tokenizer.encode(chat_text, add_special_tokens=False) + + ref_chat = f"<|im_start|>assistant\n{ref_text}<|im_end|>\n" + ref_ids = tokenizer.encode(ref_chat, add_special_tokens=False) + + print(f" Text: '{text}' ({len(input_ids)} tokens)") + print(f" Ref text: '{ref_text}' ({len(ref_ids)} tokens)") + print(f" Language: {language}") + + # ------------------------------------------------------------------ + # 5. Build prefill embeddings (ICL, non-streaming) + # ------------------------------------------------------------------ + language_id = config["codec_language_id"].get(language.lower()) + if language_id is not None: + codec_prefix_ids = [ + config["codec_think_id"], + config["codec_think_bos_id"], + language_id, + config["codec_think_eos_id"], + ] + else: + codec_prefix_ids = [ + config["codec_nothink_id"], + config["codec_think_bos_id"], + config["codec_think_eos_id"], + ] + + tts_pad_embed = text_proj([config["tts_pad_token_id"]])[0] # (hidden,) + tts_bos_embed = text_proj([config["tts_bos_token_id"]])[0] + tts_eos_embed = text_proj([config["tts_eos_token_id"]])[0] + codec_pad_id = config["codec_pad_id"] + codec_bos_id = config["codec_bos_id"] + codec_pad_embed = codec_emb[codec_pad_id] + codec_bos_embed = codec_emb[codec_bos_id] + + embeds_list = [] + + # 5a. Role prefix: first 3 tokens (text proj only) + role_embed = text_proj(input_ids[:3]) # (3, hidden) + embeds_list.append(role_embed) + + # 5b. Codec prefix: tts_pad + codec_embed for each prefix token + for cid in codec_prefix_ids: + embeds_list.append((tts_pad_embed + codec_emb[cid]).reshape(1, -1)) + + # 5c. Speaker slot: speaker_embed (1024-d projected into hidden_size via sum) + # The reference code does: speaker_embed.view(1, 1, -1) in the codec embedding position + # Since enc_dim == hidden_size for 0.6B (both 1024), it's direct. + embeds_list.append(spk_embed.reshape(1, -1)) + + # 5d. Transition: codec_pad + codec_bos (tts_pad + tts_bos on the text side) + embeds_list.append((tts_pad_embed + codec_pad_embed).reshape(1, -1)) + # Note: the reference builds this as: + # tts_pad * (num_codec_prefix + 1 speaker) | tts_bos + codec_prefix[:-1] | codec[last] + # But we already laid out codec prefix above. The last two before ICL are: + # tts_bos + codec_pad (transition into text/codec interleave) + # Let me re-derive from the reference code... + + # Actually, let me redo this more carefully following the reference: + # The reference builds: + # _talker_input_embed = tts_pad.expand(N-2) | tts_bos + codec_input_embedding[:, :-1] + # where codec_input_embedding = [codec_prefix(think,think_bos,lang,think_eos), speaker, codec_pad, codec_bos] + # and [:, :-1] means all except codec_bos + # Then the first text token is paired with codec_bos separately outside of ICL path. + # But in ICL mode, the text path is different... + + # Let me restart the prefill construction cleanly. + embeds_list = [] + + # Part A: Role prefix (3 tokens, text_proj only) + role_embed = text_proj(input_ids[:3]) # (3, hidden) + embeds_list.append(role_embed) + + # Part B: Codec prefix + # codec_input_embedding = [think, think_bos, lang, think_eos, speaker, codec_pad, codec_bos] + # ^speaker slot + # _talker_input_embed = concat(tts_pad × (len-2), tts_bos) + codec_input_embedding[:-1] + # So text side: [tts_pad, tts_pad, tts_pad, tts_pad, tts_pad, tts_bos] + # codec side: [think, think_bos, lang, think_eos, speaker, codec_pad] + # (codec_bos is the last element, used separately) + + codec_full = list(codec_prefix_ids) # [think, think_bos, lang, think_eos] + # Add speaker + codec_pad + codec_bos + num_tts_pad = len(codec_full) + 1 # +1 for speaker slot, then tts_bos pairs with codec_pad + # Total codec_input: [think, think_bos, lang, think_eos, speaker, codec_pad, codec_bos] + # [:-1] = [think, think_bos, lang, think_eos, speaker, codec_pad] + # text side = [tts_pad × 5, tts_bos] (5 = len(codec_prefix) + 1 for speaker) + + for cid in codec_full: + e = tts_pad_embed + codec_emb[cid] + embeds_list.append(e.reshape(1, -1)) + + # Speaker slot: tts_pad + speaker_embed + embeds_list.append((tts_pad_embed + spk_embed[0]).reshape(1, -1)) + + # Transition: tts_bos + codec_pad + embeds_list.append((tts_bos_embed + codec_pad_embed).reshape(1, -1)) + + # Part C: ICL block (non-streaming) + # text_embed = text_proj(ref_ids[3:-2] ++ input_ids[3:-5]) | tts_eos_embed + ref_content_ids = ref_ids[3:-2] # strip role prefix + trailing <|im_end|>\n + text_content_ids = input_ids[3:-5] # strip role prefix + trailing markers + combined_text_ids = list(ref_content_ids) + list(text_content_ids) + text_embed = text_proj(combined_text_ids) # (T1-1, hidden) + text_embed = np.concatenate([text_embed, tts_eos_embed.reshape(1, -1)], axis=0) # (T1, hidden) + T1 = text_embed.shape[0] + + # codec_embed = [codec_bos] | sum_over_groups(codec_embed_g[ref_code]) + # For each ref code frame, sum all 16 group embeddings + ref_frames = ref_codes.shape[2] + codec_frame_embeds = np.zeros((ref_frames, hidden_size), dtype=np.float32) + for f in range(ref_frames): + # Group 0: talker codec embedding + codec_frame_embeds[f] += codec_emb[ref_codes[0, 0, f]] + # Groups 1-15: CP codec embeddings + for g in range(num_code_groups - 1): + codec_frame_embeds[f] += cp_codec_embs[g][ref_codes[0, g + 1, f]] + + # Prepend codec_bos + codec_embed = np.concatenate([ + codec_bos_embed.reshape(1, -1), + codec_frame_embeds, + ], axis=0) # (T2, hidden) where T2 = 1 + ref_frames + T2 = codec_embed.shape[0] + + # Non-streaming ICL interleave: + # icl_input = (text_embed + codec_pad_embed × T1) | (codec_embed + tts_pad × T2) + text_with_pad = text_embed + np.tile(codec_pad_embed, (T1, 1)) + codec_with_pad = codec_embed + np.tile(tts_pad_embed, (T2, 1)) + icl_embed = np.concatenate([text_with_pad, codec_with_pad], axis=0) # (T1+T2, hidden) + embeds_list.append(icl_embed) + + # Trailing text hidden for decode loop = tts_pad_embed (non-streaming) + trailing_hidden = tts_pad_embed.reshape(1, -1) + + # Stack prefill + prefill_embeds = np.concatenate(embeds_list, axis=0)[np.newaxis, :, :].astype(np.float32) + T = prefill_embeds.shape[1] + attention_mask = np.ones((1, T), dtype=np.int64) + position_ids = np.arange(T).reshape(1, 1, T).repeat(3, axis=0) + + print(f" Prefill: {T} tokens (role=3, codec_prefix={len(codec_full)+2}, " + f"ICL text={T1}, ICL codec={T2})") + + # ------------------------------------------------------------------ + # 6. Run prefill + decode loop (same as generate_onnx.py) + # ------------------------------------------------------------------ + num_layers = config["talker_num_layers"] + vocab_size = config["talker_vocab_size"] + codec_eos = config["codec_eos_token_id"] + cp_num_layers = config["cp_num_layers"] + cp_num_kv_heads = config["cp_num_kv_heads"] + cp_head_dim = config["cp_head_dim"] + + suppress_mask = np.zeros(vocab_size, dtype=bool) + suppress_mask[vocab_size - 1024:vocab_size] = True + suppress_mask[codec_eos] = False + + print(" Running prefill ...") + t0 = time.time() + + prefill_out = prefill_sess.run(None, { + "inputs_embeds": prefill_embeds, + "attention_mask": attention_mask, + "position_ids": position_ids, + }) + + logits = prefill_out[0] + hidden_states = prefill_out[1] + kv_outputs = prefill_out[2:] + past_keys = np.stack([kv_outputs[i * 2] for i in range(num_layers)]) + past_values = np.stack([kv_outputs[i * 2 + 1] for i in range(num_layers)]) + + all_codes = [] + current_pos = T + generated_tokens = [] + + print(" Decoding ...") + for step in range(max_new_tokens): + last_logits = logits[0, -1, :].copy() + last_logits[suppress_mask] = -np.inf + if step < 2: + last_logits[codec_eos] = -np.inf + + if repetition_penalty != 1.0 and generated_tokens: + seen = np.array(generated_tokens) + scores = last_logits[seen] + scores = np.where(scores > 0, scores / repetition_penalty, + scores * repetition_penalty) + last_logits[seen] = scores + + group0_token = sample_top_k(last_logits, top_k, temperature) + if group0_token == codec_eos: + break + generated_tokens.append(group0_token) + + # Code predictor: groups 1-15 + frame_codes = [group0_token] + talker_hidden = hidden_states[0, -1:, :] + group0_embed = codec_emb[group0_token].reshape(1, -1) + cp_input = np.concatenate([talker_hidden, group0_embed], axis=0) + cp_input = cp_input[np.newaxis, :, :].astype(np.float32) + cp_past_keys = np.zeros((cp_num_layers, 1, cp_num_kv_heads, 0, cp_head_dim), dtype=np.float32) + cp_past_values = np.zeros((cp_num_layers, 1, cp_num_kv_heads, 0, cp_head_dim), dtype=np.float32) + + for g in range(num_code_groups - 1): + cp_out = cp_sess.run(None, { + "inputs_embeds": cp_input, + "generation_steps": np.array([g], dtype=np.int64), + "past_keys": cp_past_keys, + "past_values": cp_past_values, + }) + cp_past_keys = cp_out[1] + cp_past_values = cp_out[2] + token = sample_top_k(cp_out[0][0, -1, :], top_k, temperature) + frame_codes.append(token) + cp_input = cp_codec_embs[g][token].reshape(1, 1, -1).astype(np.float32) + + all_codes.append(frame_codes) + + # Next talker input + next_embed = codec_emb[group0_token].copy() + for g in range(num_code_groups - 1): + next_embed += cp_codec_embs[g][frame_codes[g + 1]] + next_embed += trailing_hidden[0] + next_embed = next_embed.reshape(1, 1, -1).astype(np.float32) + + decode_mask = np.ones((1, current_pos + 1), dtype=np.int64) + decode_pos = np.array([[[current_pos]]]).repeat(3, axis=0) + + decode_out = decode_sess.run(None, { + "inputs_embeds": next_embed, + "attention_mask": decode_mask, + "position_ids": decode_pos, + "past_keys": past_keys, + "past_values": past_values, + }) + logits = decode_out[0] + hidden_states = decode_out[1] + past_keys = decode_out[2] + past_values = decode_out[3] + current_pos += 1 + + if (step + 1) % 50 == 0: + print(f" ... {step + 1} frames") + + gen_time = time.time() - t0 + num_gen = len(all_codes) + print(f" Generated {num_gen} frames in {gen_time:.1f}s") + + if num_gen == 0: + print(" ERROR: no frames generated") + return + + # ------------------------------------------------------------------ + # 7. Vocoder: prepend ref_codes, decode, trim reference portion + # ------------------------------------------------------------------ + gen_codes = np.array(all_codes, dtype=np.int64) # (gen_frames, 16) + ref_codes_t = ref_codes[0].T # (ref_frames, 16) — was (16, ref_frames) + + all_codes_arr = np.concatenate([ref_codes_t, gen_codes], axis=0) # (total, 16) + codes_input = all_codes_arr.T[np.newaxis, :, :] # (1, 16, total) + total_frames = codes_input.shape[2] + + print(f" Vocoder: {total_frames} frames (ref={ref_frames}, gen={num_gen})") + t0 = time.time() + wav = vocoder_sess.run(None, {"codes": codes_input})[0].flatten() + voc_time = time.time() - t0 + + # Trim leading reference portion (proportional cut like the Python reference) + cut = int(ref_frames / max(total_frames, 1) * len(wav)) + wav = wav[cut:] + + duration = len(wav) / MEL_SR + print(f" Vocoder: {voc_time:.1f}s, output: {duration:.1f}s (trimmed {cut} samples)") + + sf.write(output_path, wav, MEL_SR) + print(f" Saved: {output_path}") + + +# --------------------------------------------------------------------------- +# CLI +# --------------------------------------------------------------------------- + +def main(): + parser = argparse.ArgumentParser( + description="Voice clone via ONNX-only Qwen3-TTS 0.6B Base pipeline", + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog="""\ +Examples: + python generate_clone_onnx.py --ref-audio ref.wav --ref-text "Hello world" --text "New text" + python generate_clone_onnx.py --variant int4 --ref-audio ref.wav --ref-text "Hello" --text "New" +""", + ) + parser.add_argument("--text", required=True, help="Text to synthesize in cloned voice") + parser.add_argument("--ref-audio", required=True, help="Reference audio WAV path") + parser.add_argument("--ref-text", required=True, help="Transcript of reference audio") + parser.add_argument("--lang", default="english", help="Language (default: english)") + parser.add_argument("--model-dir", default="./output/qwen3-tts-0.6b-base", + help="Root model directory") + parser.add_argument("--variant", default="fp32", help="fp32 or int4 (default: fp32)") + parser.add_argument("-o", "--output", default="clone_output.wav", help="Output WAV path") + parser.add_argument("--max-tokens", type=int, default=2048) + parser.add_argument("--temperature", type=float, default=0.9) + parser.add_argument("--top-k", type=int, default=50) + parser.add_argument("--repetition-penalty", type=float, default=1.05) + parser.add_argument("--seed", type=int, default=None) + args = parser.parse_args() + + generate_clone_onnx( + model_dir=args.model_dir, + variant=args.variant, + text=args.text, + ref_audio_path=args.ref_audio, + ref_text=args.ref_text, + language=args.lang, + output_path=args.output, + max_new_tokens=args.max_tokens, + temperature=args.temperature, + top_k=args.top_k, + repetition_penalty=args.repetition_penalty, + seed=args.seed, + ) + + +if __name__ == "__main__": + main() diff --git a/tools/qwen3-tts-onnx/mask_patch.py b/tools/qwen3-tts-onnx/mask_patch.py index 374730f..2a676a2 100644 --- a/tools/qwen3-tts-onnx/mask_patch.py +++ b/tools/qwen3-tts-onnx/mask_patch.py @@ -104,4 +104,16 @@ def simple_sliding_window_causal_mask(config, input_embeds, attention_mask, except ImportError: pass + # Patch the Mimi encoder model (used by the speech tokenizer encoder). + # It has its own `from transformers.masking_utils import create_causal_mask` + # local binding that the module-level patch above doesn't reach. + try: + import transformers.models.mimi.modeling_mimi as mimi_mod + if hasattr(mimi_mod, 'create_causal_mask'): + mimi_mod.create_causal_mask = simple_causal_mask + if hasattr(mimi_mod, 'create_sliding_window_causal_mask'): + mimi_mod.create_sliding_window_causal_mask = simple_sliding_window_causal_mask + except ImportError: + pass + print(" Patched create_causal_mask (vmap-free)")