Files
rustytorch/crates/models/rtx-csm/examples/audioseal_apply.rs
T
osobhandClaude Opus 4.7 c0ca1e1f8a rtx-csm: Phase 6f.lora + 6f.wm — voice clone + watermark in converse_server
Wires two more capabilities into examples/converse_server.rs:

  --lora <path> [--lora-rank N --lora-alpha A]   inject voice-clone adapter
  --watermark-generator <path>                    \
  --watermark-detector  <path>                     embed AudioSeal on every
  [--watermark-message 0xCAFE]                     assistant utterance

LoRA path: parallels examples/generate.rs — load Generator (FP), build
LoraConfig, add_lora_to_backbone(VarMap), load_lora_adapter, refresh_lora.
Voice clones now work in the conversation pipeline. Combining with
--quantized-gguf is rejected at boot; the LoRA-on-Q8 path works in
generate but the server hasn't been audited so it's gated for now.

Watermark path: parallels examples/generate.rs — load AudioSeal generator
+ detector safetensors via VarBuilder::from_mmaped_safetensors, build
AudioSealWatermarker, wrap in ResampledWatermarker(24k↔16k), install via
generator.set_watermarker(...). converse.rs::synthesize already calls
the watermarker per-utterance, so no plumbing changes needed downstream.

Verified end-to-end:
  * boot server with --watermark-generator/-detector --watermark-message 0xCAFE
  * single conversation turn (10.43s LibriSpeech in -> 4.24s assistant out)
  * detect on response WAV: mean_presence=0.9995, decoded=0xCAFE,
    16/16 message bits matching after full STT->LLM->TTS->post->watermark
    ->24k->WS->wav round-trip.

Also adds audioseal_apply --detect-only flag (skip embed, run detector
against arbitrary WAV) — used to verify the round-trip above. --out is
now optional and only required when not in detect-only mode.

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

162 lines
5.5 KiB
Rust

