//! Inspect specific tensors inside a converted WavLM-SV safetensors. //! Useful for sanity-checking weight loading and key naming. use anyhow::Result; use candle_core::Device; use clap::Parser; use std::path::PathBuf; #[derive(Debug, Parser)] struct Cli { #[arg(long)] weights: PathBuf, /// Tensor name to dump (e.g. "layer_weights", "projector.weight"). #[arg(long)] key: String, /// Apply softmax along dim=0 before printing (for layer_weights). #[arg(long)] softmax: bool, /// Number of leading elements to print (default: all). #[arg(long)] head: Option, } fn main() -> Result<()> { let cli = Cli::parse(); let tensors = candle_core::safetensors::load(&cli.weights, &Device::Cpu)?; let t = tensors .get(&cli.key) .ok_or_else(|| anyhow::anyhow!("key not found: {}", cli.key))?; println!("{}: dtype={:?}, shape={:?}", cli.key, t.dtype(), t.dims()); let to_print = if cli.softmax { candle_nn::ops::softmax(t, 0)? } else { t.clone() }; let v: Vec = to_print .flatten_all()? .to_dtype(candle_core::DType::F32)? .to_vec1()?; let limit = cli.head.unwrap_or(v.len()).min(v.len()); let head: Vec<&f32> = v.iter().take(limit).collect(); println!("first {limit} values: {head:?}"); if v.len() > limit { let tail: Vec<&f32> = v.iter().rev().take(8).collect(); println!("(last 8 values, reversed): {tail:?}"); } let sum: f32 = v.iter().sum(); let max = v.iter().cloned().fold(f32::NEG_INFINITY, f32::max); let min = v.iter().cloned().fold(f32::INFINITY, f32::min); println!("sum={sum:.6}, min={min:.6}, max={max:.6}"); Ok(()) }