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]>
607 lines
21 KiB
Rust
607 lines
21 KiB
Rust
//! 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(())
|
||
}
|