1use burn::tensor::{Int, Tensor, TensorData};
7use maolan_generate::acestep::condition::AceStepCondition;
8use maolan_generate::acestep::config::AceStepConfig;
9use maolan_generate::acestep::pipeline::SilenceLatent;
10use maolan_generate::acestep::vae::{OobleckDecoder, OobleckVaeConfig};
11use std::path::Path;
12
13type B = burn::backend::NdArray<f32>;
14
15fn rms(values: &[f32]) -> f32 {
16 (values.iter().map(|v| v * v).sum::<f32>() / values.len() as f32).sqrt()
17}
18
19fn run_case(
20 name: &str,
21 codes: &[u32],
22 condition: &AceStepCondition<B>,
23 vae: &OobleckDecoder<B>,
24 silence: &SilenceLatent<B>,
25 device: &burn::tensor::Device<B>,
26) -> anyhow::Result<()> {
27 let n = codes.len();
28 let tensor = Tensor::<B, 2, Int>::from_data(TensorData::new(codes.to_vec(), [1, n]), device);
29 let hints = condition.codes_to_hints(tensor);
30 let silence_frames = silence.slice(n * 5);
31 let diff = (hints.clone() - silence_frames)
32 .abs()
33 .into_data()
34 .convert::<f32>()
35 .to_vec::<f32>()
36 .map_err(|e| anyhow::anyhow!("{e}"))?;
37 let hint_values: Vec<f32> = hints
38 .clone()
39 .into_data()
40 .convert::<f32>()
41 .to_vec()
42 .map_err(|e| anyhow::anyhow!("{e}"))?;
43 let audio = vae.decode(hints);
44 let audio_values: Vec<f32> = audio
45 .into_data()
46 .convert::<f32>()
47 .to_vec()
48 .map_err(|e| anyhow::anyhow!("{e}"))?;
49 let peak = audio_values.iter().fold(0.0_f32, |a, v| a.max(v.abs()));
50 let mean_abs_diff = diff.iter().sum::<f32>() / diff.len() as f32;
51 println!(
52 "{name}: hints rms {:.4} mean|hints-silence| {:.4} audio rms {:.4} audio peak {:.4}",
53 rms(&hint_values),
54 mean_abs_diff,
55 rms(&audio_values),
56 peak
57 );
58 Ok(())
59}
60
61fn main() -> anyhow::Result<()> {
62 let model_dir = std::env::args()
63 .nth(1)
64 .unwrap_or_else(|| "/home/meka/repos/ace".to_string());
65 let model_dir = Path::new(&model_dir);
66 let device = Default::default();
67
68 let dit_config = AceStepConfig::load(&model_dir.join("dit_config.json"))?;
69 let condition = AceStepCondition::<B>::from_burnpack(
70 &dit_config,
71 &model_dir.join("acestep-condition.bpk"),
72 &device,
73 )?;
74 let vae_config = OobleckVaeConfig::load(&model_dir.join("vae_config.json"))?;
75 let vae = OobleckDecoder::<B>::from_burnpack(
76 &vae_config,
77 &model_dir.join("acestep-vae.bpk"),
78 &device,
79 )?;
80 let silence =
81 SilenceLatent::<B>::from_burnpack(&model_dir.join("silence_latent.bpk"), &device)?;
82
83 let mut state = 42_u64;
85 let mut rand_codes = Vec::new();
86 for _ in 0..20 {
87 state = state
88 .wrapping_mul(6364136223846793005)
89 .wrapping_add(1442695040888963407);
90 rand_codes.push(((state >> 33) % 64000) as u32);
91 }
92
93 let lm_codes: Vec<u32> = vec![
94 61890, 51649, 53753, 53753, 56314, 53754, 55802, 4538, 4538, 4538, 5050, 5050, 5050, 5050,
95 5050, 5050, 5050, 5050, 5050, 43513,
96 ];
97
98 for code in [0u32, 12345, 53754] {
101 let tensor = Tensor::<B, 2, Int>::from_data(TensorData::new(vec![code], [1, 1]), &device);
102 let hints = condition.codes_to_hints(tensor); let values: Vec<f32> = hints
104 .into_data()
105 .convert::<f32>()
106 .to_vec()
107 .map_err(|e| anyhow::anyhow!("{e}"))?;
108 let frame_means: Vec<f32> = values
109 .as_chunks::<64>()
110 .0
111 .iter()
112 .map(|f| f.iter().sum::<f32>() / 64.0)
113 .collect();
114 let frame_rms: Vec<f32> = values.as_chunks::<64>().0.iter().map(|f| rms(f)).collect();
115 println!(
116 "code {code:>6}: frame means {:?}",
117 frame_means
118 .iter()
119 .map(|v| format!("{v:.4}"))
120 .collect::<Vec<_>>()
121 );
122 println!(
123 " frame rms {:?}",
124 frame_rms
125 .iter()
126 .map(|v| format!("{v:.4}"))
127 .collect::<Vec<_>>()
128 );
129 }
130
131 let tensor =
133 Tensor::<B, 2, Int>::from_data(TensorData::new(rand_codes.clone(), [1, 20]), &device);
134 let hints = condition.codes_to_hints(tensor); let values: Vec<f32> = hints
136 .into_data()
137 .convert::<f32>()
138 .to_vec()
139 .map_err(|e| anyhow::anyhow!("{e}"))?;
140 let (frames, _) = values.as_chunks::<64>();
141 let frame_means: Vec<f32> = frames
143 .iter()
144 .map(|f| f.iter().sum::<f32>() / 64.0)
145 .collect();
146 let tm = frame_means.iter().sum::<f32>() / frame_means.len() as f32;
147 let time_var =
148 frame_means.iter().map(|v| (v - tm) * (v - tm)).sum::<f32>() / frame_means.len() as f32;
149 let chan_var = frames
150 .iter()
151 .map(|f| {
152 let m = f.iter().sum::<f32>() / 64.0;
153 f.iter().map(|v| (v - m) * (v - m)).sum::<f32>() / 64.0
154 })
155 .sum::<f32>()
156 / frames.len() as f32;
157 println!("random-code hints: time var {time_var:.6} channel var {chan_var:.6}");
158
159 run_case(
160 "random codes ",
161 &rand_codes,
162 &condition,
163 &vae,
164 &silence,
165 &device,
166 )?;
167 run_case(
168 "LM codes ",
169 &lm_codes,
170 &condition,
171 &vae,
172 &silence,
173 &device,
174 )?;
175 run_case(
176 "constant 0 ",
177 &[0u32; 20],
178 &condition,
179 &vae,
180 &silence,
181 &device,
182 )?;
183 Ok(())
184}