//! List all tensor keys + shapes from facebook/audioseal generator/detector //! .pth checkpoints. Used to drive the converter's key remapping table. //! //! Usage: //! ``` //! cargo run -p rtx-csm --release --example audioseal_inspect -- --which generator //! cargo run -p rtx-csm --release --example audioseal_inspect -- --which detector //! cargo run -p rtx-csm --release --example audioseal_inspect -- --path /local.pth //! ``` use anyhow::Result; use candle_core::pickle; use clap::{Parser, ValueEnum}; use rtx_csm::hub; use std::path::PathBuf; #[derive(Debug, Clone, ValueEnum)] enum Which { Generator, Detector, WavlmSv, } #[derive(Debug, Parser)] #[command(name = "audioseal_inspect", about = "Dump AudioSeal .pth tensor keys + shapes")] struct Cli { /// Which checkpoint to fetch from facebook/audioseal. #[arg(long, value_enum, default_value = "generator")] which: Which, /// Path override; if set, ignore --which and read this file directly. #[arg(long)] path: Option, /// Show only keys matching this substring. #[arg(long)] filter: Option, /// Cap on number of keys printed (0 = unlimited). #[arg(long, default_value_t = 0)] limit: usize, /// Optional dict key to descend into (e.g. "model", "best_state", "xp.cfg"). #[arg(long)] key: Option, /// Print the raw pickle object tree before tensor extraction. #[arg(long)] verbose: bool, } fn main() -> Result<()> { tracing_subscriber::fmt().init(); let cli = Cli::parse(); let path = match cli.path { Some(p) => p, None => match cli.which { Which::Generator => hub::resolve_audioseal_generator()?, Which::Detector => hub::resolve_audioseal_detector()?, Which::WavlmSv => hub::resolve_wavlm_sv()?, }, }; println!("inspecting: {}", path.display()); let infos = pickle::read_pth_tensor_info(&path, cli.verbose, cli.key.as_deref())?; println!("found {} tensor entries", infos.len()); let mut printed = 0usize; for info in &infos { if let Some(f) = cli.filter.as_ref() { if !info.name.contains(f) { continue; } } println!( " {:<70} dtype={:?} shape={:?}", info.name, info.dtype, info.layout ); printed += 1; if cli.limit > 0 && printed >= cli.limit { println!(" ... (truncated at --limit {})", cli.limit); break; } } Ok(()) }