From d0bdc4a9379a2b663e4ba8f83862afd1ec9e3f01 Mon Sep 17 00:00:00 2001 From: Sandip Dey Date: Tue, 14 Jul 2026 18:07:53 +0100 Subject: [PATCH] feat(speech): add pocket-tts CPU text-to-speech pipeline (ported from babybirdprd/pocket-tts) --- .gitignore | 4 + docs/src/content/docs/examples/index.md | 2 +- .../examples/rust/models/speech-pockettts.md | 55 +++ .../docs/guides/models/use-speech-models.mdx | 109 ++++- .../content/docs/reference/python/enums.md | 1 + .../docs/reference/supported-models.md | 1 + mistralrs-cli/src/args/mod.rs | 14 +- mistralrs-cli/src/commands/serve.rs | 12 +- mistralrs-cli/src/config/mod.rs | 27 +- mistralrs-core/src/model_loader.rs | 2 + mistralrs-core/src/model_metadata.rs | 5 + mistralrs-core/src/model_selected.rs | 4 + mistralrs-core/src/pipeline/auto.rs | 8 + mistralrs-core/src/pipeline/speech.rs | 219 ++++++--- mistralrs-core/src/speech_models/dia/mod.rs | 5 +- mistralrs-core/src/speech_models/mod.rs | 32 +- .../pockettts/conditioners/mod.rs | 1 + .../pockettts/conditioners/text.rs | 390 ++++++++++++++++ .../src/speech_models/pockettts/config.rs | 141 ++++++ .../src/speech_models/pockettts/mod.rs | 84 ++++ .../speech_models/pockettts/models/flow_lm.rs | 151 ++++++ .../speech_models/pockettts/models/mimi.rs | 262 +++++++++++ .../src/speech_models/pockettts/models/mod.rs | 4 + .../speech_models/pockettts/models/seanet.rs | 403 ++++++++++++++++ .../pockettts/models/transformer.rs | 252 ++++++++++ .../pockettts/modules/attention.rs | 291 ++++++++++++ .../speech_models/pockettts/modules/conv.rs | 346 ++++++++++++++ .../speech_models/pockettts/modules/mlp.rs | 418 +++++++++++++++++ .../speech_models/pockettts/modules/mod.rs | 5 + .../speech_models/pockettts/modules/rope.rs | 79 ++++ .../speech_models/pockettts/modules/sdpa.rs | 385 +++++++++++++++ .../src/speech_models/pockettts/pause.rs | 249 ++++++++++ .../src/speech_models/pockettts/tts_model.rs | 442 ++++++++++++++++++ .../speech_models/pockettts/voice_state.rs | 168 +++++++ mistralrs-pyo3/mistralrs.pyi | 1 + mistralrs-pyo3/src/lib.rs | 1 + mistralrs-pyo3/src/which.rs | 2 + .../src/speech_generation.rs | 14 +- mistralrs/Cargo.toml | 4 + .../examples/models/speech_pockettts/main.rs | 38 ++ mistralrs/src/model_builder_trait.rs | 2 + mistralrs/src/speech_model.rs | 8 + 42 files changed, 4563 insertions(+), 78 deletions(-) create mode 100644 docs/src/content/docs/examples/rust/models/speech-pockettts.md create mode 100644 mistralrs-core/src/speech_models/pockettts/conditioners/mod.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/conditioners/text.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/config.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/mod.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/models/flow_lm.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/models/mimi.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/models/mod.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/models/seanet.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/models/transformer.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/modules/attention.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/modules/conv.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/modules/mlp.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/modules/mod.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/modules/rope.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/modules/sdpa.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/pause.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/tts_model.rs create mode 100644 mistralrs-core/src/speech_models/pockettts/voice_state.rs create mode 100644 mistralrs/examples/models/speech_pockettts/main.rs diff --git a/.gitignore b/.gitignore index eeb50307f5..45c0b25bc6 100644 --- a/.gitignore +++ b/.gitignore @@ -16,3 +16,7 @@ docs/node_modules/ docs/dist/ docs/.astro/ out/ + +dist +dist-build-arm64 +.memsearch diff --git a/docs/src/content/docs/examples/index.md b/docs/src/content/docs/examples/index.md index 7b6935e7fb..b9e86c42e5 100644 --- a/docs/src/content/docs/examples/index.md +++ b/docs/src/content/docs/examples/index.md @@ -9,6 +9,6 @@ Every page in this section is generated from a runnable example in the repositor | Tree | Source | Pages | | --- | --- | --- | -| Rust SDK | [`mistralrs/examples`](https://github.com/EricLBuehler/mistral.rs/blob/master/mistralrs/examples) | 57 | +| Rust SDK | [`mistralrs/examples`](https://github.com/EricLBuehler/mistral.rs/blob/master/mistralrs/examples) | 58 | | Python SDK | [`examples/python`](https://github.com/EricLBuehler/mistral.rs/blob/master/examples/python) | 72 | | HTTP server | [`examples/server`](https://github.com/EricLBuehler/mistral.rs/blob/master/examples/server) | 61 | diff --git a/docs/src/content/docs/examples/rust/models/speech-pockettts.md b/docs/src/content/docs/examples/rust/models/speech-pockettts.md new file mode 100644 index 0000000000..031a9be84b --- /dev/null +++ b/docs/src/content/docs/examples/rust/models/speech-pockettts.md @@ -0,0 +1,55 @@ +--- +title: "CPU-fast text-to-speech with pocket-tts (Kyutai Mimi codec + FlowLM)" +description: "CPU-fast text-to-speech with pocket-tts (Kyutai Mimi codec + FlowLM)" +sidebar: + label: "speech_pockettts" +--- + + + +CPU-fast text-to-speech with pocket-tts (Kyutai Mimi codec + FlowLM). + +Run with: `cargo run --release --example speech_pockettts -p mistralrs` + +```rust +//! CPU-fast text-to-speech with pocket-tts (Kyutai Mimi codec + FlowLM). +//! +//! Run with: `cargo run --release --example speech_pockettts -p mistralrs` + +use std::time::Instant; + +use anyhow::Result; +use mistralrs::{speech_utils, SpeechLoaderType, SpeechModelBuilder}; + +#[tokio::main] +async fn main() -> Result<()> { + let model = SpeechModelBuilder::new( + "kyutai/pocket-tts-without-voice-cloning", + SpeechLoaderType::PocketTts, + ) + .with_logging() + .build() + .await?; + + let start = Instant::now(); + + let text_to_speak = + "Pocket TTS runs on the CPU in seconds, so mistral rs can serve speech without a GPU."; + + let (pcm, rate, channels) = model.generate_speech(text_to_speak).await?; + + let finished = Instant::now(); + + let mut output = std::fs::File::create("out.wav").unwrap(); + speech_utils::write_pcm_as_wav(&mut output, &pcm, rate as u32, channels as u16).unwrap(); + + println!( + "Done! Took {} s. Audio saved at `out.wav`.", + finished.duration_since(start).as_secs_f32(), + ); + + Ok(()) +} +``` + +Source: [`mistralrs/examples/models/speech_pockettts/main.rs`](https://github.com/EricLBuehler/mistral.rs/blob/master/mistralrs/examples/models/speech_pockettts/main.rs) diff --git a/docs/src/content/docs/guides/models/use-speech-models.mdx b/docs/src/content/docs/guides/models/use-speech-models.mdx index 4b6199743b..87b4b20366 100644 --- a/docs/src/content/docs/guides/models/use-speech-models.mdx +++ b/docs/src/content/docs/guides/models/use-speech-models.mdx @@ -1,16 +1,19 @@ --- title: Speech models -description: Voxtral for audio understanding, Dia for text-to-speech. +description: Voxtral for audio understanding, Dia and pocket-tts for text-to-speech. --- import { Tabs, TabItem } from '@astrojs/starlight/components'; -mistral.rs supports two speech-related model families: +mistral.rs supports these speech-related model families: - **Voxtral**: multimodal model accepting audio input. Used for transcription and audio understanding through `/v1/chat/completions`. It uses a Whisper-style audio encoder. -- **Dia**: dedicated text-to-speech model served via `/v1/audio/speech`. +- **Dia**: dedicated text-to-speech model served via `/v1/audio/speech`. GPU-oriented; expressive dialogue. +- **pocket-tts**: dedicated text-to-speech model served via `/v1/audio/speech`. Kyutai-style Mimi codec + FlowLM transformer; runs fast on CPU. -Voxtral is classified as a multimodal model (audio is one of its input modalities); Dia is classified as a dedicated speech model. +Voxtral is classified as a multimodal model (audio is one of its input modalities); Dia and pocket-tts are classified as dedicated speech models. + +Both TTS models share the `/v1/audio/speech` contract, and multiple speech models can be loaded at once (e.g. pocket-tts on CPU and Dia on GPU) and selected per request by the `model` field. ## Voxtral: audio in, text out @@ -174,3 +177,101 @@ speech_utils::write_pcm_as_wav(&mut output, &pcm, rate as u32, channels as u16)? + +## pocket-tts: CPU-fast text-to-speech + +`/v1/audio/speech`, same OpenAI shape as Dia, but runs in seconds on CPU: + +```bash +mistralrs serve -m kyutai/pocket-tts-without-voice-cloning +``` + +The architecture is auto-detected from the model id (or set it explicitly with `--speech-arch pockettts`). pocket-tts speaks with a stock voice; the default is `alba` and you can pick another with `--voice` (stock names: `alba`, `marius`, `javert`, `jean`, `fantine`, `cosette`, `eponine`, `azelma`). + + + + +```bash +curl http://localhost:1234/v1/audio/speech \ + -H "Content-Type: application/json" \ + -d '{ + "model": "default", + "input": "Pocket TTS runs on the CPU in seconds.", + "response_format": "wav" + }' \ + --output out.wav +``` + +- Output: raw audio bytes (24 kHz mono). +- `response_format`: only `wav` and `pcm` are read; other formats return a validation error. +- The voice is fixed at load time (`--voice`); the OpenAI `voice`/`speed` request fields are ignored. + + + + +```python +import struct +import wave +from pathlib import Path + +from mistralrs import Runner, SpeechLoaderType, Which + +runner = Runner( + which=Which.Speech( + model_id="kyutai/pocket-tts-without-voice-cloning", + arch=SpeechLoaderType.PocketTts, + ) +) + +response = runner.generate_audio("Pocket TTS runs on the CPU in seconds.") + +output_path = Path("out.wav") +pcm_ints = [int(max(-32768, min(32767, sample * 32767))) for sample in response.pcm] +with wave.open(output_path, "wb") as wav: + wav.setnchannels(response.channels) + wav.setsampwidth(2) + wav.setframerate(response.rate) + wav.writeframes(b"".join(struct.pack(" + + +```rust +use mistralrs::{speech_utils, SpeechLoaderType, SpeechModelBuilder}; + +let model = SpeechModelBuilder::new( + "kyutai/pocket-tts-without-voice-cloning", + SpeechLoaderType::PocketTts, +) +.with_voice("alba") +.build() +.await?; + +let (pcm, rate, channels) = model + .generate_speech("Pocket TTS runs on the CPU in seconds.") + .await?; + +let mut output = std::fs::File::create("out.wav")?; +speech_utils::write_pcm_as_wav(&mut output, &pcm, rate as u32, channels as u16)?; +``` + + + + +### Serving both TTS models from one binary + +With `mistralrs from-config`, load pocket-tts (CPU) and Dia (GPU) together and select per request by `model`: + +```toml +command = "serve" +default_model_id = "pockettts" + +[[models]] +kind = "speech" +model_id = "kyutai/pocket-tts-without-voice-cloning" +# speech_arch = "pockettts" # auto-detected from model_id +# voice = "alba" +[models.device] +cpu = true +``` diff --git a/docs/src/content/docs/reference/python/enums.md b/docs/src/content/docs/reference/python/enums.md index 358adb48f3..207d02131c 100644 --- a/docs/src/content/docs/reference/python/enums.md +++ b/docs/src/content/docs/reference/python/enums.md @@ -92,6 +92,7 @@ Members and their wire/config names where relevant. The members are fieldless Py | Member | Wire/config name | | --- | --- | | `SpeechLoaderType.Dia` | `'Dia'` | +| `SpeechLoaderType.PocketTts` | `'PocketTts'` | ## `ModelDType` diff --git a/docs/src/content/docs/reference/supported-models.md b/docs/src/content/docs/reference/supported-models.md index a7e00da450..8d266a64a8 100644 --- a/docs/src/content/docs/reference/supported-models.md +++ b/docs/src/content/docs/reference/supported-models.md @@ -91,6 +91,7 @@ The `Architecture` column is the `config.json` `architectures` value. Per-family | Architecture | Model families | Example | |---|---|---| | `Dia` | Dia |
nari-labs/Dia-1.6Bmistralrs run -m nari-labs/Dia-1.6B
| +| `PocketTts` | PocketTts |
kyutai/pocket-tts-without-voice-cloningmistralrs run -m kyutai/pocket-tts-without-voice-cloning
| ## Embedding diff --git a/mistralrs-cli/src/args/mod.rs b/mistralrs-cli/src/args/mod.rs index 982a00b431..6ae2086ef7 100644 --- a/mistralrs-cli/src/args/mod.rs +++ b/mistralrs-cli/src/args/mod.rs @@ -17,7 +17,7 @@ pub use server::*; use clap::{Parser, Subcommand, ValueEnum}; use clap_complete::Shell; -use mistralrs_core::TokenSource; +use mistralrs_core::{SpeechLoaderType, TokenSource}; use serde::Deserialize; use std::path::PathBuf; @@ -396,6 +396,10 @@ fn parse_arch(s: &str) -> Result { s.parse() } +fn parse_speech_arch(s: &str) -> Result { + s.parse() +} + fn parse_dtype(s: &str) -> Result { s.parse() } @@ -488,6 +492,14 @@ pub enum ModelType { #[command(flatten)] device: DeviceOptions, + + /// Speech architecture (`dia` or `pockettts`). Auto-detected from the model id if omitted. + #[arg(long = "speech-arch", value_parser = parse_speech_arch)] + arch: Option, + + /// Speaker voice for pocket-tts (a stock name like `alba`). Ignored by Dia. + #[arg(long)] + voice: Option, }, /// Embedding model diff --git a/mistralrs-cli/src/commands/serve.rs b/mistralrs-cli/src/commands/serve.rs index d96b118edb..c246d3b437 100644 --- a/mistralrs-cli/src/commands/serve.rs +++ b/mistralrs-cli/src/commands/serve.rs @@ -7,7 +7,6 @@ use tracing::{debug, info}; use mistralrs_core::{ initialize_logging, DiffusionLoaderType, McpClientConfig, ModelSelected, PagedCacheType, - SpeechLoaderType, }; use mistralrs_server_core::{ approvals::ApprovalBroker, @@ -24,6 +23,7 @@ use crate::args::{ GlobalOptions, MatformerSelection, ModelFormat, ModelSourceOptions, ModelType, QuantizationOptions, RuntimeOptions, SandboxMode, SandboxOptions, ServerOptions, }; +use crate::config::detect_speech_arch; use crate::ui::build_ui_router; /// Run the HTTP server with the specified model @@ -402,10 +402,16 @@ pub(crate) fn convert_to_model_selected( dtype: model.dtype, }), - ModelType::Speech { model, device: _ } => Ok(ModelSelected::Speech { + ModelType::Speech { + model, + device: _, + arch, + voice, + } => Ok(ModelSelected::Speech { model_id: model.model_id.clone(), dac_model_id: None, - arch: SpeechLoaderType::Dia, + arch: arch.unwrap_or_else(|| detect_speech_arch(&model.model_id)), + voice: voice.clone(), dtype: model.dtype, }), diff --git a/mistralrs-cli/src/config/mod.rs b/mistralrs-cli/src/config/mod.rs index 598fe55e66..399f4a8576 100644 --- a/mistralrs-cli/src/config/mod.rs +++ b/mistralrs-cli/src/config/mod.rs @@ -12,7 +12,7 @@ use crate::args::{ ModelType, MultimodalOptions, PagedAttentionOptions, QuantizationOptions, RuntimeOptions, SandboxOptions, ServerOptions, }; -use mistralrs_core::{ModelDType, NormalLoaderType, TokenSource}; +use mistralrs_core::{ModelDType, NormalLoaderType, SpeechLoaderType, TokenSource}; #[derive(Deserialize)] #[serde(tag = "command", rename_all = "kebab-case")] @@ -86,6 +86,13 @@ pub struct ModelEntry { pub tokenizer: Option, #[serde(default)] pub arch: Option, + /// Speech architecture (`dia` or `pockettts`). Only meaningful for `kind = "speech"`. + /// Auto-detected from `model_id` when omitted. + #[serde(default)] + pub speech_arch: Option, + /// Speaker voice for pocket-tts (a stock name like `alba`). Only meaningful for `kind = "speech"`. + #[serde(default)] + pub voice: Option, #[serde(default)] pub dtype: ModelDType, #[serde(default)] @@ -145,6 +152,14 @@ pub fn load_cli_config(path: &Path) -> Result { Ok(config) } +pub(crate) fn detect_speech_arch(model_id: &str) -> SpeechLoaderType { + if model_id.to_lowercase().contains("pocket") { + SpeechLoaderType::PocketTts + } else { + SpeechLoaderType::Dia + } +} + fn validate_config(config: &CliConfig) -> Result<()> { let (models, default_model_id) = match config { CliConfig::Serve(cfg) => (&cfg.models, cfg.default_model_id.as_ref()), @@ -255,7 +270,15 @@ impl ModelEntry { multimodal: self.multimodal.clone(), }, ModelKind::Diffusion => ModelType::Diffusion { model, device }, - ModelKind::Speech => ModelType::Speech { model, device }, + ModelKind::Speech => ModelType::Speech { + arch: Some( + self.speech_arch + .unwrap_or_else(|| detect_speech_arch(&self.model_id)), + ), + voice: self.voice.clone(), + model, + device, + }, ModelKind::Embedding => ModelType::Embedding { model, format: self.format.clone(), diff --git a/mistralrs-core/src/model_loader.rs b/mistralrs-core/src/model_loader.rs index 9aa8ca03ad..2f5f293d35 100644 --- a/mistralrs-core/src/model_loader.rs +++ b/mistralrs-core/src/model_loader.rs @@ -403,12 +403,14 @@ fn loader_from_model_selected(args: LoaderBuilder) -> anyhow::Result Box::new(SpeechLoader { model_id, dac_model_id, arch, cfg: None, + voice, }), ModelSelected::XLora { model_id, diff --git a/mistralrs-core/src/model_metadata.rs b/mistralrs-core/src/model_metadata.rs index 83e2ba9f74..6e19b2874a 100644 --- a/mistralrs-core/src/model_metadata.rs +++ b/mistralrs-core/src/model_metadata.rs @@ -363,6 +363,11 @@ impl SpeechLoaderType { modalities: &[Text, Audio], examples: &[ex!("nari-labs/Dia-1.6B")], }, + Self::PocketTts => ArchMetadata { + families: &["PocketTts"], + modalities: &[Text, Audio], + examples: &[ex!("kyutai/pocket-tts-without-voice-cloning")], + }, } } } diff --git a/mistralrs-core/src/model_selected.rs b/mistralrs-core/src/model_selected.rs index bac2f6a377..55a76ce36b 100644 --- a/mistralrs-core/src/model_selected.rs +++ b/mistralrs-core/src/model_selected.rs @@ -688,6 +688,10 @@ pub enum ModelSelected { #[arg(short, long, value_parser = parse_speech_arch)] arch: SpeechLoaderType, + /// Speaker voice for pocket-tts (a stock name like `alba`). Ignored by Dia. + #[arg(long)] + voice: Option, + /// Model data type. Defaults to `auto`. #[arg(long, default_value_t = ModelDType::Auto, value_parser = parse_model_dtype)] dtype: ModelDType, diff --git a/mistralrs-core/src/pipeline/auto.rs b/mistralrs-core/src/pipeline/auto.rs index 50124d3ca2..80c78f8ce8 100644 --- a/mistralrs-core/src/pipeline/auto.rs +++ b/mistralrs-core/src/pipeline/auto.rs @@ -343,6 +343,13 @@ impl AutoLoader { } } + // Speech models that ship no `config.json` (e.g. pocket-tts) are detected by their files. + if let Some(tp) = + crate::speech_models::SpeechLoaderType::auto_detect_from_files(&artifacts.repo_files) + { + return Ok(Detected::Speech(tp)); + } + if artifacts.sentence_transformers_present { if let Some(ref config) = artifacts.contents { let cfg: AutoConfig = serde_json::from_str(config)?; @@ -440,6 +447,7 @@ impl AutoLoader { dac_model_id: None, arch: tp, cfg: None, + voice: None, }); *guard = Some(loader); } diff --git a/mistralrs-core/src/pipeline/speech.rs b/mistralrs-core/src/pipeline/speech.rs index 41f9bccebb..0597fab91d 100644 --- a/mistralrs-core/src/pipeline/speech.rs +++ b/mistralrs-core/src/pipeline/speech.rs @@ -10,7 +10,10 @@ use crate::distributed::{use_ring, WorkerTransferData}; use crate::pipeline::{ChatTemplate, EmbeddingModulePaths, Modalities, SupportedModality}; use crate::prefix_cacher::PrefixCacheManagerV2; use crate::sequence::Sequence; -use crate::speech_models::{DiaConfig, DiaPipeline, SpeechGenerationOutput, SpeechLoaderType}; +use crate::speech_models::{ + DiaConfig, DiaPipeline, PocketTtsConfig, PocketTtsPipeline, SpeechGenerationOutput, + SpeechLoaderType, POCKETTTS_WEIGHTS_FILE, +}; use crate::utils::progress::ProgressScopeGuard; use crate::utils::varbuilder_utils::DeviceForLoadTensor; use crate::utils::{tokens::get_token, varbuilder_utils::from_mmaped_safetensors}; @@ -33,10 +36,40 @@ use std::sync::Arc; use tokenizers::Tokenizer; use tokio::sync::Mutex; +const POCKETTTS_TOKENIZER_FILE: &str = "tokenizer.model"; +const POCKETTTS_DEFAULT_VOICE: &str = "alba"; + #[derive(Clone, Debug)] pub struct SpeechModelPaths { weights: Vec, config: PathBuf, + tokenizer: Option, + voice_prompt: Option, +} + +enum SpeechModel { + Dia(DiaPipeline), + PocketTts(PocketTtsPipeline), +} + +impl SpeechModel { + fn generate( + &self, + prompt: &str, + cfg: &SpeechGenerationConfig, + ) -> candle_core::Result { + match self { + Self::Dia(m) => m.generate(prompt, cfg), + Self::PocketTts(m) => m.generate(prompt, cfg), + } + } + + fn device(&self) -> &Device { + match self { + Self::Dia(m) => m.device(), + Self::PocketTts(m) => m.device(), + } + } } impl ModelPaths for SpeechModelPaths { @@ -143,7 +176,7 @@ impl InputsProcessor for SpeechInputsProcessor { pub struct SpeechPipeline { model_id: String, - model: DiaPipeline, + model: SpeechModel, metadata: Arc, dummy_cache: EitherCache, cfg: SpeechGenerationConfig, @@ -154,6 +187,8 @@ pub struct SpeechLoader { pub dac_model_id: Option, pub arch: SpeechLoaderType, pub cfg: Option, + /// Speaker voice for pocket-tts (a stock name like `alba`). Ignored by Dia. + pub voice: Option, } impl Loader for SpeechLoader { @@ -170,17 +205,70 @@ impl Loader for SpeechLoader { paged_attn_config: Option, ) -> Result>> { let _progress_guard = ProgressScopeGuard::new(silent); - let paths: anyhow::Result> = { - // Main weights first, DAC is the final one. - let mut weights = Vec::new(); - - // Main model - let config = { + let paths: anyhow::Result> = match self.arch { + SpeechLoaderType::Dia => { + // Main weights first, DAC is the final one. + let mut weights = Vec::new(); + + // Main model + let config = { + let api = ApiBuilder::new() + .with_progress(!silent) + .with_token(get_token(&token_source)?) + .build()?; + let revision = revision.clone().unwrap_or("main".to_string()); + let api = api.repo(Repo::with_revision( + self.model_id.to_string(), + RepoType::Model, + revision.clone(), + )); + let model_id = std::path::Path::new(&self.model_id); + + let weight = api_get_file!(api, "model.safetensors", &model_id, &revision); + let config = api_get_file!(api, "config.json", &model_id, &revision); + weights.push(weight); + config + }; + + // DAC model + { + let api = ApiBuilder::new() + .with_progress(!silent) + .with_token(get_token(&token_source)?) + .build()?; + let revision = revision.unwrap_or("main".to_string()); + + let dac_model = self + .dac_model_id + .clone() + .unwrap_or_else(|| "EricB/dac_44khz".to_string()); + + let api = api.repo(Repo::with_revision( + dac_model.clone(), + RepoType::Model, + revision.clone(), + )); + let model_id = std::path::Path::new(&dac_model); + + let weight = api_get_file!(api, "model.safetensors", &model_id, &revision); + weights.push(weight); + } + + Ok(Box::new(SpeechModelPaths { + weights, + config, + tokenizer: None, + voice_prompt: None, + })) + } + SpeechLoaderType::PocketTts => { + // pocket-tts ships weights + a SentencePiece tokenizer + a set of stock speaker + // embeddings under `embeddings/.safetensors`; the config is baked in. let api = ApiBuilder::new() .with_progress(!silent) .with_token(get_token(&token_source)?) .build()?; - let revision = revision.clone().unwrap_or("main".to_string()); + let revision = revision.unwrap_or("main".to_string()); let api = api.repo(Repo::with_revision( self.model_id.to_string(), RepoType::Model, @@ -188,40 +276,20 @@ impl Loader for SpeechLoader { )); let model_id = std::path::Path::new(&self.model_id); - let weight = api_get_file!(api, "model.safetensors", &model_id, &revision); - let config = api_get_file!(api, "config.json", &model_id, &revision); - weights.push(weight); - config - }; + let weight = api_get_file!(api, POCKETTTS_WEIGHTS_FILE, &model_id, &revision); + let tokenizer = api_get_file!(api, POCKETTTS_TOKENIZER_FILE, &model_id, &revision); - // DAC model - { - let api = ApiBuilder::new() - .with_progress(!silent) - .with_token(get_token(&token_source)?) - .build()?; - let revision = revision.unwrap_or("main".to_string()); - - // Apply default here - let dac_model = self - .dac_model_id - .clone() - .unwrap_or_else(|| match self.arch { - SpeechLoaderType::Dia => "EricB/dac_44khz".to_string(), - }); + let voice = self.voice.as_deref().unwrap_or(POCKETTTS_DEFAULT_VOICE); + let voice_file = format!("embeddings/{voice}.safetensors"); + let voice_prompt = api_get_file!(api, &voice_file, &model_id, &revision); - let api = api.repo(Repo::with_revision( - dac_model.clone(), - RepoType::Model, - revision.clone(), - )); - let model_id = std::path::Path::new(&dac_model); - - let weight = api_get_file!(api, "model.safetensors", &model_id, &revision); - weights.push(weight); + Ok(Box::new(SpeechModelPaths { + weights: vec![weight.clone()], + config: weight, + tokenizer: Some(tokenizer), + voice_prompt: Some(voice_prompt), + })) } - - Ok(Box::new(SpeechModelPaths { weights, config })) }; self.load_model_from_path( &paths?, @@ -262,8 +330,6 @@ impl Loader for SpeechLoader { mistralrs_quant::IsqCaptureMode::Immediate, ); - let cfg: DiaConfig = serde_json::from_str(&std::fs::read_to_string(&paths.config)?)?; - #[cfg(feature = "cuda")] if let Device::Cuda(dev) = &device { unsafe { dev.disable_event_tracking() }; @@ -283,29 +349,56 @@ impl Loader for SpeechLoader { DeviceMapSetting::dummy().into_mapper(usize::MAX, device, None, &available_devices)?; let dtype = mapper.get_min_dtype(dtype)?; - // Last weight is the dac. - let model_weights = paths.weights[..paths.weights.len() - 1].to_vec(); - let vb = from_mmaped_safetensors( - model_weights, - Vec::new(), - Some(dtype), - device, - vec![None], - silent, - None, - |_| true, - Arc::new(|_| DeviceForLoadTensor::Base), - )?; - - let dac_vb = unsafe { - VarBuilder::from_mmaped_safetensors(&[paths.weights.last().unwrap()], dtype, device)? + let model = match self.arch { + SpeechLoaderType::Dia => { + let cfg: DiaConfig = + serde_json::from_str(&std::fs::read_to_string(&paths.config)?)?; + + // Last weight is the dac. + let model_weights = paths.weights[..paths.weights.len() - 1].to_vec(); + let vb = from_mmaped_safetensors( + model_weights, + Vec::new(), + Some(dtype), + device, + vec![None], + silent, + None, + |_| true, + Arc::new(|_| DeviceForLoadTensor::Base), + )?; + + let dac_vb = unsafe { + VarBuilder::from_mmaped_safetensors( + &[paths.weights.last().unwrap()], + dtype, + device, + )? + }; + + SpeechModel::Dia(DiaPipeline::new(&cfg, vb, dac_vb)?) + } + SpeechLoaderType::PocketTts => { + let cfg = PocketTtsConfig::b6369a24(); + let tokenizer = paths + .tokenizer + .as_ref() + .expect("pocket-tts requires a tokenizer path"); + let voice_prompt = paths + .voice_prompt + .as_ref() + .expect("pocket-tts requires a voice prompt path"); + let vb = unsafe { + VarBuilder::from_mmaped_safetensors( + &paths.weights, + candle_core::DType::F32, + device, + )? + }; + SpeechModel::PocketTts(PocketTtsPipeline::new(&cfg, vb, tokenizer, voice_prompt)?) + } }; - // Only Dia is supported for now. - assert_eq!(self.arch, SpeechLoaderType::Dia); - - let model = DiaPipeline::new(&cfg, vb, dac_vb)?; - Ok(Arc::new(Mutex::new(SpeechPipeline { model_id: self.model_id.clone(), model, diff --git a/mistralrs-core/src/speech_models/dia/mod.rs b/mistralrs-core/src/speech_models/dia/mod.rs index 0e087ef33e..ebbde94c43 100644 --- a/mistralrs-core/src/speech_models/dia/mod.rs +++ b/mistralrs-core/src/speech_models/dia/mod.rs @@ -388,7 +388,10 @@ impl DiaPipeline { temperature, top_p, top_k, - } = cfg; + } = cfg + else { + unreachable!("DiaPipeline requires a Dia speech generation config"); + }; let audio_pad_value = self.cfg.data.audio_pad_value as u32; let audio_eos_value = self.cfg.data.audio_eos_value as u32; diff --git a/mistralrs-core/src/speech_models/mod.rs b/mistralrs-core/src/speech_models/mod.rs index fd821de6b7..35d7c9c529 100644 --- a/mistralrs-core/src/speech_models/mod.rs +++ b/mistralrs-core/src/speech_models/mod.rs @@ -1,16 +1,20 @@ mod bs1770; mod dia; +mod pockettts; pub mod utils; use std::{str::FromStr, sync::Arc}; pub use dia::{DiaConfig, DiaPipeline}; +pub use pockettts::{PocketTtsConfig, PocketTtsPipeline}; use serde::{Deserialize, Serialize}; #[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, strum::EnumIter)] pub enum SpeechLoaderType { #[serde(rename = "dia")] Dia, + #[serde(rename = "pockettts")] + PocketTts, } impl FromStr for SpeechLoaderType { @@ -18,13 +22,17 @@ impl FromStr for SpeechLoaderType { fn from_str(s: &str) -> Result { match s { "dia" => Ok(Self::Dia), + "pockettts" => Ok(Self::PocketTts), a => Err(format!( - "Unknown architecture `{a}`. Possible architectures: `dia`." + "Unknown architecture `{a}`. Possible architectures: `dia`, `pockettts`." )), } } } +/// Marker file that identifies a pocket-tts repo (which ships no `config.json`). +pub const POCKETTTS_WEIGHTS_FILE: &str = "tts_b6369a24.safetensors"; + impl SpeechLoaderType { /// Auto-detect speech loader type from a config.json string. /// Extend this when adding new speech pipelines. @@ -34,6 +42,18 @@ impl SpeechLoaderType { } None } + + /// Auto-detect speech loader type from the repo file list, for models that ship no + /// `config.json` (e.g. pocket-tts). Extend this when adding new such pipelines. + pub fn auto_detect_from_files(files: &[String]) -> Option { + if files + .iter() + .any(|f| f.rsplit('/').next() == Some(POCKETTTS_WEIGHTS_FILE)) + { + return Some(Self::PocketTts); + } + None + } } #[derive(Clone, Copy, Debug)] @@ -45,6 +65,11 @@ pub enum SpeechGenerationConfig { top_p: f32, top_k: Option, }, + PocketTts { + temperature: f32, + lsd_decode_steps: usize, + eos_threshold: f32, + }, } impl SpeechGenerationConfig { @@ -57,6 +82,11 @@ impl SpeechGenerationConfig { top_p: 0.95, top_k: Some(35), }, + SpeechLoaderType::PocketTts => Self::PocketTts { + temperature: 0.7, + lsd_decode_steps: 1, + eos_threshold: -4.0, + }, } } } diff --git a/mistralrs-core/src/speech_models/pockettts/conditioners/mod.rs b/mistralrs-core/src/speech_models/pockettts/conditioners/mod.rs new file mode 100644 index 0000000000..481c63accf --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/conditioners/mod.rs @@ -0,0 +1 @@ +pub mod text; diff --git a/mistralrs-core/src/speech_models/pockettts/conditioners/text.rs b/mistralrs-core/src/speech_models/pockettts/conditioners/text.rs new file mode 100644 index 0000000000..826be518b3 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/conditioners/text.rs @@ -0,0 +1,390 @@ +use candle_core::Tensor; +use candle_nn::{Embedding, Module, VarBuilder}; + +// Use tokenizers crate for all platforms (no protobuf dependency) +use tokenizers::Tokenizer; + +use anyhow::Result; +use std::path::Path; + +use std::sync::Arc; + +#[derive(Clone)] +pub struct LUTConditioner { + tokenizer: Arc, + embed: Embedding, +} + +impl LUTConditioner { + pub fn new( + n_bins: usize, + tokenizer_path: &Path, + dim: usize, + _output_dim: usize, + vb: VarBuilder, + ) -> Result { + // Load SentencePiece model using tokenizers crate + // The tokenizers crate can load .model files directly via from_file + // For .model files, we need to use the unigram model loader + let tokenizer = if tokenizer_path.extension().is_some_and(|e| e == "model") { + // SentencePiece .model file - use unigram loader + Self::load_sentencepiece_model(tokenizer_path)? + } else { + // JSON tokenizer file + Tokenizer::from_file(tokenizer_path) + .map_err(|e| anyhow::anyhow!("Failed to load tokenizer from file: {:?}", e))? + }; + + // Verify vocab size matches + let vocab_size = tokenizer.get_vocab_size(true); + if vocab_size != n_bins { + anyhow::bail!( + "Tokenizer vocab size {} doesn't match n_bins {}", + vocab_size, + n_bins + ); + } + + // n_bins + 1 for padding + let embed = candle_nn::embedding(n_bins + 1, dim, vb.pp("embed"))?; + + Ok(Self { + tokenizer: Arc::new(tokenizer), + embed, + }) + } + + /// Load a SentencePiece .model file using tokenizers crate + fn load_sentencepiece_model(path: &Path) -> Result { + use tokenizers::models::unigram::Unigram; + use tokenizers::pre_tokenizers::metaspace::{Metaspace, PrependScheme}; + + // Read the protobuf file and extract vocab manually + // The tokenizers crate's Unigram model can be built from vocab + let model_bytes = + std::fs::read(path).map_err(|e| anyhow::anyhow!("Failed to read model file: {}", e))?; + + // Parse SentencePiece model protobuf to extract vocab + let (vocab, unk_id) = Self::parse_sentencepiece_vocab(&model_bytes)?; + + // Build Unigram model from vocab + let unigram = Unigram::from(vocab, Some(unk_id), true) + .map_err(|e| anyhow::anyhow!("Failed to create unigram model: {:?}", e))?; + + // Build tokenizer with SentencePiece-style settings + let mut tokenizer = Tokenizer::new(unigram); + tokenizer.with_pre_tokenizer(Some(Metaspace::new('▁', PrependScheme::Always, false))); + tokenizer.with_decoder(Some(Metaspace::new('▁', PrependScheme::Always, false))); + + Ok(tokenizer) + } + + /// Parse SentencePiece model protobuf to extract vocabulary + /// SentencePiece uses a simple protobuf format we can parse manually + fn parse_sentencepiece_vocab(data: &[u8]) -> Result<(Vec<(String, f64)>, usize)> { + // SentencePiece protobuf structure (simplified): + // message ModelProto { + // repeated SentencePiece pieces = 1; + // ... + // } + // message SentencePiece { + // optional string piece = 1; + // optional float score = 2; + // ... + // } + // + // We parse field 1 (pieces) which contains repeated messages with piece (field 1) and score (field 2) + + let mut vocab = Vec::new(); + let mut unk_id = 0usize; + let mut pos = 0; + + while pos < data.len() { + // Read field tag + let (tag, new_pos) = Self::read_varint(data, pos)?; + pos = new_pos; + + let field_number = tag >> 3; + let wire_type = tag & 0x7; + + match (field_number, wire_type) { + (1, 2) => { + // Field 1 (pieces), wire type 2 (length-delimited) - this is a SentencePiece message + let (len, new_pos) = Self::read_varint(data, pos)?; + pos = new_pos; + let end = pos + len as usize; + + // Parse the nested SentencePiece message + let mut piece = String::new(); + let mut score = 0.0f64; + let mut inner_pos = pos; + + while inner_pos < end { + let (inner_tag, new_inner_pos) = Self::read_varint(data, inner_pos)?; + inner_pos = new_inner_pos; + + let inner_field = inner_tag >> 3; + let inner_wire = inner_tag & 0x7; + + match (inner_field, inner_wire) { + (1, 2) => { + // piece string + let (len, new_pos) = Self::read_varint(data, inner_pos)?; + inner_pos = new_pos; + piece = String::from_utf8_lossy( + &data[inner_pos..inner_pos + len as usize], + ) + .to_string(); + inner_pos += len as usize; + } + (2, 5) => { + // score (float, wire type 5 = 32-bit) + if inner_pos + 4 <= data.len() { + let bytes: [u8; 4] = + data[inner_pos..inner_pos + 4].try_into().unwrap(); + score = f32::from_le_bytes(bytes) as f64; + inner_pos += 4; + } + } + (3, 0) => { + // type (varint) + let (type_val, new_pos) = Self::read_varint(data, inner_pos)?; + inner_pos = new_pos; + // type 2 = UNKNOWN + if type_val == 2 { + unk_id = vocab.len(); + } + } + (_, 0) => { + // Other varint field - skip + let (_, new_pos) = Self::read_varint(data, inner_pos)?; + inner_pos = new_pos; + } + (_, 2) => { + // Other length-delimited field - skip + let (len, new_pos) = Self::read_varint(data, inner_pos)?; + inner_pos = new_pos + len as usize; + } + (_, 5) => { + // 32-bit field - skip + inner_pos += 4; + } + (_, 1) => { + // 64-bit field - skip + inner_pos += 8; + } + _ => { + // Unknown wire type - try to skip + inner_pos = end; + } + } + } + + if !piece.is_empty() { + vocab.push((piece, score)); + } + pos = end; + } + (_, 0) => { + // Varint - skip + let (_, new_pos) = Self::read_varint(data, pos)?; + pos = new_pos; + } + (_, 2) => { + // Length-delimited - skip + let (len, new_pos) = Self::read_varint(data, pos)?; + pos = new_pos + len as usize; + } + (_, 5) => { + // 32-bit - skip + pos += 4; + } + (_, 1) => { + // 64-bit - skip + pos += 8; + } + _ => { + break; // Unknown wire type + } + } + } + + if vocab.is_empty() { + anyhow::bail!("No vocabulary found in SentencePiece model"); + } + + Ok((vocab, unk_id)) + } + + /// Read a varint from the buffer + fn read_varint(data: &[u8], mut pos: usize) -> Result<(u64, usize)> { + let mut result = 0u64; + let mut shift = 0; + + loop { + if pos >= data.len() { + anyhow::bail!("Unexpected end of data while reading varint"); + } + let byte = data[pos]; + pos += 1; + result |= ((byte & 0x7F) as u64) << shift; + if byte & 0x80 == 0 { + break; + } + shift += 7; + if shift >= 64 { + anyhow::bail!("Varint too large"); + } + } + + Ok((result, pos)) + } + + /// Create LUTConditioner from pre-loaded tokenizer bytes (useful for WASM) + pub fn new_from_bytes( + n_bins: usize, + tokenizer_bytes: &[u8], + dim: usize, + _output_dim: usize, + vb: VarBuilder, + ) -> Result { + // Try to parse as JSON tokenizer first + let tokenizer = if let Ok(t) = Tokenizer::from_bytes(tokenizer_bytes) { + t + } else { + // Try as SentencePiece model + let (vocab, unk_id) = Self::parse_sentencepiece_vocab(tokenizer_bytes)?; + + use tokenizers::models::unigram::Unigram; + use tokenizers::pre_tokenizers::metaspace::{Metaspace, PrependScheme}; + + let unigram = Unigram::from(vocab, Some(unk_id), true) + .map_err(|e| anyhow::anyhow!("Failed to create unigram model: {:?}", e))?; + + let mut tok = Tokenizer::new(unigram); + tok.with_pre_tokenizer(Some(Metaspace::new('▁', PrependScheme::Always, false))); + tok.with_decoder(Some(Metaspace::new('▁', PrependScheme::Always, false))); + tok + }; + + // n_bins + 1 for padding + let embed = candle_nn::embedding(n_bins + 1, dim, vb.pp("embed"))?; + + Ok(Self { + tokenizer: Arc::new(tokenizer), + embed, + }) + } + + pub fn prepare(&self, text: &str, device: &candle_core::Device) -> Result { + let encoding = self + .tokenizer + .encode(text, true) + .map_err(|e| anyhow::anyhow!("Failed to encode text: {:?}", e))?; + + let ids = encoding.get_ids(); + Ok(Tensor::from_vec(ids.to_vec(), (1, ids.len()), device)?) + } + + pub fn forward(&self, tokens: &Tensor) -> Result { + // Handle empty token tensors (e.g., shape [1, 0]) which cause Metal kernel issues + // The embedding dimension is the hidden size of the embed layer + let dims = tokens.dims(); + if dims.len() >= 2 && dims[1] == 0 { + // Return empty embeddings with correct shape [batch, 0, embed_dim] + let embed_dim = self.embed.embeddings().dims()[1]; + return Ok(Tensor::zeros( + (dims[0], 0, embed_dim), + candle_core::DType::F32, + tokens.device(), + )?); + } + Ok(self.embed.forward(tokens)?) + } + + /// Count tokens in a text string without creating tensors. + /// Used for accurate text splitting to avoid oversized chunks. + pub fn count_tokens(&self, text: &str) -> Result { + let encoding = self + .tokenizer + .encode(text, true) + .map_err(|e| anyhow::anyhow!("Failed to encode text: {:?}", e))?; + Ok(encoding.get_ids().len()) + } +} + +#[cfg(test)] +mod tests { + use super::LUTConditioner; + + fn encode_varint(mut value: u64) -> Vec { + let mut out = Vec::new(); + loop { + if value < 0x80 { + out.push(value as u8); + return out; + } + out.push(((value as u8) & 0x7f) | 0x80); + value >>= 7; + } + } + + fn encode_piece(piece: &str, score: f32, piece_type: Option) -> Vec { + let mut msg = Vec::new(); + + // field 1: piece (string) + msg.push(0x0a); + msg.extend(encode_varint(piece.len() as u64)); + msg.extend(piece.as_bytes()); + + // field 2: score (float, wire type 5) + msg.push(0x15); + msg.extend(score.to_le_bytes()); + + // field 3: type (varint) + if let Some(piece_type) = piece_type { + msg.push(0x18); + msg.extend(encode_varint(piece_type)); + } + + let mut outer = Vec::new(); + // outer field 1: repeated SentencePiece message + outer.push(0x0a); + outer.extend(encode_varint(msg.len() as u64)); + outer.extend(msg); + outer + } + + #[test] + fn test_read_varint_multibyte() { + let data = [0xac, 0x02, 0x01]; + let (first, pos) = LUTConditioner::read_varint(&data, 0).expect("first varint"); + let (second, end) = LUTConditioner::read_varint(&data, pos).expect("second varint"); + assert_eq!(first, 300); + assert_eq!(second, 1); + assert_eq!(end, data.len()); + } + + #[test] + fn test_parse_sentencepiece_vocab_extracts_pieces_and_unk() { + let mut model = Vec::new(); + model.extend(encode_piece("", -1.0, Some(2))); + model.extend(encode_piece("hello", -2.5, Some(1))); + + let (vocab, unk_id) = + LUTConditioner::parse_sentencepiece_vocab(&model).expect("parse sentencepiece vocab"); + + assert_eq!(unk_id, 0); + assert_eq!(vocab.len(), 2); + assert_eq!(vocab[0].0, ""); + assert_eq!(vocab[1].0, "hello"); + assert!((vocab[0].1 + 1.0).abs() < 1e-6); + assert!((vocab[1].1 + 2.5).abs() < 1e-6); + } + + #[test] + fn test_parse_sentencepiece_vocab_rejects_empty_vocab() { + let err = LUTConditioner::parse_sentencepiece_vocab(&[]).expect_err("expected empty error"); + assert!(err.to_string().contains("No vocabulary found")); + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/config.rs b/mistralrs-core/src/speech_models/pockettts/config.rs new file mode 100644 index 0000000000..ac4a8a699c --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/config.rs @@ -0,0 +1,141 @@ +use serde::Deserialize; + +#[derive(Debug, Clone, Deserialize)] +pub struct FlowConfig { + pub dim: usize, + pub depth: usize, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct FlowLMTransformerConfig { + pub hidden_scale: usize, + pub max_period: usize, + pub d_model: usize, + pub num_heads: usize, + pub num_layers: usize, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct LookupTableConfig { + pub dim: usize, + pub n_bins: usize, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct FlowLMConfig { + pub flow: FlowConfig, + pub transformer: FlowLMTransformerConfig, + pub lookup_table: LookupTableConfig, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct SEANetConfig { + pub dimension: usize, + pub channels: usize, + pub n_filters: usize, + pub n_residual_layers: usize, + pub ratios: Vec, + pub kernel_size: usize, + pub residual_kernel_size: usize, + pub last_kernel_size: usize, + pub dilation_base: usize, + pub pad_mode: String, + pub compress: usize, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct MimiTransformerConfig { + pub d_model: usize, + pub input_dimension: usize, + pub output_dimensions: Vec, + pub num_heads: usize, + pub num_layers: usize, + pub layer_scale: f64, + pub context: usize, + pub max_period: f64, + pub dim_feedforward: usize, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct QuantizerConfig { + pub dimension: usize, + pub output_dimension: usize, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct MimiConfig { + pub sample_rate: usize, + pub channels: usize, + pub frame_rate: f64, + pub seanet: SEANetConfig, + pub transformer: MimiTransformerConfig, + pub quantizer: QuantizerConfig, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct PocketTtsConfig { + pub flow_lm: FlowLMConfig, + pub mimi: MimiConfig, +} + +impl PocketTtsConfig { + /// The single published pocket-tts variant `b6369a24`. Values mirror `b6369a24.yaml` in the + /// upstream crate; the HF repo ships only weights + tokenizer, no `config.json`. + pub fn b6369a24() -> Self { + Self { + flow_lm: FlowLMConfig { + flow: FlowConfig { dim: 512, depth: 6 }, + transformer: FlowLMTransformerConfig { + hidden_scale: 4, + max_period: 10000, + d_model: 1024, + num_heads: 16, + num_layers: 6, + }, + lookup_table: LookupTableConfig { + dim: 1024, + n_bins: 4000, + }, + }, + mimi: MimiConfig { + sample_rate: 24000, + channels: 1, + frame_rate: 12.5, + seanet: SEANetConfig { + dimension: 512, + channels: 1, + n_filters: 64, + n_residual_layers: 1, + ratios: vec![6, 5, 4], + kernel_size: 7, + residual_kernel_size: 3, + last_kernel_size: 3, + dilation_base: 2, + pad_mode: "constant".to_string(), + compress: 2, + }, + transformer: MimiTransformerConfig { + d_model: 512, + input_dimension: 512, + output_dimensions: vec![512], + num_heads: 8, + num_layers: 2, + layer_scale: 0.01, + context: 250, + max_period: 10000.0, + dim_feedforward: 2048, + }, + quantizer: QuantizerConfig { + dimension: 32, + output_dimension: 512, + }, + }, + } + } +} + +pub mod defaults { + pub const TEMPERATURE: f32 = 0.7; + pub const LSD_DECODE_STEPS: usize = 1; + pub const EOS_THRESHOLD: f32 = -4.0; +} diff --git a/mistralrs-core/src/speech_models/pockettts/mod.rs b/mistralrs-core/src/speech_models/pockettts/mod.rs new file mode 100644 index 0000000000..cbadb34b0f --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/mod.rs @@ -0,0 +1,84 @@ +// Model + inference code ported near-verbatim from the pocket-tts Rust crate +// (https://github.com/babybirdprd/pocket-tts), compiled against candle 0.11 with `crate::` paths +// rewritten. The raw-audio voice-cloning path (audio encoder / pause helpers) is retained but +// unused here, so allow the dead code. +#![allow(dead_code)] + +mod conditioners; +mod config; +mod models; +mod modules; +mod pause; +mod tts_model; +mod voice_state; + +use std::path::Path; +use std::sync::Arc; + +use candle_core::Device; +use candle_nn::VarBuilder; + +pub use config::PocketTtsConfig; +use tts_model::TTSModel; +use voice_state::ModelState; + +use super::{SpeechGenerationConfig, SpeechGenerationOutput}; + +pub struct PocketTtsPipeline { + model: TTSModel, + voice_state: ModelState, + device: Device, +} + +impl PocketTtsPipeline { + pub fn new( + cfg: &PocketTtsConfig, + vb: VarBuilder, + tokenizer_path: &Path, + voice_prompt_path: &Path, + ) -> candle_core::Result { + let device = vb.device().clone(); + let model = TTSModel::new(cfg, vb, tokenizer_path).map_err(candle_core::Error::wrap)?; + let voice_state = model + .voice_state_from_prompt_file(voice_prompt_path) + .map_err(candle_core::Error::wrap)?; + Ok(Self { + model, + voice_state, + device, + }) + } + + pub fn device(&self) -> &Device { + &self.device + } + + pub fn generate( + &self, + prompt: &str, + cfg: &SpeechGenerationConfig, + ) -> candle_core::Result { + let mut model = self.model.clone(); + if let SpeechGenerationConfig::PocketTts { + temperature, + lsd_decode_steps, + eos_threshold, + } = cfg + { + model.temp = *temperature; + model.lsd_decode_steps = *lsd_decode_steps; + model.eos_threshold = *eos_threshold; + } + + let audio = model + .generate(prompt, &self.voice_state) + .map_err(candle_core::Error::wrap)?; + let pcm = audio.flatten_all()?.to_vec1::()?; + + Ok(SpeechGenerationOutput { + pcm: Arc::new(pcm), + rate: model.sample_rate, + channels: self.model.mimi.channels, + }) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/models/flow_lm.rs b/mistralrs-core/src/speech_models/pockettts/models/flow_lm.rs new file mode 100644 index 0000000000..69eb882f5e --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/models/flow_lm.rs @@ -0,0 +1,151 @@ +use crate::speech_models::pockettts::models::transformer::StreamingTransformer; +use crate::speech_models::pockettts::modules::mlp::{LayerNorm, ModulationParams, SimpleMLPAdaLN}; +use crate::speech_models::pockettts::voice_state::ModelState; +use candle_core::{Result, Tensor}; +use candle_nn::{Linear, Module, VarBuilder}; + +pub fn lsd_decode( + flow_net: &SimpleMLPAdaLN, + modulations: &[Vec], + x_0: &Tensor, +) -> Result { + let mut current = x_0.clone(); + let num_steps = modulations.len(); + + let step_factor = 1.0 / num_steps as f64; + for step_mod in modulations { + // Use forward_step_cached with pre-computed modulation batch for this ODE step + let flow_dir = flow_net.forward_step_cached(¤t, step_mod)?; + current = (current + flow_dir.affine(step_factor, 0.0)?)?; + } + Ok(current) +} + +#[derive(Clone)] +pub struct FlowLMModel { + pub flow_net: SimpleMLPAdaLN, + pub transformer: StreamingTransformer, + pub input_linear: Linear, + pub out_norm: LayerNorm, + pub out_eos: Linear, + pub bos_emb: Tensor, + pub emb_mean: Tensor, + pub emb_std: Tensor, + pub ldim: usize, + pub dim: usize, + pub noise_clamp: Option, +} + +fn sample_noise( + device: &candle_core::Device, + shape: (usize, usize), + temp: f32, + clamp: Option, +) -> Result { + let std = temp.sqrt(); + let noise = Tensor::randn(0.0f32, std, shape, device)?; + match clamp { + None => Ok(noise), + Some(limit) => noise.clamp(-limit, limit), + } +} + +impl FlowLMModel { + pub fn new( + flow_net: SimpleMLPAdaLN, + transformer: StreamingTransformer, + ldim: usize, + dim: usize, + vb: VarBuilder, + ) -> Result { + let input_linear = candle_nn::linear_no_bias(ldim, dim, vb.pp("input_linear"))?; + let out_norm = LayerNorm::new(dim, 1e-5, true, vb.pp("out_norm"))?; + let out_eos = candle_nn::linear(dim, 1, vb.pp("out_eos"))?; + let bos_emb = vb.get(ldim, "bos_emb")?; + let emb_mean = vb.get(ldim, "emb_mean")?; + let emb_std = vb.get(ldim, "emb_std")?; + + Ok(Self { + flow_net, + transformer, + input_linear, + out_norm, + out_eos, + bos_emb, + emb_mean, + emb_std, + ldim, + dim, + noise_clamp: None, // Default to no clamp + }) + } + + #[allow(clippy::too_many_arguments)] + pub fn forward( + &self, + sequence: &Tensor, + text_embeddings: &Tensor, + model_state: &mut ModelState, + time_embeddings: &Tensor, + temp: f32, + eos_threshold: f32, + step: usize, + ) -> Result<(Tensor, bool)> { + // sequence is [B, T, ldim] + // text_embeddings is [B, S, dim] + + // Handle BOS (if NaN, use bos_emb) - simplistic check for NaN + // In Candle we can use `Tensor::where_cond` + // But for now let's assume sequence passed in doesn't have NaNs or handled upstream. + // Original: sequence = torch.where(torch.isnan(sequence), self.bos_emb, sequence) + + // Let's assume BOS is handled by caller for now or if sequence empty. + + let x = self.input_linear.forward(sequence)?; + let s_len = text_embeddings.dims()[1]; + + // Cat text embeddings and sequence embeddings only if text_embeddings is not empty + let transformer_out_pre_norm = if s_len > 0 { + let input = Tensor::cat(&[text_embeddings, &x], 1)?; + let mut out = self.transformer.forward(&input, model_state, step)?; + // Remove prefix (text embeddings length) + out = out.narrow(1, s_len, out.dims()[1] - s_len)?; + out + } else { + self.transformer.forward(&x, model_state, step)? + }; + + let transformer_out = self.out_norm.forward(&transformer_out_pre_norm)?; + + // Only use the last frame for generation + let last_frame = transformer_out + .narrow(1, transformer_out.dims()[1] - 1, 1)? + .squeeze(1)?; + + let eos_score = self + .out_eos + .forward(&last_frame)? + .squeeze(0)? + .squeeze(0)? + .to_scalar::()?; + let is_eos = eos_score > eos_threshold; + + // Generate noise with optional clamping + let noise = sample_noise( + last_frame.device(), + (last_frame.dims()[0], self.ldim), + temp, + self.noise_clamp, + )?; + + // Pre-compute all modulations for this frame's ODE steps (8 steps * N blocks) in batch + let c_emb = self.flow_net.embed_condition(&last_frame)?; + let modulations = self + .flow_net + .precompute_modulations(&c_emb, time_embeddings)?; + + let next_latent = lsd_decode(&self.flow_net, &modulations, &noise)?; + + Ok((next_latent, is_eos)) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/models/mimi.rs b/mistralrs-core/src/speech_models/pockettts/models/mimi.rs new file mode 100644 index 0000000000..9a11de4c9e --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/models/mimi.rs @@ -0,0 +1,262 @@ +use crate::speech_models::pockettts::models::seanet::{SEANetDecoder, SEANetEncoder}; +use crate::speech_models::pockettts::models::transformer::ProjectedTransformer; +use crate::speech_models::pockettts::modules::conv::{ConvDownsample1d, ConvTrUpsample1d}; +use crate::speech_models::pockettts::voice_state::ModelState; +use candle_core::{Result, Tensor}; +use candle_nn::{Conv1d, Conv1dConfig, Module, VarBuilder}; + +#[derive(Clone)] +pub struct Quantizer { + output_proj: Conv1d, +} + +impl Quantizer { + pub fn new(dimension: usize, output_dimension: usize, vb: VarBuilder) -> Result { + let config = Conv1dConfig { + groups: 1, + padding: 0, + stride: 1, + dilation: 1, + ..Default::default() + }; + let output_proj = candle_nn::conv1d_no_bias( + dimension, + output_dimension, + 1, + config, + vb.pp("output_proj"), + )?; + Ok(Self { output_proj }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + // x is [B, C, T] + // Conv1d expects [B, C, T] and returns [B, C_out, T] + self.output_proj.forward(x) + } +} + +#[derive(Clone)] +pub struct MimiModel { + pub encoder: SEANetEncoder, + pub decoder: SEANetDecoder, + pub encoder_transformer: ProjectedTransformer, + pub decoder_transformer: ProjectedTransformer, + pub quantizer: Quantizer, + pub downsample: Option, + pub upsample: Option, + pub frame_rate: f64, + pub encoder_frame_rate: f64, + pub sample_rate: usize, + pub channels: usize, + pub dimension: usize, +} + +impl MimiModel { + #[allow(clippy::too_many_arguments)] + pub fn new( + encoder: SEANetEncoder, + decoder: SEANetDecoder, + encoder_transformer: ProjectedTransformer, + decoder_transformer: ProjectedTransformer, + frame_rate: f64, + encoder_frame_rate: f64, + sample_rate: usize, + channels: usize, + dimension: usize, // The quantizer input dimension (32) + output_dimension: usize, // The decoder input dimension (512) + name: &str, + vb: VarBuilder, + ) -> Result { + let quantizer = Quantizer::new(dimension, output_dimension, vb.pp("quantizer"))?; + + let (downsample, upsample) = if encoder_frame_rate != frame_rate { + let stride = (encoder_frame_rate / frame_rate) as usize; + ( + Some(ConvDownsample1d::new( + stride, + output_dimension, + &format!("{}.downsample", name), + vb.pp("downsample"), + )?), + Some(ConvTrUpsample1d::new( + stride, + output_dimension, + &format!("{}.upsample", name), + vb.pp("upsample"), + )?), + ) + } else { + (None, None) + }; + + Ok(Self { + encoder, + decoder, + encoder_transformer, + decoder_transformer, + quantizer, + downsample, + upsample, + frame_rate, + encoder_frame_rate, + sample_rate, + channels, + dimension, + }) + } + + pub fn frame_size(&self) -> usize { + (self.sample_rate as f64 / self.frame_rate) as usize + } + + pub fn encode_to_latent( + &self, + x: &Tensor, + model_state: &mut ModelState, + step: usize, + ) -> Result { + // x shape [B, C, T] + let _frame_size = self.frame_size(); + let (b, c, _t_orig) = x.dims3()?; + + let t = x.dims()[2]; + let hop = self.frame_size(); + let x = if !t.is_multiple_of(hop) { + let padding = hop - (t % hop); + let pad = Tensor::zeros((b, c, padding), x.dtype(), x.device())?; + Tensor::cat(&[x, &pad], 2)? + } else { + x.clone() + }; + + let mut emb = self.encoder.forward(&x, model_state, step)?; + let mut embs = self.encoder_transformer.forward(&emb, model_state, step)?; + emb = embs.remove(0); + + if let Some(down) = &self.downsample { + emb = down.forward(&emb, model_state, step)?; + } + Ok(emb) + } + + pub fn decode_from_latent( + &self, + latent: &Tensor, + model_state: &mut ModelState, + step: usize, + ) -> Result { + let mut emb = latent.clone(); + if let Some(up) = &self.upsample { + emb = up.forward(&emb, model_state, step)?; + } + let mut embs = self.decoder_transformer.forward(&emb, model_state, step)?; + emb = embs.remove(0); + let out = self.decoder.forward(&emb, model_state, step)?; + Ok(out) + } + pub fn quantize(&self, x: &Tensor) -> Result { + self.quantizer.forward(x) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use candle_core::{DType, Device, Tensor}; + use candle_nn::VarBuilder; + use std::collections::HashMap; + + #[test] + fn test_mimi_shapes() -> Result<()> { + let device = Device::Cpu; + let vb = VarBuilder::zeros(DType::F32, &device); + + let encoder = SEANetEncoder::new( + 1, + 128, + 32, + 1, + &[2, 2], + 7, + 7, + 3, + 2, + "constant", + 2, + "encoder", + vb.pp("encoder"), + )?; + let decoder = SEANetDecoder::new( + 1, + 128, + 32, + 1, + &[2, 2], + 7, + 7, + 3, + 2, + "constant", + 2, + "decoder", + vb.pp("decoder"), + )?; + + let encoder_transformer = ProjectedTransformer::new( + 128, + vec![128], + 128, + 4, + 1, + 0.1, + 10, + 10000.0, + 512, + "enc_tr", + vb.pp("enc_tr"), + )?; + let decoder_transformer = ProjectedTransformer::new( + 128, + vec![128], + 128, + 4, + 1, + 0.1, + 10, + 10000.0, + 512, + "dec_tr", + vb.pp("dec_tr"), + )?; + + let mimi = MimiModel::new( + encoder, + decoder, + encoder_transformer, + decoder_transformer, + 12.5, + 50.0, + 16000, + 1, + 128, + 512, + "mimi", + vb.pp("mimi"), + )?; + + let _audio = Tensor::zeros((1, 1, 1280), DType::F32, &device)?; // 1280 samples = 0.08s + + // Mock state + let mut _model_state: HashMap> = HashMap::new(); + // We need to initialize state for all submodules. This is complex manually. + // For shape test, we might want a simpler way or just skip stateful forward for now if init complex. + // But our forward REQUIRES state. + + // I'll skip the actual forward test here because initializing state for ALL sub-layers is tedious. + // I'll implement a helper to init states in Phase 3. + + assert_eq!(mimi.frame_size(), 1280); + Ok(()) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/models/mod.rs b/mistralrs-core/src/speech_models/pockettts/models/mod.rs new file mode 100644 index 0000000000..647045a849 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/models/mod.rs @@ -0,0 +1,4 @@ +pub mod flow_lm; +pub mod mimi; +pub mod seanet; +pub mod transformer; diff --git a/mistralrs-core/src/speech_models/pockettts/models/seanet.rs b/mistralrs-core/src/speech_models/pockettts/models/seanet.rs new file mode 100644 index 0000000000..13c8e3f057 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/models/seanet.rs @@ -0,0 +1,403 @@ +use crate::speech_models::pockettts::modules::conv::{StreamingConv1d, StreamingConvTranspose1d}; +use crate::speech_models::pockettts::voice_state::ModelState; +use candle_core::{Result, Tensor}; +use candle_nn::VarBuilder; + +#[derive(Clone)] +pub struct SEANetResnetBlock { + pub layers: Vec>, + pub _name: String, +} + +pub trait StreamingLayer: Send + Sync { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result; + fn clone_box(&self) -> Box; +} + +impl Clone for Box { + fn clone(&self) -> Self { + self.clone_box() + } +} + +impl StreamingLayer for StreamingConv1d { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.forward(x, model_state, step) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +#[derive(Clone)] +pub struct EluLayer; +impl StreamingLayer for EluLayer { + fn forward(&self, x: &Tensor, _model_state: &mut ModelState, _step: usize) -> Result { + x.elu(1.0) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl SEANetResnetBlock { + pub fn new( + dim: usize, + kernel_sizes: &[usize], + dilations: &[usize], + pad_mode: &str, + compress: usize, + name: &str, + vb: VarBuilder, + ) -> Result { + let hidden = dim / compress; + let mut layers: Vec> = Vec::new(); + for i in 0..kernel_sizes.len() { + let in_chs = if i == 0 { dim } else { hidden }; + let out_chs = if i == kernel_sizes.len() - 1 { + dim + } else { + hidden + }; + layers.push(Box::new(EluLayer)); + layers.push(Box::new(StreamingConv1d::new( + in_chs, + out_chs, + kernel_sizes[i], + 1, + dilations[i], + 1, + true, + pad_mode, + &format!("{}.block.{}", name, i * 2 + 1), + vb.pp(format!("block.{}", i * 2 + 1)), + )?)); + } + Ok(Self { + layers, + _name: name.to_string(), + }) + } + + pub fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + let mut v = x.clone(); + for layer in &self.layers { + v = layer.forward(&v, model_state, step)?; + } + x + v + } +} + +#[derive(Clone)] +pub struct SEANetEncoder { + pub layers: Vec>, + pub hop_length: usize, + pub _name: String, +} + +pub trait StreamingLayerWrapper: Send + Sync { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result; + fn clone_box(&self) -> Box; + fn weight(&self) -> Option<&Tensor> { + None + } + fn bias(&self) -> Option<&Tensor> { + None + } +} + +impl Clone for Box { + fn clone(&self) -> Self { + self.clone_box() + } +} + +impl StreamingLayerWrapper for StreamingConv1d { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.forward(x, model_state, step) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } + fn weight(&self) -> Option<&Tensor> { + Some(self.weight()) + } + fn bias(&self) -> Option<&Tensor> { + self.bias() + } +} + +impl StreamingLayerWrapper for SEANetResnetBlock { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.forward(x, model_state, step) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl StreamingLayerWrapper for EluLayer { + fn forward(&self, x: &Tensor, _model_state: &mut ModelState, _step: usize) -> Result { + x.elu(1.0) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl SEANetEncoder { + #[allow(clippy::too_many_arguments)] + pub fn new( + channels: usize, + dimension: usize, + n_filters: usize, + n_residual_layers: usize, + ratios: &[usize], + kernel_size: usize, + last_kernel_size: usize, + residual_kernel_size: usize, + dilation_base: usize, + pad_mode: &str, + compress: usize, + name: &str, + vb: VarBuilder, + ) -> Result { + let ratios: Vec = ratios.iter().copied().rev().collect(); + let hop_length = ratios.iter().product(); + let mut layers: Vec> = Vec::new(); + + let mut mult = 1; + layers.push(Box::new(StreamingConv1d::new( + channels, + mult * n_filters, + kernel_size, + 1, + 1, + 1, + true, + pad_mode, + &format!("{}.model.0", name), + vb.pp("model.0"), + )?)); + + let mut layer_idx = 1; + for ratio in ratios { + for j in range(n_residual_layers) { + layers.push(Box::new(SEANetResnetBlock::new( + mult * n_filters, + &[residual_kernel_size, 1], + &[dilation_base.pow(j as u32), 1], + pad_mode, + compress, + &format!("{}.model.{}", name, layer_idx), + vb.pp(format!("model.{}", layer_idx)), + )?)); + layer_idx += 1; + } + + layers.push(Box::new(EluLayer)); + layer_idx += 1; + + layers.push(Box::new(StreamingConv1d::new( + mult * n_filters, + mult * n_filters * 2, + ratio * 2, + ratio, + 1, + 1, + true, + pad_mode, + &format!("{}.model.{}", name, layer_idx), + vb.pp(format!("model.{}", layer_idx)), + )?)); + layer_idx += 1; + mult *= 2; + } + + layers.push(Box::new(EluLayer)); + layer_idx += 1; + + layers.push(Box::new(StreamingConv1d::new( + mult * n_filters, + dimension, + last_kernel_size, + 1, + 1, + 1, + true, + pad_mode, + &format!("{}.model.{}", name, layer_idx), + vb.pp(format!("model.{}", layer_idx)), + )?)); + + Ok(Self { + layers, + hop_length, + _name: name.to_string(), + }) + } + + pub fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + let mut x = x.clone(); + for layer in &self.layers { + x = layer.forward(&x, model_state, step)?; + } + Ok(x) + } +} + +fn range(n: usize) -> std::ops::Range { + 0..n +} + +#[derive(Clone)] +pub struct SEANetDecoder { + pub layers: Vec>, + pub hop_length: usize, + pub _name: String, +} + +pub trait StreamingLayerDecoderWrapper: Send + Sync { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result; + fn clone_box(&self) -> Box; +} + +impl Clone for Box { + fn clone(&self) -> Self { + self.clone_box() + } +} + +impl StreamingLayerDecoderWrapper for StreamingConv1d { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.forward(x, model_state, step) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl StreamingLayerDecoderWrapper for StreamingConvTranspose1d { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.forward(x, model_state, step) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl StreamingLayerDecoderWrapper for SEANetResnetBlock { + fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.forward(x, model_state, step) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl StreamingLayerDecoderWrapper for EluLayer { + fn forward(&self, x: &Tensor, _model_state: &mut ModelState, _step: usize) -> Result { + x.elu(1.0) + } + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl SEANetDecoder { + #[allow(clippy::too_many_arguments)] + pub fn new( + channels: usize, + dimension: usize, + n_filters: usize, + n_residual_layers: usize, + ratios: &[usize], + kernel_size: usize, + last_kernel_size: usize, + residual_kernel_size: usize, + dilation_base: usize, + pad_mode: &str, + compress: usize, + name: &str, + vb: VarBuilder, + ) -> Result { + let hop_length = ratios.iter().product(); + let mut layers: Vec> = Vec::new(); + + let mut mult = 2usize.pow(ratios.len() as u32); + layers.push(Box::new(StreamingConv1d::new( + dimension, + mult * n_filters, + kernel_size, + 1, + 1, + 1, + true, + pad_mode, + &format!("{}.model.0", name), + vb.pp("model.0"), + )?)); + + let mut layer_idx = 1; + for ratio in ratios { + layers.push(Box::new(EluLayer)); + layer_idx += 1; + + layers.push(Box::new(StreamingConvTranspose1d::new( + mult * n_filters, + mult * n_filters / 2, + ratio * 2, + *ratio, + 1, + true, + &format!("{}.model.{}", name, layer_idx), + vb.pp(format!("model.{}", layer_idx)), + )?)); + layer_idx += 1; + + for j in range(n_residual_layers) { + layers.push(Box::new(SEANetResnetBlock::new( + mult * n_filters / 2, + &[residual_kernel_size, 1], + &[dilation_base.pow(j as u32), 1], + pad_mode, + compress, + &format!("{}.model.{}", name, layer_idx), + vb.pp(format!("model.{}", layer_idx)), + )?)); + layer_idx += 1; + } + mult /= 2; + } + + layers.push(Box::new(EluLayer)); + layer_idx += 1; + + layers.push(Box::new(StreamingConv1d::new( + n_filters, + channels, + last_kernel_size, + 1, + 1, + 1, + true, + pad_mode, + &format!("{}.model.{}", name, layer_idx), + vb.pp(format!("model.{}", layer_idx)), + )?)); + + Ok(Self { + layers, + hop_length, + _name: name.to_string(), + }) + } + + pub fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + let mut x = x.clone(); + for layer in &self.layers { + x = layer.forward(&x, model_state, step)?; + } + Ok(x) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/models/transformer.rs b/mistralrs-core/src/speech_models/pockettts/models/transformer.rs new file mode 100644 index 0000000000..b3785f0e02 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/models/transformer.rs @@ -0,0 +1,252 @@ +use crate::speech_models::pockettts::modules::attention::StreamingMultiheadAttention; +use crate::speech_models::pockettts::modules::mlp::{LayerNorm, LayerScale}; +use crate::speech_models::pockettts::modules::rope::RotaryEmbedding; +use crate::speech_models::pockettts::voice_state::get_attention_cursor; +use crate::speech_models::pockettts::voice_state::ModelState; +use candle_core::{Result, Tensor}; +use candle_nn::{Linear, Module, VarBuilder}; + +#[derive(Clone)] +pub struct StreamingTransformerLayer { + self_attn: StreamingMultiheadAttention, + norm1: LayerNorm, + norm2: LayerNorm, + linear1: Linear, + linear2: Linear, + layer_scale_1: Option, + layer_scale_2: Option, +} + +impl StreamingTransformerLayer { + #[allow(clippy::too_many_arguments)] + pub fn new( + d_model: usize, + num_heads: usize, + dim_feedforward: usize, + context: Option, + rope: RotaryEmbedding, + layer_scale: Option, + _attention_kind: &str, + name: &str, + vb: VarBuilder, + ) -> Result { + let self_attn = StreamingMultiheadAttention::new( + d_model, + num_heads, + rope, + context, + &format!("{}.self_attn", name), + vb.pp("self_attn"), + )?; + let norm1 = LayerNorm::new(d_model, 1e-5, true, vb.pp("norm1"))?; + let norm2 = LayerNorm::new(d_model, 1e-5, true, vb.pp("norm2"))?; + let linear1 = candle_nn::linear_no_bias(d_model, dim_feedforward, vb.pp("linear1"))?; + let linear2 = candle_nn::linear_no_bias(dim_feedforward, d_model, vb.pp("linear2"))?; + + let (layer_scale_1, layer_scale_2) = if let Some(init) = layer_scale { + ( + Some(LayerScale::new(d_model, init, vb.pp("layer_scale_1"))?), + Some(LayerScale::new(d_model, init, vb.pp("layer_scale_2"))?), + ) + } else { + (None, None) + }; + + Ok(Self { + self_attn, + norm1, + norm2, + linear1, + linear2, + layer_scale_1, + layer_scale_2, + }) + } + + pub fn forward( + &self, + x: &Tensor, + model_state: &mut ModelState, + current_pos: usize, + current_len: usize, + ) -> Result { + let x_orig = x.clone(); + let h = self.norm1.forward(x)?; + let mut update = self + .self_attn + .forward(&h, model_state, current_pos, current_len)?; + if let Some(ls) = &self.layer_scale_1 { + update = ls.forward(&update)?; + } + let x = (x_orig + update)?; + + let x_orig = x.clone(); + let h = self.norm2.forward(&x)?; + let mut update = self.linear2.forward(&self.linear1.forward(&h)?.gelu()?)?; + if let Some(ls) = &self.layer_scale_2 { + update = ls.forward(&update)?; + } + x_orig + update + } +} + +#[derive(Clone)] +pub struct StreamingTransformer { + layers: Vec, + _rope: RotaryEmbedding, + name: String, +} + +impl StreamingTransformer { + #[allow(clippy::too_many_arguments)] + pub fn new( + d_model: usize, + num_heads: usize, + num_layers: usize, + layer_scale: Option, + dim_feedforward: usize, + context: Option, + max_period: f32, + kind: &str, + name: &str, + vb: VarBuilder, + ) -> Result { + let rope = RotaryEmbedding::new(max_period, d_model / num_heads, vb.device())?; + let mut layers = Vec::new(); + for i in 0..num_layers { + layers.push(StreamingTransformerLayer::new( + d_model, + num_heads, + dim_feedforward, + context, + rope.clone(), + layer_scale, + kind, + &format!("{}.layers.{}", name, i), + vb.pp(format!("layers.{}", i)), + )?); + } + Ok(Self { + layers, + _rope: rope, + name: name.to_string(), + }) + } + + pub fn forward( + &self, + x: &Tensor, + model_state: &mut ModelState, + _step: usize, + ) -> Result { + let mut x = x.clone(); + // Fetch current_pos once from the first attention layer's state to avoid redundant to_scalar calls. + let first_layer_name = format!("{}.layers.0.self_attn", self.name); + let cursor = get_attention_cursor(model_state, &first_layer_name); + let current_pos = cursor.pos; + let current_len = cursor.len; + + for layer in &self.layers { + x = layer.forward(&x, model_state, current_pos, current_len)?; + } + Ok(x) + } +} + +#[derive(Clone)] +pub struct ProjectedTransformer { + transformer: StreamingTransformer, + input_proj: Option, + output_projs: Vec>, + _input_dimension: usize, + _output_dimensions: Vec, + _d_model: usize, +} + +impl ProjectedTransformer { + #[allow(clippy::too_many_arguments)] + pub fn new( + input_dimension: usize, + output_dimensions: Vec, + d_model: usize, + num_heads: usize, + num_layers: usize, + layer_scale: f32, + context: usize, + max_period: f32, + dim_feedforward: usize, + name: &str, + vb: VarBuilder, + ) -> Result { + let transformer = StreamingTransformer::new( + d_model, + num_heads, + num_layers, + Some(layer_scale), + dim_feedforward, + Some(context), + max_period, + "mimi", + &format!("{}.transformer", name), + vb.pp("transformer"), + )?; + + let input_proj = if d_model != input_dimension { + Some(candle_nn::linear_no_bias( + input_dimension, + d_model, + vb.pp("input_proj"), + )?) + } else { + None + }; + + let mut output_projs = Vec::new(); + for (i, &output_dim) in output_dimensions.iter().enumerate() { + if d_model == output_dim { + output_projs.push(None); + } else { + output_projs.push(Some(candle_nn::linear_no_bias( + d_model, + output_dim, + vb.pp(format!("output_projs.{}", i)), + )?)); + } + } + + Ok(Self { + transformer, + input_proj, + output_projs, + _input_dimension: input_dimension, + _output_dimensions: output_dimensions, + _d_model: d_model, + }) + } + + pub fn forward( + &self, + x: &Tensor, + model_state: &mut ModelState, + step: usize, + ) -> Result> { + // x is [B, C, T] + let mut x = x.transpose(1, 2)?; // [B, T, C] + if let Some(proj) = &self.input_proj { + x = proj.forward(&x)?; + } + let z = self.transformer.forward(&x, model_state, step)?; + + let mut ys = Vec::new(); + for output_proj in &self.output_projs { + let mut y = if let Some(proj) = output_proj { + proj.forward(&z)? + } else { + z.clone() + }; + y = y.transpose(1, 2)?; // [B, C_out, T] + ys.append(&mut vec![y]); + } + Ok(ys) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/modules/attention.rs b/mistralrs-core/src/speech_models/pockettts/modules/attention.rs new file mode 100644 index 0000000000..4de6a08a8b --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/modules/attention.rs @@ -0,0 +1,291 @@ +use crate::speech_models::pockettts::modules::rope::RotaryEmbedding; +use crate::speech_models::pockettts::voice_state::ModelState; +use crate::speech_models::pockettts::voice_state::{ + read_attention_cursor, write_attention_cursor, AttentionCursor, ATTN_K_BUF_KEY, ATTN_LEN_KEY, + ATTN_POS_KEY, ATTN_V_BUF_KEY, +}; +use candle_core::{DType, Result, Tensor}; +use candle_nn::{Linear, Module, VarBuilder}; +use std::collections::HashMap; + +fn ring_chunks(buf: &Tensor, head: usize, len: usize) -> Result> { + let cap = buf.dim(2)?; + if cap == 0 { + return Ok(Vec::new()); + } + + let len = len.min(cap); + if len == 0 { + return Ok(Vec::new()); + } + + let head = head % cap; + let first_len = std::cmp::min(len, cap - head); + let second_len = len - first_len; + let mut chunks = Vec::with_capacity(if second_len > 0 { 2 } else { 1 }); + chunks.push(buf.narrow(2, head, first_len)?); + if second_len > 0 { + chunks.push(buf.narrow(2, 0, second_len)?); + } + Ok(chunks) +} + +#[derive(Clone)] +pub struct StreamingMultiheadAttention { + embed_dim: usize, + num_heads: usize, + rope: RotaryEmbedding, + in_proj: Linear, + out_proj: Linear, + context: Option, + name: String, +} + +impl StreamingMultiheadAttention { + pub fn new( + embed_dim: usize, + num_heads: usize, + rope: RotaryEmbedding, + context: Option, + name: &str, + vb: VarBuilder, + ) -> Result { + // out_dim = embed_dim + 2 * kv_dim (GQA/MHA logic in original) + // Original code: + // out_dim = embed_dim + // num_kv = num_heads + // kv_dim = (embed_dim // num_heads) * num_kv -> so embed_dim + // out_dim += 2 * kv_dim -> so 3 * embed_dim + let in_proj = candle_nn::linear_no_bias(embed_dim, 3 * embed_dim, vb.pp("in_proj"))?; + let out_proj = candle_nn::linear_no_bias(embed_dim, embed_dim, vb.pp("out_proj"))?; + + Ok(Self { + embed_dim, + num_heads, + rope, + in_proj, + out_proj, + context, + name: name.to_string(), + }) + } + + pub fn init_state( + &self, + batch_size: usize, + _sequence_length: usize, + device: &candle_core::Device, + ) -> Result> { + let dim_per_head = self.embed_dim / self.num_heads; + let mut state = HashMap::new(); + + // Initial capacity: match context if windowed, otherwise reasonable default + let cap = self.context.unwrap_or(64); + state.insert( + ATTN_K_BUF_KEY.to_string(), + Tensor::zeros( + (batch_size, self.num_heads, cap, dim_per_head), + DType::F32, + device, + )?, + ); + state.insert( + ATTN_V_BUF_KEY.to_string(), + Tensor::zeros( + (batch_size, self.num_heads, cap, dim_per_head), + DType::F32, + device, + )?, + ); + write_attention_cursor(&mut state, AttentionCursor::default(), device)?; + Ok(state) + } + + pub fn forward( + &self, + query: &Tensor, + model_state: &mut ModelState, + current_pos: usize, + current_len: usize, + ) -> Result { + let (b, t, _) = query.dims3()?; + let d = self.embed_dim / self.num_heads; + let window_size = self.context; + + // Auto-initialize state if missing + if !model_state.contains_key(&self.name) { + model_state.insert(self.name.clone(), self.init_state(b, 0, query.device())?); + } + + let module_state = model_state.get_mut(&self.name).unwrap(); + let mut cursor = read_attention_cursor(module_state); + if !module_state.contains_key(ATTN_POS_KEY) { + cursor.pos = current_pos; + } + if !module_state.contains_key(ATTN_LEN_KEY) { + cursor.len = current_len; + } + + let projected = self.in_proj.forward(query)?; + + // Reshape to (b, t, 3, h, d) + let packed = projected.reshape((b, t, 3, self.num_heads, d))?; + let mut q = packed.narrow(2, 0, 1)?.squeeze(2)?; // (b, t, h, d) + let mut k = packed.narrow(2, 1, 1)?.squeeze(2)?; // (b, t, h, d) + let mut v = packed.narrow(2, 2, 1)?.squeeze(2)?; // (b, t, h, d) + + // current_pos passed as argument + + // Apply RoPE + // RoPE expects (B, T, H, D) + (q, k) = self.rope.forward(&q, &k, current_pos)?; + + // Transpose q, k, v to (B, H, T, D) for SDPA and KV cache + q = q.transpose(1, 2)?; + k = k.transpose(1, 2)?; + v = v.transpose(1, 2)?; + + // KV cache management. + // We take ownership from the state to avoid clones and ensure uniqueness for slice_set. + let (mut k_buf, mut v_buf) = match ( + module_state.remove(ATTN_K_BUF_KEY), + module_state.remove(ATTN_V_BUF_KEY), + ) { + (Some(kb), Some(vb)) => (kb, vb), + _ => { + let initial_cap = window_size.unwrap_or(64); + let kb = Tensor::zeros((b, self.num_heads, initial_cap, d), q.dtype(), q.device())?; + let vb = Tensor::zeros((b, self.num_heads, initial_cap, d), q.dtype(), q.device())?; + (kb, vb) + } + }; + + let mut cap = k_buf.dim(2)?; // Current capacity of the buffer + let mut cache_len = cursor.len.min(cap); + let mut cache_head = if cap > 0 { cursor.head % cap } else { 0 }; + + let x = if let Some(window_size) = self.context { + // Ensure fixed ring capacity for windowed attention. + if cap != window_size { + if cap > window_size { + k_buf = k_buf.narrow(2, 0, window_size)?.contiguous()?; + v_buf = v_buf.narrow(2, 0, window_size)?.contiguous()?; + } else { + let zeros_shape = (b, self.num_heads, window_size - cap, d); + let k_zeros = Tensor::zeros(zeros_shape, q.dtype(), q.device())?; + let v_zeros = Tensor::zeros(zeros_shape, q.dtype(), q.device())?; + k_buf = Tensor::cat(&[k_buf, k_zeros], 2)?; + v_buf = Tensor::cat(&[v_buf, v_zeros], 2)?; + } + cap = window_size; + cache_len = cache_len.min(cap); + cache_head = 0; + } + + // Build chronological KV chunks from ring cache + current K/V chunk. + let mut k_chunks = ring_chunks(&k_buf, cache_head, cache_len)?; + let mut v_chunks = ring_chunks(&v_buf, cache_head, cache_len)?; + k_chunks.push(k.clone()); + v_chunks.push(v.clone()); + + let scale = 1.0 / (d as f64).sqrt(); + if k_chunks.len() == 1 { + crate::speech_models::pockettts::modules::sdpa::sdpa( + &q, + &k_chunks[0], + &v_chunks[0], + scale, + true, + self.context, + )? + } else { + crate::speech_models::pockettts::modules::sdpa::sdpa_chunked( + &q, + &k_chunks, + &v_chunks, + scale, + true, + self.context, + )? + } + } else { + // Linear attention (FlowLM) with doubling contiguous buffer. + if cache_len + t > cap { + let new_cap = (cache_len + t).next_power_of_two(); + let zeros_shape = (b, self.num_heads, new_cap - cap, d); + let k_zeros = Tensor::zeros(zeros_shape, q.dtype(), q.device())?; + let v_zeros = Tensor::zeros(zeros_shape, q.dtype(), q.device())?; + k_buf = Tensor::cat(&[k_buf, k_zeros], 2)?; + v_buf = Tensor::cat(&[v_buf, v_zeros], 2)?; + } + k_buf.slice_set(&k.contiguous()?, 2, cache_len)?; + v_buf.slice_set(&v.contiguous()?, 2, cache_len)?; + cache_len += t; + cache_head = 0; + + // Get current KV for attention + let kc = k_buf.narrow(2, 0, cache_len)?; + let vc = v_buf.narrow(2, 0, cache_len)?; + let scale = 1.0 / (d as f64).sqrt(); + crate::speech_models::pockettts::modules::sdpa::sdpa( + &q, + &kc, + &vc, + scale, + true, + self.context, + )? + }; + + if let Some(window_size) = self.context { + if t >= window_size { + k_buf = k.narrow(2, t - window_size, window_size)?.contiguous()?; + v_buf = v.narrow(2, t - window_size, window_size)?.contiguous()?; + cache_head = 0; + cache_len = window_size; + } else if window_size > 0 { + let evict = (cache_len + t).saturating_sub(window_size); + if evict > 0 { + cache_head = (cache_head + evict) % window_size; + cache_len -= evict; + } + + let write_start = (cache_head + cache_len) % window_size; + let first = std::cmp::min(t, window_size - write_start); + let second = t - first; + + if first > 0 { + let k_first = k.narrow(2, 0, first)?.contiguous()?; + let v_first = v.narrow(2, 0, first)?.contiguous()?; + k_buf.slice_set(&k_first, 2, write_start)?; + v_buf.slice_set(&v_first, 2, write_start)?; + } + if second > 0 { + let k_second = k.narrow(2, first, second)?.contiguous()?; + let v_second = v.narrow(2, first, second)?.contiguous()?; + k_buf.slice_set(&k_second, 2, 0)?; + v_buf.slice_set(&v_second, 2, 0)?; + } + cache_len += t; + } + } + + module_state.insert(ATTN_K_BUF_KEY.to_string(), k_buf); + module_state.insert(ATTN_V_BUF_KEY.to_string(), v_buf); + write_attention_cursor( + module_state, + AttentionCursor { + pos: current_pos + t, + len: cache_len, + head: cache_head, + }, + q.device(), + )?; + + // Transpose back to [B, T, H, D] and project out + let x = x.transpose(1, 2)?.reshape((b, t, self.embed_dim))?; + let x = self.out_proj.forward(&x)?; + + Ok(x) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/modules/conv.rs b/mistralrs-core/src/speech_models/pockettts/modules/conv.rs new file mode 100644 index 0000000000..7335a056d7 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/modules/conv.rs @@ -0,0 +1,346 @@ +use crate::speech_models::pockettts::voice_state::ModelState; +use candle_core::{DType, Result, Tensor}; +use candle_nn::{Conv1d, Conv1dConfig, ConvTranspose1d, ConvTranspose1dConfig, Module, VarBuilder}; +use std::collections::HashMap; + +#[derive(Clone)] +pub struct StreamingConv1d { + conv: Conv1d, + padding_mode: String, + stride: usize, + kernel_size: usize, + dilation: usize, + in_channels: usize, + name: String, +} + +impl StreamingConv1d { + #[allow(clippy::too_many_arguments)] + pub fn new( + in_channels: usize, + out_channels: usize, + kernel_size: usize, + stride: usize, + dilation: usize, + groups: usize, + bias: bool, + padding_mode: &str, + name: &str, + vb: VarBuilder, + ) -> Result { + let config = Conv1dConfig { + stride, + padding: 0, + dilation, + groups, + ..Default::default() + }; + let conv = if bias { + candle_nn::conv1d( + in_channels, + out_channels, + kernel_size, + config, + vb.pp("conv"), + )? + } else { + candle_nn::conv1d_no_bias( + in_channels, + out_channels, + kernel_size, + config, + vb.pp("conv"), + )? + }; + + Ok(Self { + conv, + padding_mode: padding_mode.to_string(), + stride, + kernel_size, + dilation, + in_channels, + name: name.to_string(), + }) + } + + pub fn effective_kernel_size(&self) -> usize { + (self.kernel_size - 1) * self.dilation + 1 + } + + pub fn init_state( + &self, + batch_size: usize, + _sequence_length: usize, + device: &candle_core::Device, + ) -> Result> { + let kernel = self.effective_kernel_size(); + let mut state = HashMap::new(); + if kernel > self.stride { + let previous = Tensor::zeros( + (batch_size, self.in_channels, kernel - self.stride), + DType::F32, + device, + )?; + state.insert("previous".to_string(), previous); + } + Ok(state) + } + + pub fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + let (b, c, t) = x.dims3()?; + let s = self.stride; + if t == 0 || t % s != 0 { + return Err(candle_core::Error::Msg(format!( + "Steps must be multiple of stride {}, got {}", + s, t + ))); + } + + // Auto-initialize state if missing + if !model_state.contains_key(&self.name) { + let init = self.init_state(b, t, x.device())?; + model_state.insert(self.name.clone(), init); + } + + let module_state = model_state.get_mut(&self.name).unwrap(); + let kernel = self.effective_kernel_size(); + let pad_left = kernel.saturating_sub(s); + + if pad_left > 0 { + let previous = module_state + .remove("previous") + .ok_or_else(|| candle_core::Error::Msg("previous state not found".to_string()))?; + let is_first = step == 0; + + let x_with_padding = if is_first && self.padding_mode == "replicate" { + // Replicate the first frame for the initial padding + let first_frame = x.narrow(2, 0, 1)?; + let replicated_padding = first_frame.broadcast_as((b, c, pad_left))?; + Tensor::cat(&[replicated_padding, x.clone()], 2)? + } else { + Tensor::cat(&[previous, x.clone()], 2)? + }; + + let y = self.conv.forward(&x_with_padding)?; + + // Update previous state for next call + let total_len = x_with_padding.dims()[2]; + let new_previous = x_with_padding.narrow(2, total_len - pad_left, pad_left)?; + module_state.insert("previous".to_string(), new_previous); + + Ok(y) + } else { + self.conv.forward(x) + } + } + + pub fn weight(&self) -> &Tensor { + self.conv.weight() + } + + pub fn bias(&self) -> Option<&Tensor> { + self.conv.bias() + } +} + +#[derive(Clone)] +pub struct StreamingConvTranspose1d { + convtr: ConvTranspose1d, + stride: usize, + kernel_size: usize, + out_channels: usize, + name: String, +} + +impl StreamingConvTranspose1d { + #[allow(clippy::too_many_arguments)] + pub fn new( + in_channels: usize, + out_channels: usize, + kernel_size: usize, + stride: usize, + groups: usize, + bias: bool, + name: &str, + vb: VarBuilder, + ) -> Result { + let config = ConvTranspose1dConfig { + stride, + padding: 0, + output_padding: 0, + dilation: 1, + groups, + }; + let convtr = if bias { + candle_nn::conv_transpose1d( + in_channels, + out_channels, + kernel_size, + config, + vb.pp("convtr"), + )? + } else { + candle_nn::conv_transpose1d_no_bias( + in_channels, + out_channels, + kernel_size, + config, + vb.pp("convtr"), + )? + }; + + Ok(Self { + convtr, + stride, + kernel_size, + out_channels, + name: name.to_string(), + }) + } + + pub fn init_state( + &self, + batch_size: usize, + _sequence_length: usize, + device: &candle_core::Device, + ) -> Result> { + let mut state = HashMap::new(); + let k = self.kernel_size; + let s = self.stride; + if k > s { + let partial = + Tensor::zeros((batch_size, self.out_channels, k - s), DType::F32, device)?; + state.insert("partial".to_string(), partial); + } + Ok(state) + } + + pub fn forward( + &self, + x: &Tensor, + model_state: &mut ModelState, + _step: usize, + ) -> Result { + let (b, _c, t) = x.dims3()?; + let k = self.kernel_size; + let s = self.stride; + let trim = k.saturating_sub(s); + + // Auto-initialize state if missing + if !model_state.contains_key(&self.name) { + let init = self.init_state(b, t, x.device())?; + model_state.insert(self.name.clone(), init); + } + + let module_state = model_state.get_mut(&self.name).unwrap(); + + let mut y = self.convtr.forward(x)?; + + if trim > 0 { + if let Some(partial) = module_state.remove("partial") { + // y is (B, C, S*T + trim) + // We add partial to the start of y + let y_head = y.narrow(2, 0, trim)?; + let y_sum = (y_head + partial)?; + let y_tail = y.narrow(2, trim, y.dims()[2] - trim)?; + y = Tensor::cat(&[y_sum, y_tail], 2)?; + } + + // The last `trim` elements of `y` become the next `partial` + let len = y.dims()[2]; + let mut next_partial = y.narrow(2, len - trim, trim)?; + + // If bias exists, we need to subtract it from the partial state + // because it will be added again when we run the next forward pass. + if let Some(bias) = self.convtr.bias() { + let b_reshaped = bias.reshape((self.out_channels, 1))?; + next_partial = next_partial.broadcast_sub(&b_reshaped)?; + } + module_state.insert("partial".to_string(), next_partial); + + // The output we actually return is y MINUS the new partial tail + y = y.narrow(2, 0, len - trim)?; + } + + Ok(y) + } + + pub fn weight(&self) -> &Tensor { + self.convtr.weight() + } + + pub fn bias(&self) -> Option<&Tensor> { + self.convtr.bias() + } +} + +#[derive(Clone)] +pub struct ConvDownsample1d { + conv: StreamingConv1d, +} + +impl ConvDownsample1d { + pub fn new(stride: usize, dimension: usize, name: &str, vb: VarBuilder) -> Result { + let conv = StreamingConv1d::new( + dimension, + dimension, + 2 * stride, + stride, + 1, + 1, + false, + "replicate", + &format!("{}.conv", name), + vb.pp("conv"), + )?; + Ok(Self { conv }) + } + + pub fn init_state( + &self, + batch_size: usize, + sequence_length: usize, + device: &candle_core::Device, + ) -> Result> { + self.conv.init_state(batch_size, sequence_length, device) + } + + pub fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.conv.forward(x, model_state, step) + } +} + +#[derive(Clone)] +pub struct ConvTrUpsample1d { + convtr: StreamingConvTranspose1d, +} + +impl ConvTrUpsample1d { + pub fn new(stride: usize, dimension: usize, name: &str, vb: VarBuilder) -> Result { + let convtr = StreamingConvTranspose1d::new( + dimension, + dimension, + 2 * stride, + stride, + dimension, + false, + &format!("{}.convtr", name), + vb.pp("convtr"), + )?; + Ok(Self { convtr }) + } + + pub fn init_state( + &self, + batch_size: usize, + sequence_length: usize, + device: &candle_core::Device, + ) -> Result> { + self.convtr.init_state(batch_size, sequence_length, device) + } + + pub fn forward(&self, x: &Tensor, model_state: &mut ModelState, step: usize) -> Result { + self.convtr.forward(x, model_state, step) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/modules/mlp.rs b/mistralrs-core/src/speech_models/pockettts/modules/mlp.rs new file mode 100644 index 0000000000..c6c65e1ea0 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/modules/mlp.rs @@ -0,0 +1,418 @@ +use candle_core::{DType, Result, Tensor}; +use candle_nn::{Linear, Module, VarBuilder}; + +pub type StepFn = Box Result + Send + Sync>; + +#[derive(Clone)] +pub struct RMSNorm { + alpha: Tensor, + eps: f64, +} + +impl RMSNorm { + pub fn new(dim: usize, eps: f64, vb: VarBuilder) -> Result { + let alpha = vb.get(dim, "alpha")?; + Ok(Self { alpha, eps }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + let x_dtype = x.dtype(); + // Python's "RMSNorm" uses x.var() which IS mean((x - mean)²), NOT standard RMSNorm + // We must match Python exactly for parity + let var = x.var_keepdim(candle_core::D::Minus1)?; + let inv_rms = (var + self.eps)?.sqrt()?.recip()?; + let normalized = x.broadcast_mul(&inv_rms)?; + normalized.broadcast_mul(&self.alpha)?.to_dtype(x_dtype) + } +} + +#[derive(Clone)] +pub struct LayerNorm { + inner: candle_nn::LayerNorm, +} + +impl LayerNorm { + pub fn new(dim: usize, eps: f64, affine: bool, vb: VarBuilder) -> Result { + let (weight, bias) = if affine { + let weight = vb.get(dim, "weight")?; + let bias = vb.get(dim, "bias")?; + (Some(weight), Some(bias)) + } else { + (None, None) + }; + // candle_nn::LayerNorm::new takes (weight, bias, eps) + // If not affine, we can pass None for weight/bias but candle_nn::LayerNorm expects Tensor if present. + // Actually candle_nn has a layer_norm function. + // Let's use it. + let weight = weight.unwrap_or_else(|| Tensor::ones(dim, vb.dtype(), vb.device()).unwrap()); + let bias = bias.unwrap_or_else(|| Tensor::zeros(dim, vb.dtype(), vb.device()).unwrap()); + + Ok(Self { + inner: candle_nn::LayerNorm::new(weight, bias, eps), + }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + self.inner.forward(x) + } +} + +#[derive(Clone)] +pub struct LayerScale { + scale: Tensor, +} + +impl LayerScale { + pub fn new(channels: usize, _init: f32, vb: VarBuilder) -> Result { + let scale = vb.get(channels, "scale")?; + Ok(Self { scale }) + } + + pub fn forward(&self, x: &Tensor) -> Result { + x.broadcast_mul(&self.scale) + } +} + +#[derive(Clone)] +pub struct TimestepEmbedder { + lin1: Linear, + lin2: Linear, + norm: RMSNorm, + freqs: Tensor, +} + +impl TimestepEmbedder { + pub fn new( + hidden_size: usize, + frequency_embedding_size: usize, + max_period: f32, + vb: VarBuilder, + ) -> Result { + let lin1 = candle_nn::linear(frequency_embedding_size, hidden_size, vb.pp("mlp.0"))?; + let lin2 = candle_nn::linear(hidden_size, hidden_size, vb.pp("mlp.2"))?; + let norm = RMSNorm::new(hidden_size, 1e-5, vb.pp("mlp.3"))?; + + let half = frequency_embedding_size / 2; + let ds = Tensor::arange(0u32, half as u32, vb.device())?.to_dtype(DType::F32)?; + let freqs = ds + .affine(-(max_period.ln() as f64) / half as f64, 0.0)? + .exp()? + .to_dtype(vb.dtype())?; // Pre-convert to model dtype + + Ok(Self { + lin1, + lin2, + norm, + freqs, + }) + } + + pub fn forward(&self, t: &Tensor) -> Result { + // t is [B], freqs is [half] + // We need args to be [B, half] for MLP to process + let t = if t.dims().len() == 1 { + t.unsqueeze(1)? // [B] -> [B, 1] + } else { + t.clone() + }; + // args = t * freqs: [B, 1] * [half] -> [B, half] + let args = t.broadcast_mul(&self.freqs)?; + let cos = args.cos()?; + let sin = args.sin()?; + // [B, half] cat [B, half] -> [B, frequency_embedding_size] + let mut x = Tensor::cat(&[cos, sin], candle_core::D::Minus1)?; + + // Forward through MLP sequence: lin1 -> silu -> lin2 -> norm + x = self.lin1.forward(&x)?; + x = x.silu()?; + x = self.lin2.forward(&x)?; + x = self.norm.forward(&x)?; + + Ok(x) + } +} + +pub fn modulate(x: &Tensor, shift: &Tensor, scale: &Tensor) -> Result { + x.broadcast_mul(&(scale + 1.0)?)?.broadcast_add(shift) +} + +#[derive(Clone)] +pub struct ModulationParams { + pub shift: Tensor, + pub scale: Tensor, + pub gate: Option, +} + +#[derive(Clone)] +pub struct ResBlock { + in_ln: LayerNorm, + mlp_lin1: Linear, + mlp_lin2: Linear, + ada_ln_lin: Linear, +} + +impl ResBlock { + pub fn new(channels: usize, vb: VarBuilder) -> Result { + let in_ln = LayerNorm::new(channels, 1e-6, true, vb.pp("in_ln"))?; + let mlp_lin1 = candle_nn::linear(channels, channels, vb.pp("mlp.0"))?; + let mlp_lin2 = candle_nn::linear(channels, channels, vb.pp("mlp.2"))?; + let ada_ln_lin = candle_nn::linear(channels, 3 * channels, vb.pp("adaLN_modulation.1"))?; + Ok(Self { + in_ln, + mlp_lin1, + mlp_lin2, + ada_ln_lin, + }) + } + + pub fn forward(&self, x: &Tensor, modulation: &ModulationParams) -> Result { + let mut h = self.in_ln.forward(x)?; + h = modulate(&h, &modulation.shift, &modulation.scale)?; + h = self.mlp_lin1.forward(&h)?.silu()?; + h = self.mlp_lin2.forward(&h)?; + + if let Some(gate) = &modulation.gate { + x + h.broadcast_mul(gate) + } else { + x + h + } + } +} + +#[derive(Clone)] +pub struct FinalLayer { + norm_final: LayerNorm, + linear: Linear, + ada_ln_lin: Linear, +} + +impl FinalLayer { + pub fn new(model_channels: usize, out_channels: usize, vb: VarBuilder) -> Result { + let norm_final = LayerNorm::new(model_channels, 1e-6, false, vb.pp("norm_final"))?; + let linear = candle_nn::linear(model_channels, out_channels, vb.pp("linear"))?; + let ada_ln_lin = candle_nn::linear( + model_channels, + 2 * model_channels, + vb.pp("adaLN_modulation.1"), + )?; + Ok(Self { + norm_final, + linear, + ada_ln_lin, + }) + } + + pub fn forward(&self, x: &Tensor, modulation: &ModulationParams) -> Result { + let h = modulate( + &self.norm_final.forward(x)?, + &modulation.shift, + &modulation.scale, + )?; + self.linear.forward(&h) + } +} + +#[derive(Clone)] +pub struct SimpleMLPAdaLN { + time_embeds: Vec, + cond_embed: Linear, + input_proj: Linear, + res_blocks: Vec, + final_layer: FinalLayer, + num_time_conds: usize, +} + +impl SimpleMLPAdaLN { + #[allow(clippy::too_many_arguments)] + pub fn new( + in_channels: usize, + model_channels: usize, + out_channels: usize, + cond_channels: usize, + num_res_blocks: usize, + num_time_conds: usize, + max_period: f32, + vb: VarBuilder, + ) -> Result { + let mut time_embeds = Vec::new(); + for i in 0..num_time_conds { + time_embeds.push(TimestepEmbedder::new( + model_channels, + 256, + max_period, + vb.pp(format!("time_embed.{}", i)), + )?); + } + + let cond_embed = candle_nn::linear(cond_channels, model_channels, vb.pp("cond_embed"))?; + let input_proj = candle_nn::linear(in_channels, model_channels, vb.pp("input_proj"))?; + + let mut res_blocks = Vec::new(); + for i in 0..num_res_blocks { + res_blocks.push(ResBlock::new( + model_channels, + vb.pp(format!("res_blocks.{}", i)), + )?); + } + + let final_layer = FinalLayer::new(model_channels, out_channels, vb.pp("final_layer"))?; + + Ok(Self { + time_embeds, + cond_embed, + input_proj, + res_blocks, + final_layer, + num_time_conds, + }) + } + + pub fn forward(&self, c: &Tensor, s: &Tensor, t: &Tensor, x: &Tensor) -> Result { + let c_emb = self.embed_condition(c)?; + self.forward_step(x, &c_emb, s, t) + } + + pub fn embed_condition(&self, c: &Tensor) -> Result { + self.cond_embed.forward(c) + } + + pub fn forward_step( + &self, + x: &Tensor, + c_emb: &Tensor, + s: &Tensor, + t: &Tensor, + ) -> Result { + let y = (self.time_embeds[0].forward(s)? + self.time_embeds[1].forward(t)?)?; + let t_combined = (y / self.num_time_conds as f64)?; + + // Compute modulations on the fly for non-cached call + let mod_vec = self.precompute_modulations(c_emb, &t_combined)?; + self.forward_step_cached(x, &mod_vec[0]) + } +} + +impl SimpleMLPAdaLN { + pub fn compute_time_embeddings( + &self, + num_steps: usize, + device: &candle_core::Device, + dtype: DType, + ) -> Result { + let mut embeddings = Vec::with_capacity(num_steps); + for i in 0..num_steps { + let s = i as f64 / num_steps as f64; + let t = (i + 1) as f64 / num_steps as f64; + + // 1D Tensors [1] + let s_tensor = Tensor::new(&[s as f32], device)?.to_dtype(dtype)?; + let t_tensor = Tensor::new(&[t as f32], device)?.to_dtype(dtype)?; + + let t0 = self.time_embeds[0].forward(&s_tensor)?; + let t1 = self.time_embeds[1].forward(&t_tensor)?; + let t_combined = ((t0 + t1)? / self.num_time_conds as f64)?; + embeddings.push(t_combined); + } + // stack of [1, 512] -> [num_steps, 1, 512] + // squeeze(1) -> [num_steps, 512] + Tensor::stack(&embeddings, 0)?.squeeze(1) + } + + #[allow(clippy::needless_range_loop)] + pub fn precompute_modulations( + &self, + c_emb: &Tensor, + time_embeddings: &Tensor, + ) -> Result>> { + // c_emb: [1, 512], time_embeddings: [8, 512] + let num_steps = time_embeddings.dim(0)?; + let y = time_embeddings.broadcast_add(c_emb)?; // [8, 512] + let y_silu = y.silu()?; + + let mut all_step_modulations = + vec![Vec::with_capacity(self.res_blocks.len() + 1); num_steps]; + + // ResBlocks + for block in &self.res_blocks { + let mod_batch = block.ada_ln_lin.forward(&y_silu)?; // [8, 1536] + let dim = mod_batch.dim(candle_core::D::Minus1)? / 3; + + for s in 0..num_steps { + let modulation = mod_batch.narrow(0, s, 1)?; // [1, 1536] + let shift = modulation.narrow(candle_core::D::Minus1, 0, dim)?; + let scale = modulation.narrow(candle_core::D::Minus1, dim, dim)?; + let gate = modulation.narrow(candle_core::D::Minus1, 2 * dim, dim)?; + all_step_modulations[s].push(ModulationParams { + shift, + scale, + gate: Some(gate), + }); + } + } + + // Final layer + let mod_batch = self.final_layer.ada_ln_lin.forward(&y_silu)?; // [8, 1024] + let dim = mod_batch.dim(candle_core::D::Minus1)? / 2; + for s in 0..num_steps { + let modulation = mod_batch.narrow(0, s, 1)?; // [1, 1024] + let shift = modulation.narrow(candle_core::D::Minus1, 0, dim)?; + let scale = modulation.narrow(candle_core::D::Minus1, dim, dim)?; + all_step_modulations[s].push(ModulationParams { + shift, + scale, + gate: None, + }); + } + + Ok(all_step_modulations) + } + + pub fn forward_step_cached( + &self, + x: &Tensor, + modulations: &[ModulationParams], + ) -> Result { + let mut x = self.input_proj.forward(x)?; + + for (i, block) in self.res_blocks.iter().enumerate() { + x = block.forward(&x, &modulations[i])?; + } + + self.final_layer + .forward(&x, &modulations[self.res_blocks.len()]) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use candle_core::{Device, Tensor}; + use candle_nn::VarBuilder; + use std::collections::HashMap; + + #[test] + fn test_rmsnorm_parity() -> Result<()> { + let device = Device::Cpu; + let mut map = HashMap::new(); + map.insert( + "alpha".to_string(), + Tensor::ones((4,), DType::F32, &device)?, + ); + let vb = VarBuilder::from_tensors(map, DType::F32, &device); + let norm = RMSNorm::new(4, 1e-5, vb)?; + + // Input: [[1.0, 2.0, 3.0, 4.0]] + let x = Tensor::new(&[[1.0f32, 2.0, 3.0, 4.0]], &device)?; + let y = norm.forward(&x)?; + + // Python's "RMSNorm" uses x.var() = mean((x - mean)²) + // mean = 2.5, var = ((1-2.5)² + (2-2.5)² + (3-2.5)² + (4-2.5)²) / 3 = 1.6667 (Bessel) + // rsqrt(1.6667 + 1e-5) ≈ 0.7746 + // output = x * 0.7746 = [0.7746, 1.5492, 2.3238, 3.0984] + let expected = Tensor::new(&[[0.7746f32, 1.5492, 2.3238, 3.0984]], &device)?; + + let diff = (y - expected)?.abs()?.max_all()?.to_scalar::()?; + assert!(diff < 1e-3, "RMSNorm parity failed: diff={}", diff); + Ok(()) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/modules/mod.rs b/mistralrs-core/src/speech_models/pockettts/modules/mod.rs new file mode 100644 index 0000000000..718234c4ad --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/modules/mod.rs @@ -0,0 +1,5 @@ +pub mod attention; +pub mod conv; +pub mod mlp; +pub mod rope; +pub mod sdpa; diff --git a/mistralrs-core/src/speech_models/pockettts/modules/rope.rs b/mistralrs-core/src/speech_models/pockettts/modules/rope.rs new file mode 100644 index 0000000000..dd635b75e3 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/modules/rope.rs @@ -0,0 +1,79 @@ +use candle_core::{DType, Result, Tensor}; + +#[derive(Debug, Clone)] +pub struct RotaryEmbedding { + inv_freq: Tensor, +} + +impl RotaryEmbedding { + pub fn new(max_period: f32, head_dim: usize, device: &candle_core::Device) -> Result { + let d = head_dim / 2; + let ds = Tensor::arange(0u32, d as u32, device)?.to_dtype(DType::F32)?; + let inv_freq = ds + .affine((-max_period.ln() * 2.0 / head_dim as f32) as f64, 0.0)? + .exp()?; + Ok(Self { inv_freq }) + } + + pub fn forward(&self, q: &Tensor, k: &Tensor, offset: usize) -> Result<(Tensor, Tensor)> { + let (b, t, h, d_full) = q.dims4()?; + let (_bk, _tk, hk, _dk) = k.dims4()?; + let d = d_full / 2; + let dev = q.device(); + + // ts = (arange(T) + offset).view(-1, 1, 1) + let ts = if t == 1 { + Tensor::new(&[offset as f32], dev)? + } else { + Tensor::arange(0u32, t as u32, dev)? + .to_dtype(DType::F32)? + .affine(1.0, offset as f64)? + } + .reshape((t, 1, 1))?; + + // freqs * ts -> shape (t, 1, d) + let freqs_ts = self.inv_freq.reshape((1, 1, d))?.broadcast_mul(&ts)?; + let cos = freqs_ts.cos()?; + let sin = freqs_ts.sin()?; + + // Reshape q and k to (b, t, h, d, 2) + let q = q.reshape((b, t, h, d, 2))?; + let k = k.reshape((b, t, hk, d, 2))?; + + let qr = q.narrow(4, 0, 1)?.squeeze(4)?; + let qi = q.narrow(4, 1, 1)?.squeeze(4)?; + let kr = k.narrow(4, 0, 1)?.squeeze(4)?; + let ki = k.narrow(4, 1, 1)?.squeeze(4)?; + + // qor = qr * cos - qi * sin + // qoi = qr * sin + qi * cos + let qor = (qr.broadcast_mul(&cos)? - qi.broadcast_mul(&sin)?)?; + let qoi = (qr.broadcast_mul(&sin)? + qi.broadcast_mul(&cos)?)?; + + let kor = (kr.broadcast_mul(&cos)? - ki.broadcast_mul(&sin)?)?; + let koi = (kr.broadcast_mul(&sin)? + ki.broadcast_mul(&cos)?)?; + + let qo = Tensor::stack(&[qor, qoi], 4)?.reshape((b, t, h, d_full))?; + let ko = Tensor::stack(&[kor, koi], 4)?.reshape((b, t, hk, d_full))?; + + Ok((qo, ko)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use candle_core::{Device, Tensor}; + + #[test] + fn test_rope_shape() -> Result<()> { + let device = Device::Cpu; + let q = Tensor::zeros((1, 10, 4, 32), DType::F32, &device)?; + let k = Tensor::zeros((1, 10, 4, 32), DType::F32, &device)?; + let rope = RotaryEmbedding::new(10000.0, 32, &device)?; + let (qo, ko) = rope.forward(&q, &k, 0)?; + assert_eq!(qo.dims(), &[1, 10, 4, 32]); + assert_eq!(ko.dims(), &[1, 10, 4, 32]); + Ok(()) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/modules/sdpa.rs b/mistralrs-core/src/speech_models/pockettts/modules/sdpa.rs new file mode 100644 index 0000000000..1cbf1d97e0 --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/modules/sdpa.rs @@ -0,0 +1,385 @@ +use candle_core::{Result, Tensor, D}; + +#[inline] +fn can_skip_mask_for_single_query( + q_len: usize, + kv_len: usize, + is_causal: bool, + context_window: Option, +) -> bool { + if !is_causal || q_len != 1 { + return false; + } + + match context_window { + None => true, + Some(ctx) => kv_len <= ctx, + } +} + +/// Memory-efficient Scaled Dot Product Attention +/// +/// Computes `softmax(Q @ K.T / sqrt(d) + mask) @ V` using tiling on the query dimension +/// to avoid materializing the full N x N attention matrix. +/// +/// # Arguments +/// * `q` - Query tensor of shape [Batch, Heads, Q_Len, Dim] +/// * `k` - Key tensor of shape [Batch, Heads, KV_Len, Dim] +/// * `v` - Value tensor of shape [Batch, Heads, KV_Len, Dim] +/// * `scale` - Scaling factor (usually 1 / sqrt(dim)) +/// * `is_causal` - Whether to apply causal masking +/// * `context_window` - Optional context window size for local attention +/// +/// # Returns +/// * Tensor of shape [Batch, Heads, Q_Len, Dim] +#[inline] +pub fn sdpa( + q: &Tensor, + k: &Tensor, + v: &Tensor, + scale: f64, + is_causal: bool, + context_window: Option, +) -> Result { + let q = q.contiguous()?; + let k = k.contiguous()?; + let v = v.contiguous()?; + let (_b, _h, q_len, _dim) = q.dims4()?; + let kv_len = k.dims()[2]; + + // Adaptive strategy: + // For small Q (decoding, chunked prefill), tiling overhead hurts performance. + // Use naive implementation if Q is small enough. + // Benchmark showed naive is faster for Q=1 and comparable for Q=50/64. + const TILING_THRESHOLD: usize = 512; + + let k_t = k.transpose(2, 3)?.contiguous()?; // [B, H, D, S] + + if q_len < TILING_THRESHOLD { + // Naive path (no tiling) + let scores = (q.matmul(&k_t)? * scale)?; + + let scores = if can_skip_mask_for_single_query(q_len, kv_len, is_causal, context_window) { + scores + } else if is_causal || context_window.is_some() { + let mask = generate_mask_chunk( + 0, + q_len, + kv_len, + q_len, + is_causal, + context_window, + q.device(), + )?; + scores.broadcast_add(&mask)? + } else { + scores + }; + + let probs = candle_nn::ops::softmax(&scores, D::Minus1)?; + return probs.matmul(&v); + } + + // Tiled path for large Q + // Always tile if sequence length is significant to avoid N^2 mask allocation + let block_size = 128; // Tiling size for Q dimension. + + let mut outputs = Vec::new(); + + for start in (0..q_len).step_by(block_size) { + let end = std::cmp::min(start + block_size, q_len); + let len = end - start; + + // Slice Q: [B, H, Block, D] + let q_chunk = q.narrow(2, start, len)?; + + // Compute scores: [B, H, Block, S] = [B, H, Block, D] @ [B, H, D, S] + let scores = (q_chunk.matmul(&k_t)? * scale)?; + + // Generate and apply mask on-the-fly for this chunk + let scores = if is_causal || context_window.is_some() { + let mask_chunk = generate_mask_chunk( + start, + len, + kv_len, + q_len, + is_causal, + context_window, + q.device(), + )?; + scores.broadcast_add(&mask_chunk)? + } else { + scores + }; + + // Softmax + let probs = candle_nn::ops::softmax(&scores, D::Minus1)?; + + // Output chunk: [B, H, Block, D] = [B, H, Block, S] @ [B, H, S, D] + let out_chunk = probs.matmul(&v)?; + + outputs.push(out_chunk); + } + + // Cat along Q dimension (dim 2) + Tensor::cat(&outputs, 2) +} + +/// Helper to generate a mask chunk for a specific query range using vectorized operations +fn generate_mask_chunk( + start_q: usize, + num_q: usize, + k_len: usize, + total_q_len: usize, + is_causal: bool, + context_window: Option, + device: &candle_core::Device, +) -> Result { + let shift = k_len.saturating_sub(total_q_len); + + // pos_q: [num_q, 1] + let pos_q = (Tensor::arange(0u32, num_q as u32, device)? + .to_dtype(candle_core::DType::F32)? + .affine(1.0, (start_q + shift) as f64)? + .reshape((num_q, 1)))?; + + // pos_k: [1, k_len] + let pos_k = Tensor::arange(0u32, k_len as u32, device)? + .to_dtype(candle_core::DType::F32)? + .reshape((1, k_len))?; + + let mut mask = Tensor::zeros((num_q, k_len), candle_core::DType::F32, device)?; + + if is_causal { + let is_future = pos_k.broadcast_gt(&pos_q)?; + mask = is_future.where_cond( + &Tensor::full(f32::NEG_INFINITY, (num_q, k_len), device)?, + &mask, + )?; + } + + if let Some(ctx) = context_window { + let limit = pos_q.broadcast_sub(&Tensor::full(ctx as f32, (num_q, 1), device)?)?; + let is_out = pos_k.broadcast_le(&limit)?; + mask = is_out.where_cond( + &Tensor::full(f32::NEG_INFINITY, (num_q, k_len), device)?, + &mask, + )?; + } + + mask.reshape((1, 1, num_q, k_len)) +} + +/// Chunked version of SDPA that accepts a list of Key/Value pointers +/// to avoid concatenating the full KV cache. +pub fn sdpa_chunked( + q: &Tensor, + k_chunks: &[Tensor], + v_chunks: &[Tensor], + scale: f64, + is_causal: bool, + context_window: Option, +) -> Result { + if k_chunks.is_empty() { + let (_b, h, _q, d) = q.dims4()?; + return Tensor::zeros((_b, h, _q, d), q.dtype(), q.device()); + } + + let device = q.device(); + let dtype = q.dtype(); + let q = q.contiguous()?; + let (b, h, q_len, d) = q.dims4()?; + + // Ensure all KV chunks are contiguous for CPU matmul compatibility + let k_chunks: Vec = k_chunks + .iter() + .map(|t| t.contiguous()) + .collect::>()?; + let v_chunks: Vec = v_chunks + .iter() + .map(|t| t.contiguous()) + .collect::>()?; + + // Fast path for single chunk + if k_chunks.len() == 1 { + let k_t = k_chunks[0].transpose(2, 3)?.contiguous()?; + let scores = (q.matmul(&k_t)? * scale)?; + let kv_len = k_chunks[0].dims()[2]; + + let masked_scores = + if can_skip_mask_for_single_query(q_len, kv_len, is_causal, context_window) { + scores + } else if is_causal || context_window.is_some() { + let mask = generate_mask_chunk( + 0, + q_len, + kv_len, + q_len, + is_causal, + context_window, + device, + )?; + scores.broadcast_add(&mask)? + } else { + scores + }; + + let probs = candle_nn::ops::softmax(&masked_scores, D::Minus1)?; + return probs.matmul(&v_chunks[0]); + } + + // 1. Compute scores against all K chunks + let mut score_chunks = Vec::with_capacity(k_chunks.len()); + let mut total_kv_len = 0; + + for k_chunk in k_chunks { + total_kv_len += k_chunk.dims()[2]; + let k_t = k_chunk.transpose(2, 3)?.contiguous()?; + let score_chunk = (q.matmul(&k_t)? * scale)?; + score_chunks.push(score_chunk); + } + + // 2. Concatenate scores to apply global Softmax + let all_scores = Tensor::cat(&score_chunks, 3)?; + + // 3. Apply masking + let masked_scores = + if can_skip_mask_for_single_query(q_len, total_kv_len, is_causal, context_window) { + all_scores + } else if is_causal || context_window.is_some() { + let mask = generate_mask_chunk( + 0, + q_len, + total_kv_len, + q_len, + is_causal, + context_window, + device, + )?; + all_scores.broadcast_add(&mask)? + } else { + all_scores + }; + + // 4. Softmax + let probs = candle_nn::ops::softmax(&masked_scores, D::Minus1)?; + + // 5. Compute Weighted Sum: Probs @ V + let mut output = Tensor::zeros((b, h, q_len, d), dtype, device)?; + + let mut offset = 0; + for v_chunk in v_chunks { + let chunk_len = v_chunk.dims()[2]; + let probs_chunk = probs.narrow(3, offset, chunk_len)?.contiguous()?; + let out_chunk = probs_chunk.matmul(&v_chunk)?; + output = (output + out_chunk)?; + offset += chunk_len; + } + + Ok(output) +} + +#[cfg(test)] +mod tests { + use super::*; + use candle_core::Device; + + #[test] + fn test_generate_mask_chunk_causal() -> Result<()> { + let device = Device::Cpu; + // q_len = 1, k_len = 5, total_q = 1 + // shift = 4. pos_q = 4. pos_k = 0..5. + // is_future = j > 4. No futures. + let mask = generate_mask_chunk(0, 1, 5, 1, true, None, &device)?; + let mask_data = mask.flatten_all()?.to_vec1::()?; + assert_eq!(mask_data, vec![0.0, 0.0, 0.0, 0.0, 0.0]); + + // q_len = 3, k_len = 3, total_q = 3 (prefill) + // shift = 0. pos_q = 0..3. pos_k = 0..3. + let mask = generate_mask_chunk(0, 3, 3, 3, true, None, &device)?; + let mask_data = mask.reshape((3, 3))?.to_vec2::()?; + // Row 0: pos_q=0. k=0 ok, k=1 future, k=2 future + assert_eq!( + mask_data[0], + vec![0.0, f32::NEG_INFINITY, f32::NEG_INFINITY] + ); + // Row 1: pos_q=1. k=0,1 ok, k=2 future + assert_eq!(mask_data[1], vec![0.0, 0.0, f32::NEG_INFINITY]); + // Row 2: pos_q=2. k=0,1,2 ok + assert_eq!(mask_data[2], vec![0.0, 0.0, 0.0]); + + Ok(()) + } + + #[test] + fn test_generate_mask_chunk_window() -> Result<()> { + let device = Device::Cpu; + // ctx = 2. pos_q = 5. k_len = 6. total_q = 1. + // shift = 5. pos_q = 5. pos_k = 0..6. + // limit = 5 - 2 = 3. + // is_out = j <= 3 -> 0,1,2,3 masked. 4,5 ok. + let mask = generate_mask_chunk(0, 1, 6, 1, false, Some(2), &device)?; + let mask_data = mask.flatten_all()?.to_vec1::()?; + assert_eq!( + mask_data, + vec![ + f32::NEG_INFINITY, + f32::NEG_INFINITY, + f32::NEG_INFINITY, + f32::NEG_INFINITY, + 0.0, + 0.0 + ] + ); + + Ok(()) + } + + #[test] + fn test_can_skip_mask_for_single_query() { + assert!(can_skip_mask_for_single_query(1, 64, true, None)); + assert!(can_skip_mask_for_single_query(1, 64, true, Some(64))); + assert!(!can_skip_mask_for_single_query(1, 65, true, Some(64))); + assert!(!can_skip_mask_for_single_query(2, 64, true, None)); + assert!(!can_skip_mask_for_single_query(1, 64, false, None)); + } + + #[test] + fn test_sdpa_handles_non_contiguous_inputs() -> Result<()> { + let device = Device::Cpu; + let scale = 1.0 / (64f64).sqrt(); + + // Q built from a transposed view. + let q_base = Tensor::zeros((1, 8, 64, 128), candle_core::DType::F32, &device)?; + let q = q_base.transpose(2, 3)?; // [1, 8, 128, 64] + + // K is contiguous but K^T in SDPA is not, unless re-materialized. + let k = Tensor::zeros((1, 8, 1600, 64), candle_core::DType::F32, &device)?; + + // V built from transpose + narrow view. + let v_base = Tensor::zeros((1, 8, 64, 1601), candle_core::DType::F32, &device)?; + let v = v_base.transpose(2, 3)?.narrow(2, 1, 1600)?; // [1, 8, 1600, 64] + + let out = sdpa(&q, &k, &v, scale, true, None)?; + assert_eq!(out.dims(), &[1, 8, 128, 64]); + Ok(()) + } + + #[test] + fn test_sdpa_chunked_handles_non_contiguous_value_chunks() -> Result<()> { + let device = Device::Cpu; + let scale = 1.0 / (32f64).sqrt(); + + let q = Tensor::zeros((1, 4, 64, 32), candle_core::DType::F32, &device)?; + let k_full = Tensor::zeros((1, 4, 320, 32), candle_core::DType::F32, &device)?; + let v_base = Tensor::zeros((1, 4, 32, 321), candle_core::DType::F32, &device)?; + let v_full = v_base.transpose(2, 3)?.narrow(2, 1, 320)?; // [1, 4, 320, 32] + + let k_chunks = vec![k_full.narrow(2, 0, 160)?, k_full.narrow(2, 160, 160)?]; + let v_chunks = vec![v_full.narrow(2, 0, 160)?, v_full.narrow(2, 160, 160)?]; + + let out = sdpa_chunked(&q, &k_chunks, &v_chunks, scale, true, None)?; + assert_eq!(out.dims(), &[1, 4, 64, 32]); + Ok(()) + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/pause.rs b/mistralrs-core/src/speech_models/pockettts/pause.rs new file mode 100644 index 0000000000..94ac2b010f --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/pause.rs @@ -0,0 +1,249 @@ +//! Pause/silence handling for text-to-speech +//! +//! Supports: +//! - Explicit pause markers: `[pause:Xms]` or `[pause:Xs]` +//! - Natural pauses from punctuation: `...`, `,` + +use regex::Regex; +use std::sync::LazyLock; + +/// Pause marker found in text +#[derive(Debug, Clone, PartialEq)] +pub struct PauseMarker { + /// Original text that was matched + pub original: String, + /// Duration in milliseconds + pub duration_ms: u32, + /// Position in the original text (byte offset) + pub position: usize, +} + +/// Default pause durations (in milliseconds) for punctuation +pub mod defaults { + /// Ellipsis "..." pause duration + pub const ELLIPSIS_MS: u32 = 500; + /// Comma pause duration + pub const COMMA_MS: u32 = 200; + /// Period/sentence end pause duration + pub const PERIOD_MS: u32 = 400; + /// Semicolon pause duration + pub const SEMICOLON_MS: u32 = 300; +} + +// Regex patterns for pause parsing +static EXPLICIT_PAUSE_REGEX: LazyLock = LazyLock::new(|| { + // Matches [pause:500ms] or [pause:1s] or [pause:1.5s] + Regex::new(r"\[pause:(\d+(?:\.\d+)?)(ms|s)\]").unwrap() +}); + +static ELLIPSIS_REGEX: LazyLock = LazyLock::new(|| Regex::new(r"\.{3,}").unwrap()); + +/// Parse explicit pause markers from text +/// +/// # Example +/// ``` +/// use pocket_tts::pause::parse_explicit_pauses; +/// +/// let pauses = parse_explicit_pauses("Hello [pause:500ms] world [pause:1s] done"); +/// assert_eq!(pauses.len(), 2); +/// assert_eq!(pauses[0].duration_ms, 500); +/// assert_eq!(pauses[1].duration_ms, 1000); +/// ``` +pub fn parse_explicit_pauses(text: &str) -> Vec { + EXPLICIT_PAUSE_REGEX + .captures_iter(text) + .filter_map(|cap| { + let full_match = cap.get(0)?; + let value: f64 = cap.get(1)?.as_str().parse().ok()?; + let unit = cap.get(2)?.as_str(); + + let duration_ms = match unit { + "ms" => value as u32, + "s" => (value * 1000.0) as u32, + _ => return None, + }; + + Some(PauseMarker { + original: full_match.as_str().to_string(), + duration_ms, + position: full_match.start(), + }) + }) + .collect() +} + +/// Parse natural pauses from punctuation +pub fn parse_natural_pauses(text: &str) -> Vec { + let mut pauses = Vec::new(); + + // Find ellipses + for cap in ELLIPSIS_REGEX.find_iter(text) { + pauses.push(PauseMarker { + original: cap.as_str().to_string(), + duration_ms: defaults::ELLIPSIS_MS, + position: cap.start(), + }); + } + + // Find commas (but not inside numbers like "1,000") + for (i, c) in text.char_indices() { + if c == ',' { + // Check if it's not surrounded by digits + let prev_is_digit = + i > 0 && text[..i].chars().last().is_some_and(|c| c.is_ascii_digit()); + let next_is_digit = text[(i + 1)..] + .chars() + .next() + .is_some_and(|c| c.is_ascii_digit()); + + if !prev_is_digit || !next_is_digit { + pauses.push(PauseMarker { + original: ",".to_string(), + duration_ms: defaults::COMMA_MS, + position: i, + }); + } + } + } + + // Sort by position + pauses.sort_by_key(|p| p.position); + pauses +} + +/// Remove pause markers from text, returning clean text for TTS +pub fn strip_pause_markers(text: &str) -> String { + EXPLICIT_PAUSE_REGEX.replace_all(text, " ").to_string() +} + +/// Parsed text with pause information +#[derive(Debug, Clone)] +pub struct ParsedText { + /// Text with pause markers removed + pub clean_text: String, + /// All pause markers (explicit + natural) with adjusted positions + pub pauses: Vec, +} + +/// Parse text for all pause markers (explicit and natural) +pub fn parse_text_with_pauses(text: &str) -> ParsedText { + // First, find explicit pauses in original text + let mut all_pauses = parse_explicit_pauses(text); + + // Strip explicit markers to get clean text + let clean_text = strip_pause_markers(text); + + // Find natural pauses in clean text + let natural_pauses = parse_natural_pauses(&clean_text); + + // Note: positions in all_pauses are relative to original text + // We need to adjust them to the clean text + // For simplicity, we'll recalculate based on clean text positions + + // Clear and rebuild with correct positions + all_pauses.clear(); + all_pauses.extend(natural_pauses); + + // Re-parse explicit pauses and calculate where they would be in clean text + let mut offset = 0; + for cap in EXPLICIT_PAUSE_REGEX.captures_iter(text) { + let full_match = cap.get(0).unwrap(); + let original_pos = full_match.start(); + let adjusted_pos = original_pos.saturating_sub(offset); + let value: f64 = cap.get(1).unwrap().as_str().parse().unwrap_or(0.0); + let unit = cap.get(2).unwrap().as_str(); + + let duration_ms = match unit { + "ms" => value as u32, + "s" => (value * 1000.0) as u32, + _ => 0, + }; + + if duration_ms > 0 { + all_pauses.push(PauseMarker { + original: full_match.as_str().to_string(), + duration_ms, + position: adjusted_pos, + }); + } + + offset += full_match.len() - 1; // -1 for the space we replace with + } + + // Sort by position + all_pauses.sort_by_key(|p| p.position); + + ParsedText { + clean_text, + pauses: all_pauses, + } +} + +/// Calculate the number of silence samples for a given duration +pub fn silence_samples(duration_ms: u32, sample_rate: u32) -> usize { + ((duration_ms as u64 * sample_rate as u64) / 1000) as usize +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_explicit_pause_ms() { + let pauses = parse_explicit_pauses("Hello [pause:500ms] world"); + assert_eq!(pauses.len(), 1); + assert_eq!(pauses[0].duration_ms, 500); + assert_eq!(pauses[0].original, "[pause:500ms]"); + } + + #[test] + fn test_parse_explicit_pause_seconds() { + let pauses = parse_explicit_pauses("Test [pause:1s] and [pause:1.5s]"); + assert_eq!(pauses.len(), 2); + assert_eq!(pauses[0].duration_ms, 1000); + assert_eq!(pauses[1].duration_ms, 1500); + } + + #[test] + fn test_parse_ellipsis() { + let pauses = parse_natural_pauses("Hello... world"); + assert_eq!(pauses.len(), 1); + assert_eq!(pauses[0].duration_ms, defaults::ELLIPSIS_MS); + } + + #[test] + fn test_parse_comma() { + let pauses = parse_natural_pauses("Hello, world"); + assert_eq!(pauses.len(), 1); + assert_eq!(pauses[0].duration_ms, defaults::COMMA_MS); + } + + #[test] + fn test_comma_in_number_ignored() { + let pauses = parse_natural_pauses("That costs 1,000 dollars"); + // The comma in 1,000 should be ignored + assert_eq!(pauses.len(), 0); + } + + #[test] + fn test_strip_pause_markers() { + let clean = strip_pause_markers("Hello [pause:500ms] world [pause:1s] done"); + assert_eq!(clean, "Hello world done"); + } + + #[test] + fn test_parse_text_with_pauses() { + let parsed = parse_text_with_pauses("Hello... [pause:500ms] world, done"); + assert_eq!(parsed.clean_text, "Hello... world, done"); + // Should have: ellipsis, explicit pause, comma + assert_eq!(parsed.pauses.len(), 3); + } + + #[test] + fn test_silence_samples() { + // 500ms at 24kHz = 12000 samples + assert_eq!(silence_samples(500, 24000), 12000); + // 1s at 24kHz = 24000 samples + assert_eq!(silence_samples(1000, 24000), 24000); + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/tts_model.rs b/mistralrs-core/src/speech_models/pockettts/tts_model.rs new file mode 100644 index 0000000000..89ed813cec --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/tts_model.rs @@ -0,0 +1,442 @@ +//! High-level pocket-tts pipeline: text -> latents (FlowLM) -> audio (Mimi codec). +//! Ported from the upstream `pocket-tts` crate against candle 0.11. The +//! `without-voice-cloning` checkpoint runs with an empty voice state (default speaker). + +use std::path::Path; + +use super::conditioners::text::LUTConditioner; +use super::config::{defaults, PocketTtsConfig}; +use super::models::flow_lm::FlowLMModel; +use super::models::mimi::MimiModel; +use super::models::seanet::{SEANetDecoder, SEANetEncoder}; +use super::models::transformer::{ProjectedTransformer, StreamingTransformer}; +use super::modules::mlp::SimpleMLPAdaLN; +use super::voice_state::{increment_steps, init_states, ModelState}; + +use anyhow::Result; +use candle_core::{DType, Device, Tensor}; +use candle_nn::VarBuilder; + +#[derive(Clone)] +pub struct TTSModel { + pub flow_lm: FlowLMModel, + pub mimi: MimiModel, + pub conditioner: LUTConditioner, + pub temp: f32, + pub lsd_decode_steps: usize, + pub eos_threshold: f32, + pub sample_rate: usize, + pub dim: usize, + pub ldim: usize, + pub device: Device, +} + +impl TTSModel { + pub fn new(config: &PocketTtsConfig, vb: VarBuilder, tokenizer_path: &Path) -> Result { + let device = vb.device().clone(); + + let conditioner = LUTConditioner::new( + config.flow_lm.lookup_table.n_bins, + tokenizer_path, + config.flow_lm.lookup_table.dim, + config.flow_lm.transformer.d_model, + vb.pp("flow_lm.conditioner"), + )?; + + let dim = config.flow_lm.transformer.d_model; + let ldim = config.mimi.quantizer.dimension; + let hidden_dim = dim * config.flow_lm.transformer.hidden_scale; + + let flow_net = SimpleMLPAdaLN::new( + ldim, + config.flow_lm.flow.dim, + ldim, + dim, + config.flow_lm.flow.depth, + 2, + config.flow_lm.transformer.max_period as f32, + vb.pp("flow_lm.flow_net"), + )?; + + let transformer = StreamingTransformer::new( + dim, + config.flow_lm.transformer.num_heads, + config.flow_lm.transformer.num_layers, + None, + hidden_dim, + None, + config.flow_lm.transformer.max_period as f32, + "kv", + "flow_lm.transformer", + vb.pp("flow_lm.transformer"), + )?; + + let flow_lm = FlowLMModel::new(flow_net, transformer, ldim, dim, vb.pp("flow_lm"))?; + + let seanet_cfg = &config.mimi.seanet; + let encoder = SEANetEncoder::new( + seanet_cfg.channels, + seanet_cfg.dimension, + seanet_cfg.n_filters, + seanet_cfg.n_residual_layers, + &seanet_cfg.ratios, + seanet_cfg.kernel_size, + seanet_cfg.last_kernel_size, + seanet_cfg.residual_kernel_size, + seanet_cfg.dilation_base, + &seanet_cfg.pad_mode, + seanet_cfg.compress, + "mimi.encoder", + vb.pp("mimi.encoder"), + )?; + + let decoder = SEANetDecoder::new( + seanet_cfg.channels, + seanet_cfg.dimension, + seanet_cfg.n_filters, + seanet_cfg.n_residual_layers, + &seanet_cfg.ratios, + seanet_cfg.kernel_size, + seanet_cfg.last_kernel_size, + seanet_cfg.residual_kernel_size, + seanet_cfg.dilation_base, + &seanet_cfg.pad_mode, + seanet_cfg.compress, + "mimi.decoder", + vb.pp("mimi.decoder"), + )?; + + let mimi_tr_cfg = &config.mimi.transformer; + let encoder_transformer = ProjectedTransformer::new( + mimi_tr_cfg.input_dimension, + mimi_tr_cfg.output_dimensions.clone(), + mimi_tr_cfg.d_model, + mimi_tr_cfg.num_heads, + mimi_tr_cfg.num_layers, + mimi_tr_cfg.layer_scale as f32, + mimi_tr_cfg.context, + mimi_tr_cfg.max_period as f32, + mimi_tr_cfg.dim_feedforward, + "mimi.encoder_transformer", + vb.pp("mimi.encoder_transformer"), + )?; + + let decoder_transformer = ProjectedTransformer::new( + mimi_tr_cfg.input_dimension, + mimi_tr_cfg.output_dimensions.clone(), + mimi_tr_cfg.d_model, + mimi_tr_cfg.num_heads, + mimi_tr_cfg.num_layers, + mimi_tr_cfg.layer_scale as f32, + mimi_tr_cfg.context, + mimi_tr_cfg.max_period as f32, + mimi_tr_cfg.dim_feedforward, + "mimi.decoder_transformer", + vb.pp("mimi.decoder_transformer"), + )?; + + let hop_length: usize = seanet_cfg.ratios.iter().product(); + let encoder_frame_rate = config.mimi.sample_rate as f64 / hop_length as f64; + + let mimi = MimiModel::new( + encoder, + decoder, + encoder_transformer, + decoder_transformer, + config.mimi.frame_rate, + encoder_frame_rate, + config.mimi.sample_rate, + config.mimi.channels, + config.mimi.quantizer.dimension, + config.mimi.quantizer.output_dimension, + "mimi", + vb.pp("mimi"), + )?; + + Ok(Self { + flow_lm, + mimi, + conditioner, + temp: defaults::TEMPERATURE, + lsd_decode_steps: defaults::LSD_DECODE_STEPS, + eos_threshold: defaults::EOS_THRESHOLD, + sample_rate: config.mimi.sample_rate, + dim, + ldim, + device, + }) + } + + /// Load a precomputed speaker latent prompt (`.safetensors` with an `audio_prompt` tensor, + /// shape `[1, T, d_model]`) and prime a `ModelState` by running it through FlowLM. + pub fn voice_state_from_prompt_file(&self, path: &Path) -> Result { + let tensors = candle_core::safetensors::load(path, &self.device)?; + let prompt = tensors + .get("audio_prompt") + .ok_or_else(|| anyhow::anyhow!("'audio_prompt' not found in {path:?}"))?; + let prompt = if prompt.device().same_device(&self.device) { + prompt.clone() + } else { + prompt.to_device(&self.device)? + }; + let mut state = init_states(1, 1000); + self.run_flow_lm_prompt(&prompt, &mut state)?; + Ok(state) + } + + fn run_flow_lm_prompt(&self, conditioning: &Tensor, state: &mut ModelState) -> Result<()> { + let empty_text = Tensor::zeros((1, 0), DType::I64, &self.device)?; + let text_embeddings = self.conditioner.forward(&empty_text)?; + let input = Tensor::cat(&[conditioning, &text_embeddings], 1)?; + let _ = self.flow_lm.transformer.forward(&input, state, 0)?; + let increment_by = conditioning.dims()[1]; + increment_steps(state, "offset", increment_by); + Ok(()) + } + + /// Token-length-aware sentence chunking, keeping each chunk under `MAX_TOKENS_PER_CHUNK` to + /// preserve O(N) attention for long inputs. + pub fn split_into_best_sentences(&self, text: &str) -> Vec { + const MAX_TOKENS_PER_CHUNK: usize = 50; + + let prepared_text = prepare_text_prompt(text); + + let raw_sentences: Vec<&str> = prepared_text + .split_inclusive(['.', '!', '?', ';', ':']) + .map(|s| s.trim()) + .filter(|s| !s.is_empty()) + .collect(); + + if raw_sentences.is_empty() { + return vec![prepared_text]; + } + + let mut chunks = Vec::new(); + let mut current_chunk = String::new(); + let mut current_token_count = 0; + + for sentence in raw_sentences { + let sentence_tokens = self + .conditioner + .count_tokens(sentence) + .unwrap_or(MAX_TOKENS_PER_CHUNK); + + if sentence_tokens > MAX_TOKENS_PER_CHUNK { + if !current_chunk.is_empty() { + chunks.push(current_chunk); + current_chunk = String::new(); + current_token_count = 0; + } + + let words: Vec<&str> = sentence.split_whitespace().collect(); + const WORDS_PER_BATCH: usize = 35; + + for word_batch in words.chunks(WORDS_PER_BATCH) { + let chunk_str = word_batch.join(" "); + let actual_tokens = self + .conditioner + .count_tokens(&chunk_str) + .unwrap_or(MAX_TOKENS_PER_CHUNK); + + if actual_tokens <= MAX_TOKENS_PER_CHUNK { + chunks.push(chunk_str); + } else { + let mid = word_batch.len() / 2; + chunks.push(word_batch[..mid].join(" ")); + chunks.push(word_batch[mid..].join(" ")); + } + } + continue; + } + + if current_chunk.is_empty() { + current_chunk = sentence.to_string(); + current_token_count = sentence_tokens; + } else if current_token_count + sentence_tokens > MAX_TOKENS_PER_CHUNK { + chunks.push(current_chunk); + current_chunk = sentence.to_string(); + current_token_count = sentence_tokens; + } else { + current_chunk.push(' '); + current_chunk.push_str(sentence); + current_token_count += sentence_tokens; + } + } + + if !current_chunk.is_empty() { + chunks.push(current_chunk); + } + + chunks + } + + /// Generate mono audio for `text` conditioned on `voice_state` (a primed speaker prompt). + pub fn generate(&self, text: &str, voice_state: &ModelState) -> Result { + let mut audio_chunks = Vec::new(); + for chunk in self.generate_stream(text, voice_state) { + audio_chunks.push(chunk?); + } + if audio_chunks.is_empty() { + anyhow::bail!("No audio generated"); + } + let audio = Tensor::cat(&audio_chunks, 2)?; + Ok(audio.squeeze(0)?) + } + + fn generate_stream<'a>( + &'a self, + text: &str, + voice_state: &ModelState, + ) -> Box> + 'a> { + let chunks = self.split_into_best_sentences(text); + let voice_state_owned = voice_state.clone(); + let iterator = chunks.into_iter().flat_map(move |chunk_text| { + self.generate_stream_segment(chunk_text, &voice_state_owned) + }); + Box::new(iterator) + } + + fn generate_stream_segment( + &self, + text: String, + voice_state: &ModelState, + ) -> Box>> { + let mut state = voice_state.clone(); + let mut mimi_state = init_states(1, 1000); + + let prepared_text = prepare_text_prompt(&text); + + let tokens = match self.conditioner.prepare(&prepared_text, &self.device) { + Ok(t) => t, + Err(e) => return Box::new(std::iter::once(Err(e))), + }; + + let text_embeddings = match self.conditioner.forward(&tokens) { + Ok(e) => e, + Err(e) => return Box::new(std::iter::once(Err(e))), + }; + + if let Err(e) = self + .flow_lm + .transformer + .forward(&text_embeddings, &mut state, 0) + { + return Box::new(std::iter::once(Err(anyhow::Error::from(e)))); + } + + let max_gen_len = (prepared_text.split_whitespace().count() + 2) * 13; + let frames_after_eos = estimate_frames_after_eos(&text); + + let mut backbone_input = match self.flow_lm.bos_emb.clone().reshape((1, 1, self.ldim)) { + Ok(t) => t, + Err(e) => return Box::new(std::iter::once(Err(anyhow::Error::from(e)))), + }; + + let mut eos_step: Option = None; + let mut finished = false; + + let model = self.clone(); + + let time_embeddings = match model.flow_lm.flow_net.compute_time_embeddings( + model.lsd_decode_steps, + &model.device, + DType::F32, + ) { + Ok(te) => te, + Err(e) => return Box::new(std::iter::once(Err(anyhow::Error::from(e)))), + }; + + let empty_text_embeddings = + Tensor::zeros((1, 0, model.dim), DType::F32, &model.device).unwrap(); + + Box::new((0..max_gen_len).map_while(move |step| { + if finished { + return None; + } + + let (next_latent, is_eos) = match model.flow_lm.forward( + &backbone_input, + &empty_text_embeddings, + &mut state, + &time_embeddings, + model.temp, + model.eos_threshold, + step, + ) { + Ok(res) => res, + Err(e) => return Some(Err(anyhow::anyhow!(e))), + }; + + let audio_frame = match (|| -> Result { + let next_latent_denorm = next_latent + .broadcast_mul(&model.flow_lm.emb_std)? + .broadcast_add(&model.flow_lm.emb_mean)?; + + let mimi_input = next_latent_denorm.unsqueeze(1)?.transpose(1, 2)?; + let quantized = model.mimi.quantize(&mimi_input)?; + let audio = model + .mimi + .decode_from_latent(&quantized, &mut mimi_state, step) + .map_err(|e| anyhow::anyhow!(e))?; + Ok(audio) + })() { + Ok(frame) => frame, + Err(e) => return Some(Err(e)), + }; + + if is_eos && eos_step.is_none() { + eos_step = Some(step); + } + + if let Some(e_step) = eos_step { + if step >= e_step + frames_after_eos { + finished = true; + } + } + + backbone_input = next_latent.unsqueeze(1).unwrap(); + + Some(Ok(audio_frame)) + })) + } +} + +fn prepare_text_prompt(text: &str) -> String { + let text = super::pause::strip_pause_markers(text); + + let mut text = text.trim().to_string(); + if text.is_empty() { + return ".".to_string(); + } + + text = text.replace(['\n', '\r'], " ").replace(" ", " "); + + let word_count = text.split_whitespace().count(); + + if let Some(first) = text.chars().next() { + if !first.is_uppercase() { + text = format!("{}{}", first.to_uppercase(), &text[first.len_utf8()..]); + } + } + + if let Some(last) = text.chars().last() { + if last.is_alphanumeric() { + text.push('.'); + } + } + + if word_count < 5 { + text = format!("{}{}", " ".repeat(8), text); + } + + text +} + +fn estimate_frames_after_eos(text: &str) -> usize { + let word_count = text.split_whitespace().count(); + if word_count <= 4 { + 3 + 2 + } else { + 1 + 2 + } +} diff --git a/mistralrs-core/src/speech_models/pockettts/voice_state.rs b/mistralrs-core/src/speech_models/pockettts/voice_state.rs new file mode 100644 index 0000000000..dede831a7b --- /dev/null +++ b/mistralrs-core/src/speech_models/pockettts/voice_state.rs @@ -0,0 +1,168 @@ +//! Voice state management for streaming generation and voice cloning + +use candle_core::{Result, Tensor}; +use std::collections::HashMap; + +/// Model state type for stateful modules +pub type ModelState = HashMap>; + +/// Common per-attention state keys. +pub const ATTN_POS_KEY: &str = "pos"; +pub const ATTN_LEN_KEY: &str = "l"; +pub const ATTN_HEAD_KEY: &str = "head"; +pub const ATTN_K_BUF_KEY: &str = "k_buf"; +pub const ATTN_V_BUF_KEY: &str = "v_buf"; + +/// Cursor/scalar metadata for attention cache state. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct AttentionCursor { + pub pos: usize, + pub len: usize, + pub head: usize, +} + +/// Initialize empty model state for all stateful modules +/// +/// Creates a nested HashMap structure that will be populated +/// as modules run their forward passes. +pub fn init_states(_batch_size: usize, _seq_len: usize) -> ModelState { + // Start with empty state - modules will populate as needed + HashMap::new() +} + +/// Get or create a module's state entry +pub fn get_or_create_state<'a>( + state: &'a mut ModelState, + module_name: &str, +) -> &'a mut HashMap { + state.entry(module_name.to_string()).or_default() +} + +fn tensor_to_usize(t: &Tensor) -> Option { + if let Ok(v) = t.to_scalar::() { + return Some(v.max(0) as usize); + } + if let Ok(v) = t.to_scalar::() { + return Some(v as usize); + } + None +} + +/// Read attention cursor values from a module state map. +pub fn read_attention_cursor(module_state: &HashMap) -> AttentionCursor { + AttentionCursor { + pos: module_state + .get(ATTN_POS_KEY) + .and_then(tensor_to_usize) + .unwrap_or(0), + len: module_state + .get(ATTN_LEN_KEY) + .and_then(tensor_to_usize) + .unwrap_or(0), + head: module_state + .get(ATTN_HEAD_KEY) + .and_then(tensor_to_usize) + .unwrap_or(0), + } +} + +/// Write attention cursor values into a module state map. +pub fn write_attention_cursor( + module_state: &mut HashMap, + cursor: AttentionCursor, + device: &candle_core::Device, +) -> Result<()> { + module_state.insert( + ATTN_POS_KEY.to_string(), + Tensor::new(cursor.pos as u32, device)?, + ); + module_state.insert( + ATTN_LEN_KEY.to_string(), + Tensor::new(cursor.len as i64, device)?, + ); + module_state.insert( + ATTN_HEAD_KEY.to_string(), + Tensor::new(cursor.head as i64, device)?, + ); + Ok(()) +} + +/// Read attention cursor for a module name from a full model state. +pub fn get_attention_cursor(state: &ModelState, module_name: &str) -> AttentionCursor { + state + .get(module_name) + .map(read_attention_cursor) + .unwrap_or_default() +} + +/// Increment step counters in model state for all modules +/// +/// This is used after processing tokens to update position information +/// for streaming generation. +pub fn increment_steps(state: &mut ModelState, key: &str, increment: usize) { + for (_module_name, module_state) in state.iter_mut() { + if let Some(step_tensor) = module_state.get_mut(key) { + if let Ok(current) = step_tensor.to_scalar::() { + if let Ok(new_tensor) = + Tensor::new(current + increment as i64, step_tensor.device()) + { + *step_tensor = new_tensor; + } + } + } + } +} + +/// Get the current step/offset for a module +pub fn get_offset(state: &ModelState, module_name: &str) -> usize { + state + .get(module_name) + .and_then(|s| s.get("offset")) + .and_then(|t| t.to_scalar::().ok()) + .unwrap_or(0) as usize +} + +/// Set the offset for a module +pub fn set_offset(state: &mut ModelState, module_name: &str, offset: usize) -> Result<()> { + let module_state = get_or_create_state(state, module_name); + let device = module_state + .values() + .next() + .map(|t| t.device().clone()) + .unwrap_or(candle_core::Device::Cpu); + module_state.insert("offset".to_string(), Tensor::new(offset as i64, &device)?); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_init_states() { + let state = init_states(1, 100); + assert!(state.is_empty()); + } + + #[test] + fn test_get_or_create_state() { + let mut state = init_states(1, 100); + let module_state = get_or_create_state(&mut state, "test_module"); + assert!(module_state.is_empty()); + assert!(state.contains_key("test_module")); + } + + #[test] + fn test_offset_operations() -> Result<()> { + let mut state = init_states(1, 100); + + // Initially offset is 0 + assert_eq!(get_offset(&state, "test"), 0); + + // Set offset + set_offset(&mut state, "test", 42)?; + assert_eq!(get_offset(&state, "test"), 42); + + Ok(()) + } +} diff --git a/mistralrs-pyo3/mistralrs.pyi b/mistralrs-pyo3/mistralrs.pyi index b2f93b4f31..4caaa20eaa 100644 --- a/mistralrs-pyo3/mistralrs.pyi +++ b/mistralrs-pyo3/mistralrs.pyi @@ -284,6 +284,7 @@ class DiffusionArchitecture(Enum): @dataclass class SpeechLoaderType(Enum): Dia = "Dia" + PocketTts = "PocketTts" @dataclass class IsqOrganization(Enum): diff --git a/mistralrs-pyo3/src/lib.rs b/mistralrs-pyo3/src/lib.rs index a956706690..67d5c97ada 100644 --- a/mistralrs-pyo3/src/lib.rs +++ b/mistralrs-pyo3/src/lib.rs @@ -561,6 +561,7 @@ fn parse_which( dac_model_id, arch: arch.into(), cfg: None, + voice: None, }), }) } diff --git a/mistralrs-pyo3/src/which.rs b/mistralrs-pyo3/src/which.rs index 6bfba8a664..5de1ca89b2 100644 --- a/mistralrs-pyo3/src/which.rs +++ b/mistralrs-pyo3/src/which.rs @@ -161,12 +161,14 @@ impl From for DiffusionLoaderType { #[derive(Debug, Clone, PartialEq)] pub enum SpeechLoaderType { Dia, + PocketTts, } impl From for mistralrs_core::SpeechLoaderType { fn from(value: SpeechLoaderType) -> Self { match value { SpeechLoaderType::Dia => mistralrs_core::SpeechLoaderType::Dia, + SpeechLoaderType::PocketTts => mistralrs_core::SpeechLoaderType::PocketTts, } } } diff --git a/mistralrs-server-core/src/speech_generation.rs b/mistralrs-server-core/src/speech_generation.rs index 26627d74a8..39cb784dfe 100644 --- a/mistralrs-server-core/src/speech_generation.rs +++ b/mistralrs-server-core/src/speech_generation.rs @@ -17,7 +17,7 @@ use tokio::sync::mpsc::{Receiver, Sender}; use crate::{ handler_core::{ - base_process_non_streaming_response, create_response_channel, send_request, + base_process_non_streaming_response, create_response_channel, send_request_with_model, ErrorToResponse, JsonError, }, openai::{AudioResponseFormat, SpeechGenerationRequest}, @@ -118,6 +118,10 @@ pub async fn speech_generation( ) -> SpeechGenerationResponder { let (tx, mut rx) = create_response_channel(None); + // Route to the model named in the request (honoring `"default"` -> default model), so + // config.toml can serve multiple speech models selected per request by `model`. + let requested_model = oairequest.model.clone(); + let (request, response_format) = match parse_request(oairequest, state.clone(), tx) { Ok(x) => x, Err(e) => return handle_error(state, e.into()), @@ -133,7 +137,13 @@ pub async fn speech_generation( ))); } - if let Err(e) = send_request(&state, request).await { + let model_id = if requested_model == "default" { + None + } else { + Some(requested_model.as_str()) + }; + + if let Err(e) = send_request_with_model(&state, request, model_id).await { return handle_error(state, e.into()); } diff --git a/mistralrs/Cargo.toml b/mistralrs/Cargo.toml index 57c723e04f..51c13a75f7 100644 --- a/mistralrs/Cargo.toml +++ b/mistralrs/Cargo.toml @@ -91,6 +91,10 @@ path = "examples/models/diffusion/main.rs" name = "speech" path = "examples/models/speech/main.rs" +[[example]] +name = "speech_pockettts" +path = "examples/models/speech_pockettts/main.rs" + [[example]] name = "asr" path = "examples/models/asr/main.rs" diff --git a/mistralrs/examples/models/speech_pockettts/main.rs b/mistralrs/examples/models/speech_pockettts/main.rs new file mode 100644 index 0000000000..7dcf861b88 --- /dev/null +++ b/mistralrs/examples/models/speech_pockettts/main.rs @@ -0,0 +1,38 @@ +//! CPU-fast text-to-speech with pocket-tts (Kyutai Mimi codec + FlowLM). +//! +//! Run with: `cargo run --release --example speech_pockettts -p mistralrs` + +use std::time::Instant; + +use anyhow::Result; +use mistralrs::{speech_utils, SpeechLoaderType, SpeechModelBuilder}; + +#[tokio::main] +async fn main() -> Result<()> { + let model = SpeechModelBuilder::new( + "kyutai/pocket-tts-without-voice-cloning", + SpeechLoaderType::PocketTts, + ) + .with_logging() + .build() + .await?; + + let start = Instant::now(); + + let text_to_speak = + "Pocket TTS runs on the CPU in seconds, so mistral rs can serve speech without a GPU."; + + let (pcm, rate, channels) = model.generate_speech(text_to_speak).await?; + + let finished = Instant::now(); + + let mut output = std::fs::File::create("out.wav").unwrap(); + speech_utils::write_pcm_as_wav(&mut output, &pcm, rate as u32, channels as u16).unwrap(); + + println!( + "Done! Took {} s. Audio saved at `out.wav`.", + finished.duration_since(start).as_secs_f32(), + ); + + Ok(()) +} diff --git a/mistralrs/src/model_builder_trait.rs b/mistralrs/src/model_builder_trait.rs index edfc928ec2..3d75b81f38 100644 --- a/mistralrs/src/model_builder_trait.rs +++ b/mistralrs/src/model_builder_trait.rs @@ -980,6 +980,7 @@ pub async fn build_speech_pipeline( dac_model_id: builder.dac_model_id.clone(), arch: builder.loader_type, cfg: builder.cfg, + voice: builder.voice.clone(), }; let device = resolve_device(builder.force_cpu, None)?; @@ -1004,6 +1005,7 @@ pub async fn build_speech_pipeline( model_id: builder.model_id.clone(), dac_model_id: builder.dac_model_id.clone(), arch: builder.loader_type, + voice: builder.voice.clone(), dtype: builder.dtype, }, token_source: builder.token_source.clone(), diff --git a/mistralrs/src/speech_model.rs b/mistralrs/src/speech_model.rs index 24cffd0214..41458e57c0 100644 --- a/mistralrs/src/speech_model.rs +++ b/mistralrs/src/speech_model.rs @@ -11,6 +11,7 @@ pub struct SpeechModelBuilder { pub(crate) token_source: TokenSource, pub(crate) hf_revision: Option, pub(crate) cfg: Option, + pub(crate) voice: Option, // Model running pub(crate) loader_type: SpeechLoaderType, @@ -38,9 +39,16 @@ impl SpeechModelBuilder { with_logging: false, cfg: None, dac_model_id: None, + voice: None, } } + /// Speaker voice for pocket-tts (a stock name like `alba`). Ignored by Dia. Defaults to `alba`. + pub fn with_voice(mut self, voice: impl ToString) -> Self { + self.voice = Some(voice.to_string()); + self + } + /// DAC Model ID to load from. If not provided, this is automatically downloaded from the default path for the model. /// This may be a HF hub repo or a local path. pub fn with_dac_model_id(mut self, dac_model_id: String) -> Self {