Files
rustytorch/crates/models/rtx-csm/examples/audioseal_convert.rs
T
osobhandClaude Opus 4.7 938b54a2b0 rtx-csm: AudioSeal watermark — Rust port end-to-end
SEANet generator + detector matching `facebook/audioseal` reference layout
(weight_norm-merged via pure-Rust pickle reader). Verified on real CSM
speech: mean_presence=0.9943, 16/16 message bits decoded.

- src/audioseal.rs: SeanetEncoder (4-stage strided downsample, 2-layer
  LSTM bottleneck at 512 channels, 128-dim projection), MsgProcessor
  (16-bit message via embedding sum + broadcast-add), SeanetDecoder,
  Generator (encoder+msg+decoder), Detector (encoder + single 320×
  reverse_convolution + 1×1 head). Padding mirrors audiocraft
  _get_extra_padding_for_conv1d exactly.
- src/audioseal_convert.rs: candle_core::pickle reads .pth directly;
  merge_weight_norm computes g*v/‖v‖ over all axes except 0; writes
  flat safetensors keyed identically to what Generator/Detector read.
- examples/audioseal_inspect.rs: dumps tensor keys + shapes.
- examples/audioseal_convert.rs: HF download + convert CLI.
- examples/audioseal_demo.rs: load + embed + detect on real WAV or
  synthetic burst, optionally writes watermarked WAV.
- audio_io.rs gains generic load_mono_at_rate, resample, write_wav_mono
  (16 kHz path needed for AudioSeal).

12 new unit tests + 2 converter tests; 63 lib tests total green.

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

75 lines
2.4 KiB
Rust

//! Convert facebook/audioseal `.pth` → flat safetensors with weight_norm
//! merged. Output is consumable by `audioseal::Generator::new` /
//! `audioseal::Detector::new` via a `VarBuilder` over the safetensors file.
//!
//! Usage:
//! ```
//! cargo run -p rtx-csm --release --example audioseal_convert -- \
//! --generator-out /tmp/audioseal_generator.safetensors \
//! --detector-out /tmp/audioseal_detector.safetensors
//! ```
use anyhow::{Context, Result};
use clap::Parser;
use rtx_csm::{audioseal_convert, hub};
use std::path::PathBuf;
#[derive(Debug, Parser)]
#[command(name = "audioseal_convert")]
struct Cli {
/// Optional override; defaults to HF-fetched facebook/audioseal generator_base.pth.
#[arg(long)]
generator_in: Option<PathBuf>,
/// Optional override; defaults to HF-fetched detector_base.pth.
#[arg(long)]
detector_in: Option<PathBuf>,
/// Output safetensors for the generator (post-merge).
#[arg(long)]
generator_out: PathBuf,
/// Output safetensors for the detector (post-merge).
#[arg(long)]
detector_out: PathBuf,
}
fn main() -> Result<()> {
tracing_subscriber::fmt().init();
let cli = Cli::parse();
let gen_in = match cli.generator_in {
Some(p) => p,
None => hub::resolve_audioseal_generator()
.context("resolve_audioseal_generator")?,
};
let det_in = match cli.detector_in {
Some(p) => p,
None => hub::resolve_audioseal_detector().context("resolve_audioseal_detector")?,
};
tracing::info!(
"converting generator: {} -> {}",
gen_in.display(),
cli.generator_out.display()
);
let gen_report = audioseal_convert::convert_pth(&gen_in, &cli.generator_out, Some("model"))?;
println!(
"generator: merged {} weight_norm pairs, {} passthrough, {} total tensors",
gen_report.merged_weight_norm_pairs,
gen_report.passthrough_tensors,
gen_report.total_tensors_written
);
tracing::info!(
"converting detector: {} -> {}",
det_in.display(),
cli.detector_out.display()
);
let det_report = audioseal_convert::convert_pth(&det_in, &cli.detector_out, Some("model"))?;
println!(
"detector: merged {} weight_norm pairs, {} passthrough, {} total tensors",
det_report.merged_weight_norm_pairs,
det_report.passthrough_tensors,
det_report.total_tensors_written
);
Ok(())
}