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]>
75 lines
2.4 KiB
Rust
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(())
|
|
}
|