Skip to main content

dit_dump_cpu/
dit_dump_cpu.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, 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::NdArray<f32>;
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 = Default::default();
71
72    let dit_config = AceStepConfig::load(&model_dir.join("dit_config.json"))?;
73    let text_config = Qwen3Config::load(&model_dir.join("qwen3_config.json"))?;
74    let vae_config = OobleckVaeConfig::load(&model_dir.join("vae_config.json"))?;
75    eprintln!("loading components...");
76    let text_encoder = Qwen3Model::<B>::from_burnpack(
77        &text_config,
78        &model_dir.join("qwen3-encoder.bpk"),
79        &device,
80    )?;
81    let condition = AceStepCondition::<B>::from_burnpack(
82        &dit_config,
83        &model_dir.join("acestep-condition.bpk"),
84        &device,
85    )?;
86    let dit =
87        AceStepDiT::<B>::from_burnpack(&dit_config, &model_dir.join("acestep-dit.bpk"), &device)?;
88    let vae = OobleckDecoder::<B>::from_burnpack(
89        &vae_config,
90        &model_dir.join("acestep-vae.bpk"),
91        &device,
92    )?;
93    let silence =
94        SilenceLatent::<B>::from_burnpack(&model_dir.join("silence_latent.bpk"), &device)?;
95    let tokenizer = tokie::Tokenizer::from_json(model_dir.join("tokenizer.json"))
96        .map_err(|e| anyhow::anyhow!("tokenizer: {e}"))?;
97
98    // ---- text + lyric encoding (causal, official prompts) ----
99    let metas = build_metas_block(Some(120.0), Some("A minor"), Some("4/4"), 4);
100    let text_prompt = build_dit_text_prompt("Metal guitar with a lot of distortion", &metas);
101    let ids = tokenizer.encode(&text_prompt, false).ids;
102    let mut ids: Vec<i64> = ids.into_iter().map(i64::from).collect();
103    ids.push(151643); // explicit EOS (official add_eos=true)
104    let n_text = ids.len();
105    let ids_t = Tensor::<B, 2, Int>::from_data(TensorData::new(ids, [1, n_text]), &device);
106    let text_hidden = text_encoder.forward(ids_t, true);
107    let text_hidden_vec = vec_of(text_hidden.clone());
108    dump_tensor(
109        &out_dir,
110        "text_hidden",
111        &text_hidden_vec,
112        &[n_text as i32, 1024],
113    );
114
115    let lyric_ids = tokenizer
116        .encode(
117            maolan_generate::acestep::pipeline::INSTRUMENTAL_LYRIC_PROMPT,
118            false,
119        )
120        .ids;
121    let mut lyric_ids: Vec<i64> = lyric_ids.into_iter().map(i64::from).collect();
122    lyric_ids.push(151643); // explicit EOS (official add_eos=true)
123    let n_lyric = lyric_ids.len();
124    let lyric_ids_t =
125        Tensor::<B, 2, Int>::from_data(TensorData::new(lyric_ids, [1, n_lyric]), &device);
126    // Official: lyric branch is a raw embed_tokens lookup (no transformer).
127    let lyric_hidden = text_encoder.embed_tokens.forward(lyric_ids_t);
128    let lyric_embed_vec = vec_of(lyric_hidden.clone());
129    dump_tensor(
130        &out_dir,
131        "lyric_embed",
132        &lyric_embed_vec,
133        &[n_lyric as i32, 1024],
134    );
135    let lyric_mask = Tensor::<B, 2, Int>::ones([1, n_lyric], &device);
136
137    // ---- conditioning: use the oracle's enc_hidden/context when provided ----
138    let mut oracle_context: Option<Tensor<B, 3>> = None;
139    let enc = if let Some(dir) = std::env::var_os("MAOLAN_ORACLE_DUMP_DIR") {
140        let dir = PathBuf::from(dir);
141        let (enc_shape, enc_data) = load_bin(&dir.join("enc_hidden.bin"));
142        eprintln!("oracle enc_hidden: {enc_shape:?}");
143        let (ctx_shape, ctx_data) = load_bin(&dir.join("context.bin"));
144        eprintln!("oracle context: {ctx_shape:?}");
145        let enc = Tensor::<B, 3>::from_data(
146            TensorData::new(enc_data, [1, enc_shape[0] as usize, enc_shape[1] as usize]),
147            &device,
148        );
149        let ctx = Tensor::<B, 3>::from_data(TensorData::new(ctx_data, [1, 96, 128]), &device);
150        oracle_context = Some(ctx);
151        enc
152    } else {
153        condition.encode(
154            text_hidden,
155            lyric_hidden,
156            lyric_mask,
157            silence.timbre_reference(),
158        )
159    };
160    let [_, enc_len, enc_dim] = enc.dims();
161    let enc_vec = vec_of(enc.clone());
162    dump_tensor(
163        &out_dir,
164        "enc_hidden",
165        &enc_vec,
166        &[enc_len as i32, enc_dim as i32],
167    );
168
169    // ---- FSQ hints from the given codes (env MAOLAN_ACESTEP_CODES or the
170    // built-in oracle sequence) ----
171    let codes: Vec<u32> = std::env::var_os("MAOLAN_ACESTEP_CODES")
172        .map(|raw| {
173            raw.to_string_lossy()
174                .split(',')
175                .filter_map(|part| part.trim().parse::<u32>().ok())
176                .collect()
177        })
178        .unwrap_or_else(|| ORACLE_CODES.to_vec());
179    let n_codes = codes.len();
180    let codes_tensor =
181        Tensor::<B, 2, Int>::from_data(TensorData::new(codes, [1, n_codes]), &device);
182    let hints = condition.codes_to_hints(codes_tensor);
183    let [_, hint_len, _] = hints.dims();
184    let hints_vec = vec_of(hints.clone());
185    dump_tensor(&out_dir, "detok_output", &hints_vec, &[hint_len as i32, 64]);
186
187    // ---- noise ----
188    let noise = if let Some(path) = noise_path {
189        let (shape, data) = load_bin(&path);
190        let frames = shape[0] as usize;
191        Tensor::<B, 3>::from_data(TensorData::new(data, [1, frames, 64]), &device)
192    } else {
193        maolan_generate::acestep::pipeline::seeded_latent_noise(0, hint_len, 64, &device)
194    };
195    let [_, frames, _] = noise.dims();
196
197    // src latents: hints cropped or silence-padded to the noise's frame count
198    let src = if hint_len >= frames {
199        hints.narrow(1, 0, frames)
200    } else {
201        let padding = silence.slice(frames - hint_len);
202        Tensor::cat(vec![hints, padding], 1)
203    };
204    let chunk_mask = Tensor::ones([1, frames, 64], &device);
205    let context = Tensor::cat(vec![src.clone(), chunk_mask], 2);
206    let context = oracle_context.unwrap_or(context);
207    let context_vec = vec_of(context.clone());
208    dump_tensor(&out_dir, "context", &context_vec, &[frames as i32, 128]);
209
210    let noise_vec = vec_of(noise.clone());
211    dump_tensor(&out_dir, "noise", &noise_vec, &[frames as i32, 64]);
212
213    // ---- DiT loop with per-step dumps ----
214    let kv = dit.prepare_cross_kv(enc);
215    let mut xt = noise;
216    let total = TURBO_TIMESTEPS.len();
217    for (index, &t_cur) in TURBO_TIMESTEPS.iter().enumerate() {
218        let xt_vec = vec_of(xt.clone());
219        dump_tensor(
220            &out_dir,
221            &format!("dit_step{index}_xt"),
222            &xt_vec,
223            &[frames as i32, 64],
224        );
225        let v = dit.forward_with_kv(xt.clone(), t_cur, context.clone(), &kv);
226        let v_vec = vec_of(v.clone());
227        dump_tensor(
228            &out_dir,
229            &format!("dit_step{index}_vt"),
230            &v_vec,
231            &[frames as i32, 64],
232        );
233        let dt = if index + 1 == total {
234            t_cur
235        } else {
236            t_cur - TURBO_TIMESTEPS[index + 1]
237        };
238        xt = xt - v * dt;
239    }
240    let x0_vec = vec_of(xt.clone());
241    dump_tensor(&out_dir, "dit_x0", &x0_vec, &[frames as i32, 64]);
242
243    // ---- VAE decode ----
244    let audio = vae.decode(xt);
245    let [_, channels, samples] = audio.dims();
246    let audio_vec = vec_of(audio);
247    dump_tensor(
248        &out_dir,
249        "vae_audio",
250        &audio_vec,
251        &[channels as i32, samples as i32],
252    );
253    println!("dumps written to {}", out_dir.display());
254    Ok(())
255}