Skip to main content

dit_dump/
dit_dump.rs

1//! Dump our pipeline intermediates in acestep.cpp's debug format for direct
2//! numerical comparison against `ace-synth --dump` output.
3//!
4//! Usage:
5//!   cargo run --release --example dit_dump -- <model_dir> <out_dir> [noise.bin]
6//!
7//! Uses the oracle's 19 codes and (optionally) the oracle's noise.bin so that
8//! dit_step*_vt/xt can be compared exactly.
9
10use 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    // ---- text + lyric encoding (causal, official prompts) ----
114    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); // explicit EOS (official add_eos=true)
119    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); // explicit EOS (official add_eos=true)
138    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    // Official: lyric branch is a raw embed_tokens lookup (no transformer).
142    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    // ---- conditioning: use the oracle's enc_hidden/context when provided ----
153    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    // ---- FSQ hints from the given codes (env MAOLAN_ACESTEP_CODES or the
185    // built-in oracle sequence) ----
186    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    // ---- noise ----
203    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    // src latents: hints cropped or silence-padded to the noise's frame count
213    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    // ---- DiT loop with per-step dumps (schedule from is_turbo) ----
229    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    // ---- VAE decode ----
264    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}