//! 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, /// Optional override; defaults to HF-fetched detector_base.pth. #[arg(long)] detector_in: Option, /// 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(()) }