//! Phase 13.9 — slice 3 smoke test for the wav2vec2 candle port. //! //! Loads `facebook/wav2vec2-base-960h` from HF Hub, runs forward on a //! real audio file, and prints the greedy CTC decode (transcript). This //! is the first real-weight integration of the port: every safetensors //! key must map to a candle param of matching shape, and the resulting //! transcript should be intelligible English. //! //! Usage: //! ```bash //! cargo run -p rtx-csm --release --features metal --example wav2vec2_smoke -- \ //! --in /tmp/asr_test.flac //! ``` use anyhow::{Context, Result}; use candle_core::{Device, Tensor}; use clap::Parser; use hf_hub::api::sync::Api; use rtx_csm::audio_io; use rtx_csm::wav2vec2::{ctc_greedy_decode, Wav2Vec2, VOCAB_960H}; use std::path::PathBuf; const REPO: &str = "facebook/wav2vec2-base-960h"; const SAFETENSORS_FILE: &str = "model.safetensors"; #[derive(Debug, Parser)] struct Cli { #[arg(long = "in", default_value = "/tmp/asr_test.flac")] input: PathBuf, } fn main() -> Result<()> { tracing_subscriber::fmt().init(); let cli = Cli::parse(); let device = if candle_core::utils::metal_is_available() { Device::new_metal(0)? } else { Device::Cpu }; eprintln!("device: {device:?}"); let api = Api::new().context("hf_hub init")?; let path = api .model(REPO.to_string()) .get(SAFETENSORS_FILE) .with_context(|| format!("download {SAFETENSORS_FILE} from {REPO}"))?; eprintln!("safetensors: {}", path.display()); let load_t = std::time::Instant::now(); let model = Wav2Vec2::load_from_safetensors(&path, &device)?; eprintln!("loaded model in {:.2}s", load_t.elapsed().as_secs_f32()); // Load + resample audio to 16 kHz. let pcm = audio_io::load_mono_at_rate(&cli.input, 16_000).context("load audio")?; eprintln!( "audio: {} samples ({:.2}s @ 16 kHz)", pcm.len(), pcm.len() as f32 / 16_000.0 ); // wav2vec2 expects pre-normalized inputs (zero mean unit variance per // utterance, per HF's Wav2Vec2FeatureExtractor.do_normalize). let mean = pcm.iter().sum::() / pcm.len().max(1) as f32; let var = pcm.iter().map(|x| (x - mean).powi(2)).sum::() / pcm.len().max(1) as f32; let std = var.sqrt().max(1e-7); let norm: Vec = pcm.iter().map(|x| (x - mean) / std).collect(); let audio_t = Tensor::from_vec(norm, (1, 1, pcm.len()), &device)?; let fwd_t = std::time::Instant::now(); let logits = model.forward(&audio_t)?; eprintln!("forward in {} ms", fwd_t.elapsed().as_millis()); eprintln!("logits shape: {:?}", logits.shape().dims()); let dec_t = std::time::Instant::now(); let transcript = ctc_greedy_decode(&logits, VOCAB_960H)?; eprintln!("ctc decode in {} ms", dec_t.elapsed().as_millis()); println!(); println!("=== transcript ==="); println!("{}", transcript.trim()); println!(); Ok(()) }