Files
rustytorch/crates/models/rtx-csm/examples/tts_server.rs
T
osobhandClaude Opus 4.7 a5cedfb46a rtx-csm: emotional_speech_guide — CREMA-D vs RAVDESS firdhokk verdict
8-gen bench (4 emotions × 2 corpora) at seed=42 against firdhokk
Whisper-LV3:

  target    RAVDESS              CREMA-D
  happy     happy (0.999) ✓      happy (0.999) ✓
  angry     neutral (0.92)       sad (0.99)
  fearful   happy (0.998)        fearful (0.984) ✓
  sad       angry (0.99)         fearful (0.99)

CREMA-D 2/4 vs RAVDESS 1/4. Larger / more naturalistic corpus
produces more class-pure fearful direction. Neither corpus solves
angry or sad — recipe shifts into 'vague expressivity' rather than
class-specific corners.

Practical: prefer CREMA-D when available; A/B both per emotion if
class precision matters.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
2026-04-30 00:01:02 -07:00

607 lines
21 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Standalone HTTP server for the rtx-csm pipeline.
//!
//! Loads CSM-1B + (optionally) AudioSeal + (optionally) WavLM-SV at startup
//! and exposes:
//!
//! GET /health → "ok"
//! POST /v1/tts → audio/wav (24 kHz mono)
//! POST /v1/detect [multipart audio] → JSON { mean_presence, message }
//! POST /v1/speaker_embed [multipart audio] → JSON { embedding: [512 floats] }
//! POST /v1/speaker_compare [multipart a + b] → JSON { cosine }
//!
//! Inference is sync + GPU-heavy → spawn_blocking for every request. The
//! Generator + speaker scorer are wrapped in `std::sync::Mutex` so only
//! one inference runs per process; AudioSeal embed/detect are stateless
//! and run concurrently. Run multiple processes behind a load balancer
//! for higher throughput.
//!
//! Usage (full pipeline):
//! ```bash
//! cargo run -p rtx-csm --release --features metal --example tts_server -- \
//! --bind 127.0.0.1:8080 \
//! --audioseal-generator /tmp/audioseal_generator.safetensors \
//! --audioseal-detector /tmp/audioseal_detector.safetensors \
//! --audioseal-message 0xCAFE \
//! --wavlm-sv /tmp/wavlm_sv.safetensors
//!
//! curl -X POST http://127.0.0.1:8080/v1/tts \
//! -H 'content-type: application/json' \
//! -d '{"text":"Hello.","speaker":0}' --output /tmp/out.wav
//! curl -X POST http://127.0.0.1:8080/v1/detect \
//! -F audio=@/tmp/out.wav -F source_rate=24000
//! curl -X POST http://127.0.0.1:8080/v1/speaker_embed \
//! -F audio=@/tmp/out.wav -F source_rate=24000
//! curl -X POST http://127.0.0.1:8080/v1/speaker_compare \
//! -F a=@/tmp/a.wav -F b=@/tmp/b.wav -F source_rate=24000
//! ```
use anyhow::Result;
use axum::{
Json, Router,
body::{Body, Bytes},
extract::{Multipart, State},
http::{StatusCode, header},
response::IntoResponse,
routing::{get, post},
};
use candle_core::DType;
use clap::Parser;
use futures_util::stream::Stream;
use rtx_csm::{
GenerateOptions, Generator, PostProcess, audio_io, audioseal::AudioSealWatermarker,
speaker_sim::WavLmSimilarity, watermark::ResampledWatermarker,
};
use serde::{Deserialize, Serialize};
use std::io::Cursor;
use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::{Arc, Mutex};
const AUDIOSEAL_RATE: u32 = 16_000;
#[derive(Debug, Parser)]
struct Cli {
/// Bind address (e.g. 127.0.0.1:8080 or 0.0.0.0:8080).
#[arg(long, default_value = "127.0.0.1:8080")]
bind: SocketAddr,
/// Optional path to a quantized GGUF (output of `examples/quantize`).
#[arg(long)]
quantized_gguf: Option<PathBuf>,
/// Force CPU even if metal/cuda features are enabled.
#[arg(long)]
cpu: bool,
/// AudioSeal generator safetensors. When BOTH this and
/// --audioseal-detector are set, the inline watermarker is installed
/// on the Generator (every TTS response is watermarked) and
/// /v1/detect is enabled.
#[arg(long)]
audioseal_generator: Option<PathBuf>,
#[arg(long)]
audioseal_detector: Option<PathBuf>,
/// 16-bit watermark message (decimal or 0xHEX).
#[arg(long, default_value = "0")]
audioseal_message: String,
/// WavLM-SV safetensors (output of `wavlm_sv_convert`). When set,
/// /v1/speaker_embed and /v1/speaker_compare are enabled.
#[arg(long)]
wavlm_sv: Option<PathBuf>,
/// LoRA voice adapter (safetensors, output of `lora_train` or
/// `lora_train_emotional`). Loaded into the backbone (FP or quantized
/// — combines with --quantized-gguf per Phase 12.6) before serving so
/// every TTS request uses the trained voice.
#[arg(long)]
lora: Option<PathBuf>,
/// LoRA rank override. When omitted, auto-detected from the adapter's
/// embedded metadata (Phase 12.5); falls back to 8 for older adapters.
#[arg(long)]
lora_rank: Option<usize>,
/// LoRA alpha override. Auto-detected from metadata when omitted.
#[arg(long)]
lora_alpha: Option<f32>,
/// Force extended LoRA coverage (q+k+v+output_proj + MLP). Auto-
/// detected from metadata when omitted; setting it on a classic q+v
/// adapter just allocates extra unused B=0 slots.
#[arg(long, default_value_t = false)]
extended_lora: bool,
}
fn parse_message(s: &str) -> Result<u16> {
let s = s.trim();
if let Some(rest) = s.strip_prefix("0x").or_else(|| s.strip_prefix("0X")) {
Ok(u16::from_str_radix(rest, 16)?)
} else {
Ok(s.parse::<u16>()?)
}
}
struct AppState {
generator: Mutex<Generator>,
/// Held by the generator's set_watermarker; we keep a separate
/// detector instance for /v1/detect (mmap'd weights are shared).
audioseal_detector: Option<AudioSealWatermarker>,
wavlm: Option<Mutex<WavLmSimilarity>>,
}
#[derive(Debug, Deserialize)]
struct TtsRequest {
text: String,
#[serde(default)]
speaker: u32,
#[serde(default = "default_max_audio_ms")]
max_audio_ms: u32,
#[serde(default = "default_temperature")]
temperature: f64,
#[serde(default = "default_top_k")]
top_k: usize,
#[serde(default = "default_top_p")]
top_p: f64,
#[serde(default = "default_seed")]
seed: u64,
/// Optional per-request 16-bit watermark message override. If set
/// AND the server has AudioSeal loaded, the watermark on this
/// response carries this message instead of the server-startup
/// default. Useful for tagging each generation with a unique ID
/// (e.g. clawsample's job_id mod 0x10000) for audit trails.
/// Accepts decimal or 0xHEX as a string for safety with JSON.
#[serde(default)]
watermark_message: Option<String>,
}
fn default_max_audio_ms() -> u32 {
10_000
}
fn default_temperature() -> f64 {
0.9
}
fn default_top_k() -> usize {
50
}
fn default_top_p() -> f64 {
0.9
}
fn default_seed() -> u64 {
42
}
#[derive(Debug, Serialize)]
struct DetectResponse {
mean_presence: f32,
message: u16,
message_hex: String,
}
#[derive(Debug, Serialize)]
struct EmbedResponse {
embedding: Vec<f32>,
}
#[derive(Debug, Serialize)]
struct CompareResponse {
cosine: f32,
}
async fn health() -> &'static str {
"ok"
}
async fn tts(
State(state): State<Arc<AppState>>,
Json(req): Json<TtsRequest>,
) -> Result<impl IntoResponse, (StatusCode, String)> {
if req.text.trim().is_empty() {
return Err((StatusCode::BAD_REQUEST, "text is required".to_string()));
}
let opts = GenerateOptions {
max_audio_ms: req.max_audio_ms,
temperature: req.temperature,
top_k: req.top_k,
top_p: req.top_p,
seed: req.seed,
..GenerateOptions::default()
};
let state2 = state.clone();
let text = req.text.clone();
let speaker = req.speaker;
// Per-request message override; None falls back to startup default.
let req_msg = match req.watermark_message.as_deref() {
Some(s) => Some(
parse_message(s)
.map_err(|e| (StatusCode::BAD_REQUEST, format!("watermark_message: {e}")))?,
),
None => None,
};
let result: Result<Vec<f32>, String> = tokio::task::spawn_blocking(move || {
let mut g = state2.generator.lock().expect("generator mutex poisoned");
let mut pcm = g
.generate(&text, speaker, &[], opts)
.map_err(|e| format!("inference: {e}"))?;
PostProcess::default()
.apply(&mut pcm, g.config.sample_rate)
.map_err(|e| format!("post: {e}"))?;
if let Some(wm) = g.watermarker.as_ref() {
pcm = match req_msg {
Some(m) => wm.embed_with_message(&pcm, m),
None => wm.embed(&pcm),
}
.map_err(|e| format!("watermark: {e}"))?;
}
Ok(pcm)
})
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("join: {e}")))?;
let pcm = result.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e))?;
let mut buf: Vec<u8> = Vec::new();
{
let cursor = Cursor::new(&mut buf);
let spec = hound::WavSpec {
channels: 1,
sample_rate: 24_000,
bits_per_sample: 16,
sample_format: hound::SampleFormat::Int,
};
let mut w = hound::WavWriter::new(cursor, spec).map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("wav writer: {e}"),
)
})?;
for &s in pcm.iter() {
let v = (s.clamp(-1.0, 1.0) * i16::MAX as f32) as i16;
w.write_sample(v)
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("wav write: {e}")))?;
}
w.finalize().map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("wav finalize: {e}"),
)
})?;
}
let _ = audio_io::TARGET_SAMPLE_RATE;
Ok((
StatusCode::OK,
[(header::CONTENT_TYPE, "audio/wav")],
Bytes::from(buf),
))
}
/// Streaming TTS: chunks of raw 16-bit little-endian PCM (24 kHz mono)
/// are emitted as soon as Mimi produces them. Content-Type is
/// `audio/L16; rate=24000; channels=1` per RFC 2586. First chunk arrives
/// in roughly `chunk_frames × 80ms + per-frame compute × chunk_frames`
/// — typically ~600 ms with `chunk_frames=4` on Metal.
///
/// Clients: `curl --output - http://.../v1/tts_stream | aplay -f S16_LE -r 24000`
/// or pipe directly into a player. Note: this endpoint does NOT apply
/// post-processing or the inline watermarker; streaming watermarking
/// requires a chunked AudioSeal port (deferred). Use /v1/tts for the
/// post-processed + watermarked path.
async fn tts_stream(
State(state): State<Arc<AppState>>,
Json(req): Json<TtsRequest>,
) -> Result<impl IntoResponse, (StatusCode, String)> {
if req.text.trim().is_empty() {
return Err((StatusCode::BAD_REQUEST, "text is required".to_string()));
}
let opts = GenerateOptions {
max_audio_ms: req.max_audio_ms,
temperature: req.temperature,
top_k: req.top_k,
top_p: req.top_p,
seed: req.seed,
..GenerateOptions::default()
};
// mpsc channel: blocking inference task pushes chunks, async response
// consumer yields them as the body stream.
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Bytes, std::io::Error>>(8);
let state2 = state.clone();
let text = req.text.clone();
let speaker = req.speaker;
tokio::task::spawn_blocking(move || {
let mut g = state2.generator.lock().expect("generator mutex poisoned");
// chunk_frames=4 → ~320ms of audio per emitted chunk; first chunk
// arrives after ~4 frames × per-frame-compute. Tunable.
let chunk_frames = 4usize;
let res = g.generate_streaming(
&text,
speaker,
&[],
opts,
chunk_frames,
|samples: &[f32]| {
// Encode chunk as 16-bit LE PCM.
let mut buf = Vec::with_capacity(samples.len() * 2);
for &s in samples {
let v = (s.clamp(-1.0, 1.0) * i16::MAX as f32) as i16;
buf.extend_from_slice(&v.to_le_bytes());
}
// blocking_send is correct here — we're inside spawn_blocking.
let _ = tx.blocking_send(Ok(Bytes::from(buf)));
Ok(())
},
);
if let Err(e) = res {
let _ = tx.blocking_send(Err(std::io::Error::new(
std::io::ErrorKind::Other,
format!("stream: {e}"),
)));
}
// dropping tx closes the channel
});
let stream = ReceiverStream { rx };
let body = Body::from_stream(stream);
Ok((
StatusCode::OK,
[(header::CONTENT_TYPE, "audio/L16; rate=24000; channels=1")],
body,
))
}
/// Adapter: tokio mpsc Receiver<T> → futures Stream<Item = T>.
struct ReceiverStream {
rx: tokio::sync::mpsc::Receiver<Result<Bytes, std::io::Error>>,
}
impl Stream for ReceiverStream {
type Item = Result<Bytes, std::io::Error>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
self.rx.poll_recv(cx)
}
}
/// Read a multipart `audio` field as raw WAV bytes, decode at the given
/// `source_rate` (default 24 kHz, CSM-1B native), then resample to 16 kHz
/// for AudioSeal/WavLM-SV consumption.
async fn read_audio_at(
multipart: &mut Multipart,
field_name: &str,
source_rate: u32,
target_rate: u32,
) -> Result<Vec<f32>, (StatusCode, String)> {
while let Some(field) = multipart
.next_field()
.await
.map_err(|e| (StatusCode::BAD_REQUEST, format!("multipart: {e}")))?
{
let name = field.name().unwrap_or_default().to_string();
if name != field_name {
continue;
}
let bytes = field
.bytes()
.await
.map_err(|e| (StatusCode::BAD_REQUEST, format!("read field: {e}")))?;
let tmp = std::env::temp_dir().join(format!(
"tts_server_in_{}_{:x}.wav",
field_name,
std::process::id() as u64 ^ rand::random::<u64>()
));
std::fs::write(&tmp, &bytes).map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("temp write: {e}"),
)
})?;
let raw = audio_io::load_mono_at_rate(&tmp, source_rate)
.map_err(|e| (StatusCode::BAD_REQUEST, format!("decode: {e}")))?;
let _ = std::fs::remove_file(&tmp);
let resampled = audio_io::resample(&raw, source_rate, target_rate)
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("resample: {e}")))?;
return Ok(resampled);
}
Err((
StatusCode::BAD_REQUEST,
format!("missing field: {field_name}"),
))
}
async fn detect(
State(state): State<Arc<AppState>>,
mut multipart: Multipart,
) -> Result<Json<DetectResponse>, (StatusCode, String)> {
let det = state.audioseal_detector.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"AudioSeal not loaded — pass --audioseal-generator + --audioseal-detector at startup"
.to_string(),
))?;
// Multipart fields are consumed once; read source_rate first if present,
// then pull audio. Most clients send audio first; fall back to default 24k.
let source_rate = 24_000u32;
let samples = read_audio_at(&mut multipart, "audio", source_rate, AUDIOSEAL_RATE).await?;
let result = det
.detect(&samples)
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("detect: {e}")))?;
let message = result.message.unwrap_or(0);
Ok(Json(DetectResponse {
mean_presence: result.mean_presence,
message,
message_hex: format!("0x{:04X}", message),
}))
}
async fn speaker_embed(
State(state): State<Arc<AppState>>,
mut multipart: Multipart,
) -> Result<Json<EmbedResponse>, (StatusCode, String)> {
let scorer_mu = state.wavlm.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"WavLM-SV not loaded — pass --wavlm-sv at startup".to_string(),
))?;
let source_rate = 24_000u32;
let samples = read_audio_at(&mut multipart, "audio", source_rate, AUDIOSEAL_RATE).await?;
let scorer = scorer_mu.lock().expect("wavlm mutex poisoned");
let embedding = scorer
.embed(&samples)
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("embed: {e}")))?;
Ok(Json(EmbedResponse { embedding }))
}
async fn speaker_compare(
State(state): State<Arc<AppState>>,
mut multipart: Multipart,
) -> Result<Json<CompareResponse>, (StatusCode, String)> {
let scorer_mu = state.wavlm.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"WavLM-SV not loaded — pass --wavlm-sv at startup".to_string(),
))?;
let source_rate = 24_000u32;
// Read both `a` and `b` audio fields. We can't seek over the multipart
// stream, so collect any audio fields we see in order.
let mut a: Option<Vec<f32>> = None;
let mut b: Option<Vec<f32>> = None;
while let Some(field) = multipart
.next_field()
.await
.map_err(|e| (StatusCode::BAD_REQUEST, format!("multipart: {e}")))?
{
let name = field.name().unwrap_or_default().to_string();
match name.as_str() {
"a" | "b" => {
let bytes = field
.bytes()
.await
.map_err(|e| (StatusCode::BAD_REQUEST, format!("read field: {e}")))?;
let tmp = std::env::temp_dir().join(format!(
"tts_server_cmp_{name}_{:x}.wav",
std::process::id() as u64 ^ rand::random::<u64>()
));
std::fs::write(&tmp, &bytes).map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("temp write: {e}"),
)
})?;
let raw = audio_io::load_mono_at_rate(&tmp, source_rate)
.map_err(|e| (StatusCode::BAD_REQUEST, format!("decode {name}: {e}")))?;
let _ = std::fs::remove_file(&tmp);
let resampled = audio_io::resample(&raw, source_rate, AUDIOSEAL_RATE)
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("resample: {e}")))?;
if name == "a" {
a = Some(resampled);
} else {
b = Some(resampled);
}
}
_ => {}
}
}
let a = a.ok_or((StatusCode::BAD_REQUEST, "missing 'a' audio".to_string()))?;
let b = b.ok_or((StatusCode::BAD_REQUEST, "missing 'b' audio".to_string()))?;
let scorer = scorer_mu.lock().expect("wavlm mutex poisoned");
use rtx_csm::speaker_sim::SpeakerSimilarity;
let cosine = scorer
.score(&a, &b)
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("score: {e}")))?;
Ok(Json(CompareResponse { cosine }))
}
#[tokio::main]
async fn main() -> Result<()> {
tracing_subscriber::fmt().init();
let cli = Cli::parse();
let device = if cli.cpu {
candle_core::Device::Cpu
} else {
Generator::default_device()?
};
tracing::info!("device: {device:?}");
let mut generator = if let Some(gguf) = cli.quantized_gguf.as_ref() {
tracing::info!("loading quantized CSM from {}", gguf.display());
Generator::load_csm_1b_quantized(gguf, &device, false)?
} else {
Generator::load_csm_1b(&device)?
};
tracing::info!("CSM-1B loaded");
// Optional LoRA adapter (Phase 12.4 + 12.5 self-describing metadata).
if let Some(lora_path) = cli.lora.as_ref() {
let extended_override = if cli.extended_lora { Some(true) } else { None };
rtx_csm::training::apply_lora_adapter(
&mut generator,
lora_path,
cli.lora_rank,
cli.lora_alpha,
extended_override,
&device,
)?;
}
// Optional AudioSeal — install inline + keep a detector instance.
let audioseal_message = parse_message(&cli.audioseal_message)?;
let audioseal_detector = if let (Some(g), Some(d)) = (
cli.audioseal_generator.as_ref(),
cli.audioseal_detector.as_ref(),
) {
let g_vb =
unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[g], DType::F32, &device) }?;
let d_vb =
unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[d], DType::F32, &device) }?;
let inner =
AudioSealWatermarker::from_var_builders(g_vb, d_vb, device.clone(), audioseal_message)?;
let wm = ResampledWatermarker::new(inner, generator.config.sample_rate, AUDIOSEAL_RATE);
generator.set_watermarker(Box::new(wm));
// Build a second detector instance for /v1/detect (mmap'd weights shared).
let g_vb2 =
unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[g], DType::F32, &device) }?;
let d_vb2 =
unsafe { candle_nn::VarBuilder::from_mmaped_safetensors(&[d], DType::F32, &device) }?;
let det = AudioSealWatermarker::from_var_builders(
g_vb2,
d_vb2,
device.clone(),
audioseal_message,
)?;
tracing::info!(
"AudioSeal installed (message=0x{:04X}); /v1/tts auto-watermarks, /v1/detect enabled",
audioseal_message
);
Some(det)
} else {
None
};
let wavlm = if let Some(p) = cli.wavlm_sv.as_ref() {
let scorer = WavLmSimilarity::load(p, &device)?;
tracing::info!("WavLM-SV loaded; /v1/speaker_embed and /v1/speaker_compare enabled");
Some(Mutex::new(scorer))
} else {
None
};
let _ = audioseal_message; // captured into the watermarker; field not needed on AppState.
let state = Arc::new(AppState {
generator: Mutex::new(generator),
audioseal_detector,
wavlm,
});
let app = Router::new()
.route("/health", get(health))
.route("/v1/tts", post(tts))
.route("/v1/tts_stream", post(tts_stream))
.route("/v1/detect", post(detect))
.route("/v1/speaker_embed", post(speaker_embed))
.route("/v1/speaker_compare", post(speaker_compare))
.with_state(state);
let listener = tokio::net::TcpListener::bind(&cli.bind).await?;
tracing::info!("listening on http://{}", cli.bind);
axum::serve(listener, app).await?;
Ok(())
}