A single \`generate\` invocation now produces a watermarked WAV when AudioSeal weights are passed via CLI. End-to-end verified on real CSM speech: mean_presence=1.0000, 16/16 message bits decoded. - Generator gains \`watermarker: Option<Box<dyn Watermarker>>\` slot; \`generate_to_wav\` runs \`wm.embed(&pcm)\` after post-process, before WAV write. Field is Send+Sync so the existing Arc<Mutex<Generator>> tts_server pattern still works. - watermark.rs ships ResampledWatermarker<W> adapter for handling rate mismatches (CSM 24 kHz ↔ AudioSeal 16 kHz). Output length is normalized to input length so it's a transparent drop-in. - examples/generate.rs gains --watermark-generator/--watermark-detector/ --watermark-message flags. Loads AudioSeal, wraps in resampler, installs. Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
237 lines
7.5 KiB
Rust
237 lines
7.5 KiB
Rust
//! CSM generation CLI.
|
|
//!
|
|
//! Downloads CSM-1B + Mimi + Llama tokenizer from HuggingFace (cached), then
|
|
//! generates speech for the given `--text`.
|
|
//!
|
|
//! Usage:
|
|
//! ```
|
|
//! cargo run -p rtx-csm --release --example generate -- \
|
|
//! --text "Hello from Rust." --speaker 0 --out /tmp/hello.wav
|
|
//! ```
|
|
//!
|
|
//! Context conditioning (optional, from a prior utterance you have WAV for):
|
|
//! ```
|
|
//! cargo run -p rtx-csm --release --example generate -- \
|
|
//! --text "Nice to talk to you." \
|
|
//! --context-wav prior.wav --context-text "Previous thing I said." \
|
|
//! --context-speaker 1 \
|
|
//! --speaker 0 --out /tmp/reply.wav
|
|
//! ```
|
|
|
|
use anyhow::Result;
|
|
use clap::Parser;
|
|
use rtx_csm::{audio_io, Generator, GenerateOptions, PostProcess, Segment};
|
|
|
|
#[derive(Debug, Parser)]
|
|
#[command(name = "csm-generate", about = "Generate speech with rtx-csm")]
|
|
struct Cli {
|
|
/// Text to synthesize.
|
|
#[arg(long)]
|
|
text: String,
|
|
|
|
/// Speaker id (0 or 1).
|
|
#[arg(long, default_value_t = 0)]
|
|
speaker: u32,
|
|
|
|
/// Output WAV file (24 kHz mono 16-bit).
|
|
#[arg(long)]
|
|
out: std::path::PathBuf,
|
|
|
|
/// Max audio length in milliseconds.
|
|
#[arg(long, default_value_t = 10_000)]
|
|
max_audio_ms: u32,
|
|
|
|
/// Sampling temperature.
|
|
#[arg(long, default_value_t = 0.9)]
|
|
temperature: f64,
|
|
|
|
/// Top-K sampling.
|
|
#[arg(long, default_value_t = 50)]
|
|
top_k: usize,
|
|
|
|
/// Top-p (nucleus) cutoff applied after top-k. 1.0 disables.
|
|
#[arg(long, default_value_t = 0.9)]
|
|
top_p: f64,
|
|
|
|
/// RNG seed.
|
|
#[arg(long, default_value_t = 42)]
|
|
seed: u64,
|
|
|
|
/// Optional prior utterance audio (WAV) for context conditioning.
|
|
#[arg(long)]
|
|
context_wav: Option<std::path::PathBuf>,
|
|
|
|
/// Transcription of the context WAV.
|
|
#[arg(long)]
|
|
context_text: Option<String>,
|
|
|
|
/// Speaker of the context utterance.
|
|
#[arg(long, default_value_t = 1)]
|
|
context_speaker: u32,
|
|
|
|
/// Force CPU device even if cuda/metal features are enabled.
|
|
#[arg(long)]
|
|
cpu: bool,
|
|
|
|
/// Load a quantized GGUF model instead of the default safetensors path.
|
|
/// Path should point at the file produced by `examples/quantize`.
|
|
#[arg(long)]
|
|
quantized_gguf: Option<std::path::PathBuf>,
|
|
|
|
/// Apply a LoRA adapter (safetensors) trained via `examples/lora_train`.
|
|
/// Adapter is injected into the FP backbone before generation.
|
|
#[arg(long)]
|
|
lora: Option<std::path::PathBuf>,
|
|
|
|
/// LoRA rank — must match training. Defaults to the value from training defaults.
|
|
#[arg(long, default_value_t = 8)]
|
|
lora_rank: usize,
|
|
|
|
/// LoRA alpha — must match training.
|
|
#[arg(long, default_value_t = 16.0)]
|
|
lora_alpha: f32,
|
|
|
|
/// Disable audio post-processing (HPF + declick + LUFS normalize).
|
|
#[arg(long)]
|
|
raw: bool,
|
|
|
|
/// LUFS target for loudness normalization. Ignored if --raw.
|
|
#[arg(long, default_value_t = -16.0)]
|
|
lufs: f32,
|
|
|
|
/// Path to converted AudioSeal generator safetensors (run
|
|
/// `audioseal_convert` first). When set together with
|
|
/// `--watermark-detector`, the watermarker is wired into
|
|
/// `generate_to_wav` so output is automatically watermarked.
|
|
#[arg(long)]
|
|
watermark_generator: Option<std::path::PathBuf>,
|
|
|
|
/// Path to converted AudioSeal detector safetensors. Required for the
|
|
/// watermarker even if you only want to embed (the detector is part of
|
|
/// AudioSealWatermarker construction; future builds may make it
|
|
/// optional).
|
|
#[arg(long)]
|
|
watermark_detector: Option<std::path::PathBuf>,
|
|
|
|
/// 16-bit watermark message (decimal or 0xHEX).
|
|
#[arg(long, default_value = "0")]
|
|
watermark_message: String,
|
|
}
|
|
|
|
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)?
|
|
};
|
|
|
|
// Apply trained LoRA adapter if requested.
|
|
if let Some(lora_path) = cli.lora.as_ref() {
|
|
let lora_cfg = rtx_csm::lora::LoraConfig {
|
|
rank: cli.lora_rank,
|
|
alpha: cli.lora_alpha,
|
|
..rtx_csm::lora::LoraConfig::default()
|
|
};
|
|
let vm = candle_nn::VarMap::new();
|
|
generator.model.inner.add_lora_to_backbone(&lora_cfg, &vm)?;
|
|
rtx_csm::training::load_lora_adapter(&vm, lora_path, &device)?;
|
|
// Refresh the LoraDelta tensor handles inside the model so the loaded
|
|
// values become visible at forward time.
|
|
generator.model.inner.refresh_lora(&vm)?;
|
|
tracing::info!(
|
|
"loaded LoRA adapter from {} (rank={} alpha={})",
|
|
lora_path.display(),
|
|
cli.lora_rank,
|
|
cli.lora_alpha,
|
|
);
|
|
}
|
|
tracing::info!(
|
|
"model loaded (sr={} Hz, frame_rate={} Hz, codebooks={})",
|
|
generator.config.sample_rate,
|
|
generator.config.frame_rate_hz,
|
|
generator.config.audio_num_codebooks,
|
|
);
|
|
|
|
let mut context: Vec<Segment> = Vec::new();
|
|
if let (Some(wav), Some(txt)) = (cli.context_wav.as_ref(), cli.context_text.as_ref()) {
|
|
let audio = audio_io::load_mono_24k(wav)?;
|
|
tracing::info!("loaded context: {} samples from {}", audio.len(), wav.display());
|
|
context.push(Segment::new(cli.context_speaker, txt, audio));
|
|
}
|
|
|
|
// Optional watermarker wiring (AudioSeal + 24k↔16k resample adapter).
|
|
if let (Some(gen_path), Some(det_path)) = (
|
|
cli.watermark_generator.as_ref(),
|
|
cli.watermark_detector.as_ref(),
|
|
) {
|
|
let msg_str = cli.watermark_message.trim();
|
|
let message: u16 = if let Some(rest) = msg_str
|
|
.strip_prefix("0x")
|
|
.or_else(|| msg_str.strip_prefix("0X"))
|
|
{
|
|
u16::from_str_radix(rest, 16)?
|
|
} else {
|
|
msg_str.parse::<u16>()?
|
|
};
|
|
let gen_vb = unsafe {
|
|
candle_nn::VarBuilder::from_mmaped_safetensors(
|
|
&[gen_path],
|
|
candle_core::DType::F32,
|
|
&device,
|
|
)
|
|
}?;
|
|
let det_vb = unsafe {
|
|
candle_nn::VarBuilder::from_mmaped_safetensors(
|
|
&[det_path],
|
|
candle_core::DType::F32,
|
|
&device,
|
|
)
|
|
}?;
|
|
let inner = rtx_csm::AudioSealWatermarker::from_var_builders(
|
|
gen_vb,
|
|
det_vb,
|
|
device.clone(),
|
|
message,
|
|
)?;
|
|
// CSM produces 24 kHz; AudioSeal native is 16 kHz.
|
|
let wm = rtx_csm::ResampledWatermarker::new(inner, generator.config.sample_rate, 16_000);
|
|
generator.set_watermarker(Box::new(wm));
|
|
tracing::info!(
|
|
"watermarker installed (message=0x{:04X}, model 16 kHz, output {} Hz)",
|
|
message,
|
|
generator.config.sample_rate
|
|
);
|
|
}
|
|
|
|
let opts = GenerateOptions {
|
|
max_audio_ms: cli.max_audio_ms,
|
|
temperature: cli.temperature,
|
|
top_k: cli.top_k,
|
|
top_p: cli.top_p,
|
|
seed: cli.seed,
|
|
..GenerateOptions::default()
|
|
};
|
|
let post = if cli.raw {
|
|
PostProcess::disabled()
|
|
} else {
|
|
PostProcess {
|
|
lufs_target: Some(cli.lufs),
|
|
..PostProcess::default()
|
|
}
|
|
};
|
|
generator.generate_to_wav(&cli.text, cli.speaker, &context, opts, &post, &cli.out)?;
|
|
println!("wrote {}", cli.out.display());
|
|
Ok(())
|
|
}
|