//! Apply an AudioSeal watermark to any input WAV.
//!
//! Handles arbitrary source sample rates by resampling to 16 kHz (AudioSeal
//! native), embedding the watermark, then resampling back to the original
//! rate and writing the output. The CSM TTS pipeline produces 24 kHz audio,
//! so the typical use is:
//!
//! ```bash
//! cargo run -p rtx-csm --release --example generate -- \
//! --text "Hello." --out /tmp/hello.wav
//!
//! cargo run -p rtx-csm --release --example audioseal_apply -- \
//! --generator /tmp/audioseal_generator.safetensors \
//! --detector /tmp/audioseal_detector.safetensors \
//! --message 0xBEEF \
//! --in /tmp/hello.wav \
//! --out /tmp/hello_watermarked.wav
//! ```
//!
//! The example also runs the detector on the watermarked output to verify
//! the round-trip (mean_presence + decoded message).
use anyhow::{Context, Result};
use candle_core::{DType, Device};
use clap::Parser;
use rtx_csm::{audio_io, audioseal::AudioSealWatermarker, watermark::Watermarker};
use std::path::PathBuf;
const AUDIOSEAL_RATE: u32 = 16_000;
#[derive(Debug, Parser)]
#[command(name = "audioseal_apply")]
struct Cli {
/// Path to converted generator safetensors.
#[arg(long)]
generator: PathBuf,
/// Path to converted detector safetensors.
#[arg(long)]
detector: PathBuf,
/// 16-bit message payload (decimal or 0xHEX).
#[arg(long, default_value = "0xBEEF")]
message: String,
/// Input WAV (any rate, any channels).
#[arg(long = "in")]
input: PathBuf,
/// Output WAV path. The output is written at the SOURCE sample rate
/// (resample to 16 kHz happens internally only for the watermarker).
/// Ignored in `--detect-only` mode.
#[arg(long)]
out: Option<PathBuf>,
/// Source sample rate of the input WAV (default 24000 = CSM-1B native).
#[arg(long, default_value_t = 24_000)]
source_rate: u32,
/// Skip the embed step and just run the detector against the input
/// WAV. Useful for verifying that an externally watermarked file
/// (e.g., output of converse_server with --watermark-* flags) carries
/// a recoverable signature.
#[arg(long)]
detect_only: bool,
/// Force CPU device.
#[arg(long)]
cpu: bool,
}
fn parse_message(s: &str) -> Result<u16> {
let s = s.trim();
let v = if let Some(rest) = s.strip_prefix("0x").or_else(|| s.strip_prefix("0X")) {
u16::from_str_radix(rest, 16)?
} else {
s.parse::<u16>()?
};
Ok(v)
}
fn main() -> Result<()> {
tracing_subscriber::fmt().init();
let cli = Cli::parse();
let device = if cli.cpu {
Device::Cpu
} else if candle_core::utils::metal_is_available() {
Device::new_metal(0)?
} else {
Device::Cpu
};
let message = parse_message(&cli.message)?;
// Load source at native rate, then to 16 kHz for AudioSeal.
let src_native = audio_io::load_mono_at_rate(&cli.input, cli.source_rate)
.context("loading source at native rate")?;
let src_16k = audio_io::resample(&src_native, cli.source_rate, AUDIOSEAL_RATE)
.context("resample source -> 16 kHz")?;
println!(
"loaded {}: {} samples @ {} Hz ({} samples @ 16 kHz)",
cli.input.display(),
src_native.len(),
cli.source_rate,
src_16k.len(),
);
// Load model.
let gen_vb = unsafe {
candle_nn::VarBuilder::from_mmaped_safetensors(&[&cli.generator], DType::F32, &device)
}
.context("opening generator safetensors")?;
let det_vb = unsafe {
candle_nn::VarBuilder::from_mmaped_safetensors(&[&cli.detector], DType::F32, &device)
}
.context("opening detector safetensors")?;
let wm = AudioSealWatermarker::from_var_builders(gen_vb, det_vb, device.clone(), message)?;
println!("loaded AudioSeal (message=0x{:04X})", message);
if cli.detect_only {
let result = wm.detect(&src_16k).context("watermark detect")?;
println!(
"detect-only: mean_presence={:.4}, decoded=0x{:04X} (expected 0x{:04X})",
result.mean_presence,
result.message.unwrap_or(0),
message
);
let xor = result.message.unwrap_or(0) ^ message;
let bits_match = 16 - xor.count_ones() as usize;
println!("message bits matching: {bits_match}/16");
return Ok(());
}
let out_path = cli
.out
.as_ref()
.context("--out is required unless --detect-only is set")?;
// Embed watermark at 16 kHz.
let wm_16k = wm.embed(&src_16k).context("watermark embed")?;
// Resample back to source rate and write output.
let wm_out = audio_io::resample(&wm_16k, AUDIOSEAL_RATE, cli.source_rate)
.context("resample 16 kHz -> source rate")?;
audio_io::write_wav_mono(out_path.as_path(), &wm_out, cli.source_rate)
.context("write watermarked WAV")?;
println!(
"wrote {} ({} samples @ {} Hz)",
out_path.display(),
wm_out.len(),
cli.source_rate
);
// Verify round-trip: re-resample to 16 kHz and detect.
let probe_16k = audio_io::resample(&wm_out, cli.source_rate, AUDIOSEAL_RATE)
.context("resample for detect")?;
let result = wm.detect(&probe_16k).context("watermark detect")?;
println!(
"round-trip detect: mean_presence={:.4}, decoded=0x{:04X} (expected 0x{:04X})",
result.mean_presence,
result.message.unwrap_or(0),
message
);
let xor = result.message.unwrap_or(0) ^ message;
let bits_match = 16 - xor.count_ones() as usize;
println!("message bits matching: {bits_match}/16");
Ok(())
}