Skip to main content

enhance_parity/
enhance_parity.rs

1//! Parity harness: run the Rust FastEnhancer on a raw f32 LE mono 48 kHz input and dump
2//! the enhanced waveform as raw f32 LE, for byte-level comparison against the pinned
3//! PyTorch oracle dumps.
4//!
5//! Usage: enhance_parity <weights.safetensors> <input.f32> <output.f32>
6
7use std::io::{Read, Write};
8
9fn main() {
10    let mut args = std::env::args().skip(1);
11    let weights = args.next().expect("weights path");
12    let input = args.next().expect("input raw-f32 path");
13    let output = args.next().expect("output raw-f32 path");
14
15    let enhancer = ftts_artifacts::enhance_loader::open_enhancer(&weights)
16        .unwrap_or_else(|error| panic!("cannot load {weights}: {error}"));
17
18    let mut bytes = Vec::new();
19    std::fs::File::open(&input)
20        .and_then(|mut f| f.read_to_end(&mut bytes))
21        .unwrap_or_else(|error| panic!("cannot read {input}: {error}"));
22    assert!(bytes.len() % 4 == 0, "raw f32 input must be 4-byte aligned");
23    let wav: Vec<f32> = bytes
24        .as_chunks::<4>()
25        .0
26        .iter()
27        .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
28        .collect();
29
30    let started = std::time::Instant::now();
31    let out = enhancer.enhance_48k(&wav);
32    let elapsed = started.elapsed();
33    let audio_seconds = wav.len() as f64 / 48_000.0;
34    eprintln!(
35        "enhanced {:.2}s of audio in {:.3}s (rtf {:.4})",
36        audio_seconds,
37        elapsed.as_secs_f64(),
38        elapsed.as_secs_f64() / audio_seconds,
39    );
40
41    let mut out_bytes = Vec::with_capacity(out.len() * 4);
42    for v in &out {
43        out_bytes.extend_from_slice(&v.to_le_bytes());
44    }
45    std::fs::File::create(&output)
46        .and_then(|mut f| f.write_all(&out_bytes))
47        .unwrap_or_else(|error| panic!("cannot write {output}: {error}"));
48}