Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -17,3 +17,7 @@ docs/node_modules/
docs/dist/
docs/.astro/
out/

dist
dist-build-arm64
.memsearch
55 changes: 55 additions & 0 deletions docs/src/content/docs/examples/rust/models/speech-pockettts.md
Original file line number Diff line number Diff line change
@@ -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"
---

<!-- generated by docs/scripts/render_examples.py; edit the source example instead -->

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)
100 changes: 99 additions & 1 deletion docs/src/content/docs/guides/models/use-speech-models.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ description: Voxtral Realtime for speech transcription, Dia 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 Realtime**: multimodal model accepting audio input for speech transcription through
`/v1/chat/completions`.
Expand Down Expand Up @@ -181,3 +181,101 @@ speech_utils::write_pcm_as_wav(&mut output, &pcm, rate as u32, channels as u16)?

</TabItem>
</Tabs>

## 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`).

<Tabs>
<TabItem label="HTTP">

```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.

</TabItem>
<TabItem label="Python">

```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("<h", sample) for sample in pcm_ints))
```

</TabItem>
<TabItem label="Rust">

```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)?;
```

</TabItem>
</Tabs>

### 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
```
1 change: 1 addition & 0 deletions docs/src/content/docs/reference/python/enums.md
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,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`
Expand Down
1 change: 1 addition & 0 deletions docs/src/content/docs/reference/supported-models.md
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,7 @@ The `Architecture` column is the `config.json` `architectures` value. Per-family
| Architecture | Model families | Example |
|---|---|---|
| `Dia` | Dia | <details><summary><code>nari-labs/Dia-1.6B</code></summary><code>mistralrs run -m nari-labs/Dia-1.6B</code></details> |
| `PocketTts` | PocketTts | <details><summary><code>kyutai/pocket-tts-without-voice-cloning</code></summary><code>mistralrs run -m kyutai/pocket-tts-without-voice-cloning</code></details> |

## Embedding

Expand Down
14 changes: 13 additions & 1 deletion mistralrs-cli/src/args/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ pub use server::*;
use clap::{Parser, Subcommand, ValueEnum};
use clap_complete::Shell;
use mistralrs_core::{
ReasoningEffort, TokenSource, DEFAULT_MAX_DECODE_STEPS_BEFORE_PREFILL,
ReasoningEffort, SpeechLoaderType, TokenSource, DEFAULT_MAX_DECODE_STEPS_BEFORE_PREFILL,
DEFAULT_MAX_NUM_BATCHED_TOKENS, DEFAULT_MAX_PREFILL_CHUNK_TOKENS,
};
use serde::Deserialize;
Expand Down Expand Up @@ -439,6 +439,10 @@ fn parse_arch(s: &str) -> Result<mistralrs_core::NormalLoaderType, String> {
s.parse()
}

fn parse_speech_arch(s: &str) -> Result<SpeechLoaderType, String> {
s.parse()
}

fn parse_dtype(s: &str) -> Result<mistralrs_core::ModelDType, String> {
s.parse()
}
Expand Down Expand Up @@ -531,6 +535,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<SpeechLoaderType>,

/// Speaker voice for pocket-tts (a stock name like `alba`). Ignored by Dia.
#[arg(long)]
voice: Option<String>,
},

/// Embedding model
Expand Down
12 changes: 9 additions & 3 deletions mistralrs-cli/src/commands/serve.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@ use tracing::{debug, info, warn};

use mistralrs_core::{
initialize_logging, DiffusionLoaderType, McpClientConfig, ModelSelected, PagedCacheType,
SpeechLoaderType,
};
use mistralrs_server_core::{
approvals::ApprovalBroker,
Expand All @@ -28,6 +27,7 @@ use crate::args::{
MultimodalOptions, QuantizationOptions, RuntimeOptions, SandboxMode, SandboxOptions,
ServerOptions,
};
use crate::config::detect_speech_arch;
use crate::ui::build_ui_router;

const MEBIBYTE_BYTES: usize = 1024 * 1024;
Expand Down Expand Up @@ -493,10 +493,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,
}),

Expand Down
27 changes: 25 additions & 2 deletions mistralrs-cli/src/config/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ use crate::args::{
ModelType, MultimodalAdapterOptions, MultimodalOptions, PagedAttentionOptions,
QuantizationOptions, RuntimeOptions, SandboxOptions, ServerOptions,
};
use mistralrs_core::{ModelDType, NormalLoaderType, ReasoningEffort, TokenSource};
use mistralrs_core::{ModelDType, NormalLoaderType, ReasoningEffort, SpeechLoaderType, TokenSource};

#[derive(Deserialize)]
#[serde(tag = "command", rename_all = "kebab-case")]
Expand Down Expand Up @@ -90,6 +90,13 @@ pub struct ModelEntry {
pub tokenizer: Option<PathBuf>,
#[serde(default)]
pub arch: Option<NormalLoaderType>,
/// Speech architecture (`dia` or `pockettts`). Only meaningful for `kind = "speech"`.
/// Auto-detected from `model_id` when omitted.
#[serde(default)]
pub speech_arch: Option<SpeechLoaderType>,
/// Speaker voice for pocket-tts (a stock name like `alba`). Only meaningful for `kind = "speech"`.
#[serde(default)]
pub voice: Option<String>,
#[serde(default)]
pub dtype: ModelDType,
#[serde(default)]
Expand Down Expand Up @@ -153,6 +160,14 @@ pub fn load_cli_config(path: &Path) -> Result<CliConfig> {
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()),
Expand Down Expand Up @@ -286,7 +301,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(),
Expand Down
2 changes: 2 additions & 0 deletions mistralrs-core/src/model_loader.rs
Original file line number Diff line number Diff line change
Expand Up @@ -517,12 +517,14 @@ fn loader_from_model_selected(args: LoaderBuilder) -> anyhow::Result<Box<dyn Loa
model_id,
dac_model_id,
arch,
voice,
..
} => Box::new(SpeechLoader {
model_id,
dac_model_id,
arch,
cfg: None,
voice,
}),
ModelSelected::XLora {
model_id,
Expand Down
5 changes: 5 additions & 0 deletions mistralrs-core/src/model_metadata.rs
Original file line number Diff line number Diff line change
Expand Up @@ -373,6 +373,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")],
},
}
}
}
Expand Down
4 changes: 4 additions & 0 deletions mistralrs-core/src/model_selected.rs
Original file line number Diff line number Diff line change
Expand Up @@ -796,6 +796,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<String>,

/// Model data type. Defaults to `auto`.
#[arg(long, default_value_t = ModelDType::Auto, value_parser = parse_model_dtype)]
dtype: ModelDType,
Expand Down
8 changes: 8 additions & 0 deletions mistralrs-core/src/pipeline/auto.rs
Original file line number Diff line number Diff line change
Expand Up @@ -409,6 +409,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)?;
Expand Down Expand Up @@ -518,6 +525,7 @@ impl AutoLoader {
dac_model_id: None,
arch: tp,
cfg: None,
voice: None,
});
*guard = Some(loader);
}
Expand Down
Loading
Loading