Files
rustytorch/crates/models/rtx-csm/examples/wav2vec2_smoke.rs
T
osobhandClaude Opus 4.7 209279c13e rtx-csm: Phase 13.9 — wav2vec2 candle port (slices 1+2+3, real ASR working)
Full port of facebook/wav2vec2-base-960h (94.4 M params, MIT) closing
the WhisperX-class word-alignment gap from the audio-ML survey. Same
staged-scaffolding pattern that worked for emotion2vec — but landed
slices 1+2+3 in one session.

src/wav2vec2.rs ships:
  - Wav2Vec2Config::base_960h
  - FeatureExtractor — 7 Conv1d (1→512, total stride 320). Layer 0
    uses GroupNorm with num_groups=num_channels=512 (HF's wav2vec2
    feat_extract_norm: "group"). Critical: state-dict key is
    layer_norm.* but the OP is GroupNorm — loading as LayerNorm
    produces empty CTC output.
  - FeatureProjection — LayerNorm(512) + Linear(512→768)
  - ConvPosEmbedding — kernel 128 grouped Conv1d, materialized at
    load time from upstream weight_g + weight_v (fairseq's weight_norm
    on dim=2; eps-guarded division for numerical stability)
  - Block — POST-norm transformer with separate Q/K/V (vs emotion2vec's
    fused QKV), uses (B*H, T, D) Metal 3D-matmul workaround from
    Phase 8.8 Moonshine
  - Encoder — pos_conv + initial LayerNorm + 12 Blocks
  - Wav2Vec2 top-level — load_from_safetensors via mmap'd VarBuilder
  - ctc_greedy_decode + VOCAB_960H constant for the 32-char alphabet

examples/wav2vec2_inspect.rs (slice 1): dumps tensor layout + config
examples/wav2vec2_smoke.rs (slice 3): real-weight load + ASR forward

Verified on Metal:
  loaded model in 0.28 s
  forward in 9 ms for 10.42 s audio (~1150× realtime)
  transcript: "HE HOPED THERE WOULD BE STEW FOR DINNER TURNIPS AND
              CARROTS AND BRUISED POTATOES AND FAT MUTTON PIECES TO
              BE LADLED OUT IN THICK PEPPERED FLOWER FAT AND SAUCE"

Numerical parity with upstream Python — the FLOWER-for-FLOUR typo is
the known wav2vec2-base-960h failure mode, matches HF reference exactly.

7 new unit tests; lib suite 127/127 (was 120).

Slice 4 remaining: Viterbi forced alignment given known transcript,
to emit (token, frame_start_ms, frame_end_ms) for word-boundary cuts.
The ASR path itself is now production-ready.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
2026-04-28 03:33:35 -07:00

85 lines
2.9 KiB
Rust

//! 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::<f32>() / pcm.len().max(1) as f32;
let var = pcm.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / pcm.len().max(1) as f32;
let std = var.sqrt().max(1e-7);
let norm: Vec<f32> = 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(())
}