decode_silence/
decode_silence.rs1use 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 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}