1use burn::tensor::{Int, Tensor, TensorData};
11use maolan_generate::acestep::condition::AceStepCondition;
12use maolan_generate::acestep::config::AceStepConfig;
13use maolan_generate::acestep::dit::{AceStepDiT, SFT_TIMESTEPS, TURBO_TIMESTEPS};
14use maolan_generate::acestep::pipeline::{SilenceLatent, build_dit_text_prompt, build_metas_block};
15use maolan_generate::acestep::qwen3::{Qwen3Config, Qwen3Model};
16use maolan_generate::acestep::vae::{OobleckDecoder, OobleckVaeConfig};
17use std::path::{Path, PathBuf};
18
19type B = burn::backend::Wgpu<f32, i64, u32>;
20
21const ORACLE_CODES: [u32; 19] = [
22 37162, 24482, 53681, 61696, 48912, 36323, 57341, 14786, 18321, 50090, 36760, 36697, 36056,
23 38160, 35080, 56000, 32927, 35847, 37127,
24];
25
26fn dump_tensor(path: &Path, name: &str, data: &[f32], shape: &[i32]) {
27 let mut out = Vec::new();
28 out.extend_from_slice(&(shape.len() as i32).to_le_bytes());
29 for &d in shape {
30 out.extend_from_slice(&d.to_le_bytes());
31 }
32 for &v in data {
33 out.extend_from_slice(&v.to_le_bytes());
34 }
35 std::fs::write(path.join(format!("{name}.bin")), out).expect("dump write");
36}
37
38fn load_bin(path: &Path) -> (Vec<i32>, Vec<f32>) {
39 let raw = std::fs::read(path).expect("read bin");
40 let ndims = i32::from_le_bytes(raw[0..4].try_into().unwrap()) as usize;
41 let mut shape = Vec::new();
42 for i in 0..ndims {
43 shape.push(i32::from_le_bytes(
44 raw[4 + 4 * i..8 + 4 * i].try_into().unwrap(),
45 ));
46 }
47 let data: Vec<f32> = raw[4 + 4 * ndims..]
48 .as_chunks::<4>()
49 .0
50 .iter()
51 .map(|c| f32::from_le_bytes(*c))
52 .collect();
53 (shape, data)
54}
55
56fn vec_of<B: burn::tensor::backend::Backend>(t: Tensor<B, 3>) -> Vec<f32> {
57 t.into_data()
58 .convert::<f32>()
59 .to_vec()
60 .expect("tensor to vec")
61}
62
63fn main() -> anyhow::Result<()> {
64 let mut args = std::env::args().skip(1);
65 let model_dir = PathBuf::from(args.next().unwrap_or_else(|| "/home/meka/repos/ace".into()));
66 let out_dir = PathBuf::from(args.next().unwrap_or_else(|| "/var/tmp/dump_ours".into()));
67 let noise_path = args.next().map(PathBuf::from);
68 std::fs::create_dir_all(&out_dir)?;
69
70 let device = burn::backend::wgpu::WgpuDevice::default();
71 burn::backend::wgpu::init_setup::<burn::backend::wgpu::graphics::Vulkan>(
72 &device,
73 burn::backend::wgpu::RuntimeOptions {
74 memory_config: burn::backend::wgpu::MemoryConfiguration::ExclusivePages,
75 ..Default::default()
76 },
77 );
78
79 let variant_sft = std::env::var("MAOLAN_ACESTEP_VARIANT")
80 .map(|v| v == "sft")
81 .unwrap_or(false);
82 let prefix = if variant_sft { "sft-" } else { "" };
83
84 let dit_config = AceStepConfig::load(&model_dir.join(format!("{prefix}dit_config.json")))?;
85 let text_config = Qwen3Config::load(&model_dir.join("qwen3_config.json"))?;
86 let vae_config = OobleckVaeConfig::load(&model_dir.join("vae_config.json"))?;
87 eprintln!("loading components (prefix '{prefix}')...");
88 let text_encoder = Qwen3Model::<B>::from_burnpack(
89 &text_config,
90 &model_dir.join("qwen3-encoder.bpk"),
91 &device,
92 )?;
93 let condition = AceStepCondition::<B>::from_burnpack(
94 &dit_config,
95 &model_dir.join(format!("{prefix}acestep-condition.bpk")),
96 &device,
97 )?;
98 let dit = AceStepDiT::<B>::from_burnpack(
99 &dit_config,
100 &model_dir.join(format!("{prefix}acestep-dit.bpk")),
101 &device,
102 )?;
103 let vae = OobleckDecoder::<B>::from_burnpack(
104 &vae_config,
105 &model_dir.join("acestep-vae.bpk"),
106 &device,
107 )?;
108 let silence =
109 SilenceLatent::<B>::from_burnpack(&model_dir.join("silence_latent.bpk"), &device)?;
110 let tokenizer = tokie::Tokenizer::from_json(model_dir.join("tokenizer.json"))
111 .map_err(|e| anyhow::anyhow!("tokenizer: {e}"))?;
112
113 let metas = build_metas_block(Some(120.0), Some("A minor"), Some("4/4"), 4);
115 let text_prompt = build_dit_text_prompt("Metal guitar with a lot of distortion", &metas);
116 let ids = tokenizer.encode(&text_prompt, false).ids;
117 let mut ids: Vec<i64> = ids.into_iter().map(i64::from).collect();
118 ids.push(151643); let n_text = ids.len();
120 let ids_t = Tensor::<B, 2, Int>::from_data(TensorData::new(ids, [1, n_text]), &device);
121 let text_hidden = text_encoder.forward(ids_t, true);
122 let text_hidden_vec = vec_of(text_hidden.clone());
123 dump_tensor(
124 &out_dir,
125 "text_hidden",
126 &text_hidden_vec,
127 &[n_text as i32, 1024],
128 );
129
130 let lyric_ids = tokenizer
131 .encode(
132 maolan_generate::acestep::pipeline::INSTRUMENTAL_LYRIC_PROMPT,
133 false,
134 )
135 .ids;
136 let mut lyric_ids: Vec<i64> = lyric_ids.into_iter().map(i64::from).collect();
137 lyric_ids.push(151643); let n_lyric = lyric_ids.len();
139 let lyric_ids_t =
140 Tensor::<B, 2, Int>::from_data(TensorData::new(lyric_ids, [1, n_lyric]), &device);
141 let lyric_hidden = text_encoder.embed_tokens.forward(lyric_ids_t);
143 let lyric_embed_vec = vec_of(lyric_hidden.clone());
144 dump_tensor(
145 &out_dir,
146 "lyric_embed",
147 &lyric_embed_vec,
148 &[n_lyric as i32, 1024],
149 );
150 let lyric_mask = Tensor::<B, 2, Int>::ones([1, n_lyric], &device);
151
152 let mut oracle_context: Option<Tensor<B, 3>> = None;
154 let enc = if let Some(dir) = std::env::var_os("MAOLAN_ORACLE_DUMP_DIR") {
155 let dir = PathBuf::from(dir);
156 let (enc_shape, enc_data) = load_bin(&dir.join("enc_hidden.bin"));
157 eprintln!("oracle enc_hidden: {enc_shape:?}");
158 let (ctx_shape, ctx_data) = load_bin(&dir.join("context.bin"));
159 eprintln!("oracle context: {ctx_shape:?}");
160 let enc = Tensor::<B, 3>::from_data(
161 TensorData::new(enc_data, [1, enc_shape[0] as usize, enc_shape[1] as usize]),
162 &device,
163 );
164 let ctx = Tensor::<B, 3>::from_data(TensorData::new(ctx_data, [1, 96, 128]), &device);
165 oracle_context = Some(ctx);
166 enc
167 } else {
168 condition.encode(
169 text_hidden,
170 lyric_hidden,
171 lyric_mask,
172 silence.timbre_reference(),
173 )
174 };
175 let [_, enc_len, enc_dim] = enc.dims();
176 let enc_vec = vec_of(enc.clone());
177 dump_tensor(
178 &out_dir,
179 "enc_hidden",
180 &enc_vec,
181 &[enc_len as i32, enc_dim as i32],
182 );
183
184 let codes: Vec<u32> = std::env::var_os("MAOLAN_ACESTEP_CODES")
187 .map(|raw| {
188 raw.to_string_lossy()
189 .split(',')
190 .filter_map(|part| part.trim().parse::<u32>().ok())
191 .collect()
192 })
193 .unwrap_or_else(|| ORACLE_CODES.to_vec());
194 let n_codes = codes.len();
195 let codes_tensor =
196 Tensor::<B, 2, Int>::from_data(TensorData::new(codes, [1, n_codes]), &device);
197 let hints = condition.codes_to_hints(codes_tensor);
198 let [_, hint_len, _] = hints.dims();
199 let hints_vec = vec_of(hints.clone());
200 dump_tensor(&out_dir, "detok_output", &hints_vec, &[hint_len as i32, 64]);
201
202 let noise = if let Some(path) = noise_path {
204 let (shape, data) = load_bin(&path);
205 let frames = shape[0] as usize;
206 Tensor::<B, 3>::from_data(TensorData::new(data, [1, frames, 64]), &device)
207 } else {
208 maolan_generate::acestep::pipeline::seeded_latent_noise(0, hint_len, 64, &device)
209 };
210 let [_, frames, _] = noise.dims();
211
212 let src = if hint_len >= frames {
214 hints.narrow(1, 0, frames)
215 } else {
216 let padding = silence.slice(frames - hint_len);
217 Tensor::cat(vec![hints, padding], 1)
218 };
219 let chunk_mask = Tensor::ones([1, frames, 64], &device);
220 let context = Tensor::cat(vec![src.clone(), chunk_mask], 2);
221 let context = oracle_context.unwrap_or(context);
222 let context_vec = vec_of(context.clone());
223 dump_tensor(&out_dir, "context", &context_vec, &[frames as i32, 128]);
224
225 let noise_vec = vec_of(noise.clone());
226 dump_tensor(&out_dir, "noise", &noise_vec, &[frames as i32, 64]);
227
228 let timesteps: &[f32] = if dit_config.is_turbo {
230 &TURBO_TIMESTEPS
231 } else {
232 &SFT_TIMESTEPS
233 };
234 let kv = dit.prepare_cross_kv(enc);
235 let mut xt = noise;
236 let total = timesteps.len();
237 for (index, &t_cur) in timesteps.iter().enumerate() {
238 let xt_vec = vec_of(xt.clone());
239 dump_tensor(
240 &out_dir,
241 &format!("dit_step{index}_xt"),
242 &xt_vec,
243 &[frames as i32, 64],
244 );
245 let v = dit.forward_with_kv(xt.clone(), t_cur, context.clone(), &kv);
246 let v_vec = vec_of(v.clone());
247 dump_tensor(
248 &out_dir,
249 &format!("dit_step{index}_vt"),
250 &v_vec,
251 &[frames as i32, 64],
252 );
253 let dt = if index + 1 == total {
254 t_cur
255 } else {
256 t_cur - timesteps[index + 1]
257 };
258 xt = xt - v * dt;
259 }
260 let x0_vec = vec_of(xt.clone());
261 dump_tensor(&out_dir, "dit_x0", &x0_vec, &[frames as i32, 64]);
262
263 let audio = vae.decode(xt);
265 let [_, channels, samples] = audio.dims();
266 let audio_vec = vec_of(audio);
267 dump_tensor(
268 &out_dir,
269 "vae_audio",
270 &audio_vec,
271 &[channels as i32, samples as i32],
272 );
273 println!("dumps written to {}", out_dir.display());
274 Ok(())
275}