Skip to main content

decode_silence/
decode_silence.rs

1//! Diagnostic: decode the VAE-encoded silence latent through the ported
2//! Oobleck decoder. The output should be (near-)digital silence; a large or
3//! noisy result means the VAE decoder port is wrong.
4//!
5//! Usage: cargo run --release --example decode_silence -- <model_dir>
6
7use maolan_generate::acestep::pipeline::SilenceLatent;
8use maolan_generate::acestep::vae::{OobleckDecoder, OobleckVaeConfig};
9use std::path::Path;
10
11type B = burn::backend::NdArray<f32>;
12
13fn main() -> anyhow::Result<()> {
14    let model_dir = std::env::args()
15        .nth(1)
16        .map(std::path::PathBuf::from)
17        .unwrap_or_else(|| std::path::PathBuf::from("/home/meka/repos/ace"));
18    let device = Default::default();
19
20    let vae_config = OobleckVaeConfig::load(&model_dir.join("vae_config.json"))?;
21    let vae = OobleckDecoder::<B>::from_burnpack(
22        &vae_config,
23        &model_dir.join("acestep-vae.bpk"),
24        &device,
25    )?;
26    let silence =
27        SilenceLatent::<B>::from_burnpack(&model_dir.join("silence_latent.bpk"), &device)?;
28    let [_, frames, _] = silence.silence_latent.dims();
29    println!("silence latent: {frames} frames");
30
31    let latents = silence.slice(750);
32    let audio = vae.decode(latents);
33    let [_, channels, samples] = audio.dims();
34    let values: Vec<f32> = audio
35        .into_data()
36        .convert::<f32>()
37        .to_vec()
38        .map_err(|e| anyhow::anyhow!("{e}"))?;
39    let peak = values.iter().fold(0.0_f32, |a, v| a.max(v.abs()));
40    let rms = (values.iter().map(|v| v * v).sum::<f32>() / values.len() as f32).sqrt();
41    println!("decoded: {channels} ch x {samples} samples");
42    println!("peak {peak:.6}  rms {rms:.6}");
43    println!(
44        "{}",
45        if peak < 1e-3 {
46            "OK: VAE decodes silence correctly"
47        } else {
48            "SUSPECT: VAE decoder output is far from silence"
49        }
50    );
51
52    // Also dump the latent stats so we can eyeball them.
53    let latents = SilenceLatent::<B>::from_burnpack(
54        Path::new(&model_dir).join("silence_latent.bpk").as_path(),
55        &device,
56    )?
57    .slice(4);
58    let v: Vec<f32> = latents
59        .into_data()
60        .convert::<f32>()
61        .to_vec()
62        .map_err(|e| anyhow::anyhow!("{e}"))?;
63    let lpeak = v.iter().fold(0.0_f32, |a, v| a.max(v.abs()));
64    let lrms = (v.iter().map(|v| v * v).sum::<f32>() / v.len() as f32).sqrt();
65    println!("latent peak {lpeak:.6} rms {lrms:.6}");
66    Ok(())
67}