Skip to main content

acestep_phases/
acestep_phases.rs

1//! Phased ACE-Step diagnostic: runs one generation and dumps/analyzes every
2//! phase artifact, writing three WAVs:
3//!
4//! - `<out_prefix>_hints.wav` — LM hint latents decoded by the VAE, DiT
5//!   bypassed. Music here means LM + FSQ + detokenizer + VAE all work and any
6//!   remaining problem lives in the DiT.
7//! - `<out_prefix>_final.wav` — the normal full-pipeline output.
8//! - `<out_prefix>_silence.wav` — silence latent decoded by the VAE (should be
9//!   near-zero).
10//!
11//! Usage:
12//!   cargo run --release --example acestep_phases -- \
13//!       <model_dir> <out_prefix> [caption] [bpm] [key_scale] [time_sig] [length_ms]
14
15use maolan_generate::acestep::{
16    AceStepModelPaths, AceStepPipeline, AceStepTrace, GenerateMetadata, SilenceLatent,
17};
18use maolan_generate::heartcodec::write_wav_from_f32_interleaved;
19use std::path::Path;
20
21type B = burn::backend::Wgpu<f32, i64, u32>;
22
23fn main() -> anyhow::Result<()> {
24    let mut args = std::env::args().skip(1);
25    let model_dir = args
26        .next()
27        .unwrap_or_else(|| "/home/meka/repos/ace".to_string());
28    let out_prefix = args
29        .next()
30        .unwrap_or_else(|| "/tmp/acestep_phase".to_string());
31    let caption = args
32        .next()
33        .unwrap_or_else(|| "Metal guitar with a lot of distortion".to_string());
34    let bpm: Option<f32> = args.next().and_then(|v| v.parse().ok()).or(Some(120.0));
35    let key_scale = args.next().or_else(|| Some("A minor".to_string()));
36    let time_signature = args.next().or_else(|| Some("4/4".to_string()));
37    let length_ms: usize = args.next().and_then(|v| v.parse().ok()).unwrap_or(4000);
38
39    let device = burn::backend::wgpu::WgpuDevice::default();
40    burn::backend::wgpu::init_setup::<burn::backend::wgpu::graphics::Vulkan>(
41        &device,
42        burn::backend::wgpu::RuntimeOptions {
43            memory_config: burn::backend::wgpu::MemoryConfiguration::ExclusivePages,
44            ..Default::default()
45        },
46    );
47
48    let variant = if std::env::var("MAOLAN_ACESTEP_VARIANT")
49        .map(|v| v == "sft")
50        .unwrap_or(false)
51    {
52        maolan_generate::acestep::AceStepVariant::Sft
53    } else {
54        maolan_generate::acestep::AceStepVariant::Turbo
55    };
56    let paths = AceStepModelPaths::resolve(Path::new(&model_dir), variant)?;
57    let mut progress = |phase: &str, p: f32, op: &str| {
58        eprintln!("[{phase}] {:.0}% {op}", p * 100.0);
59    };
60    let pipeline = AceStepPipeline::<B>::load(&paths, &device, &mut progress)?;
61
62    let metadata = GenerateMetadata {
63        bpm,
64        key_scale: key_scale.as_deref(),
65        time_signature: time_signature.as_deref(),
66    };
67    let mut trace = AceStepTrace::default();
68    let (audio, meta) =
69        pipeline.generate_traced(&caption, &metadata, length_ms, 0, &mut progress, &mut trace)?;
70
71    // ---- Phase dumps ----
72    println!(
73        "\n===== TEXT PROMPT (caption branch) =====\n{}",
74        trace.text_prompt
75    );
76    println!("===== LYRIC PROMPT =====\n{}", trace.lyric_prompt);
77    println!("===== LM CoT BLOCK =====\n{}", trace.cot_block);
78    println!("===== LM CODES ({} total) =====", trace.codes.len());
79    println!("all codes: {:?}", trace.codes);
80    let mut sorted = trace.codes.clone();
81    sorted.sort_unstable();
82    sorted.dedup();
83    println!(
84        "unique: {}, min: {}, max: {}",
85        sorted.len(),
86        sorted.first().unwrap_or(&0),
87        sorted.last().unwrap_or(&0)
88    );
89    println!(
90        "\n===== CONDITIONING =====\nenc mean {:.6}  enc std {:.6}",
91        trace.enc_mean, trace.enc_std
92    );
93    println!("hints latent rms: {:.6}", trace.hints_latent_rms);
94    println!("final latent rms: {:.6}", trace.final_latent_rms);
95    println!("\n===== DIT STEPS =====\nstep  t        xt_rms   v_rms");
96    for (i, step) in trace.dit_steps.iter().enumerate() {
97        println!(
98            "{i:>4}  {:.4}   {:.4}   {:.4}",
99            step.t, step.xt_rms, step.v_rms
100        );
101    }
102
103    // ---- WAVs ----
104    if let Some((interleaved, channels, frames)) = &trace.hints_audio {
105        let path = format!("{out_prefix}_hints.wav");
106        write_wav_from_f32_interleaved(
107            interleaved,
108            *channels,
109            *frames,
110            meta.sample_rate_hz,
111            Path::new(&path),
112        )?;
113        println!("\nwrote {path} (VAE decode of LM hints, DiT bypassed)");
114    }
115
116    let [_, channels, frames] = audio.dims();
117    let channel_major: Vec<f32> = audio
118        .into_data()
119        .convert::<f32>()
120        .to_vec()
121        .map_err(|e| anyhow::anyhow!("{e}"))?;
122    let mut interleaved = vec![0.0_f32; channel_major.len()];
123    for (ch, samples) in channel_major.chunks_exact(frames).enumerate() {
124        for (frame, sample) in samples.iter().enumerate() {
125            interleaved[frame * channels + ch] = *sample;
126        }
127    }
128    let path = format!("{out_prefix}_final.wav");
129    write_wav_from_f32_interleaved(
130        &interleaved,
131        channels,
132        frames,
133        meta.sample_rate_hz,
134        Path::new(&path),
135    )?;
136    println!("wrote {path} (full pipeline)");
137
138    // Silence decode reference.
139    let silence = SilenceLatent::<B>::from_burnpack(&paths.silence_latent_bpk, &device)?;
140    let silence_audio = pipeline.vae.decode(silence.slice(frames / 1920));
141    let [_, sch, sframes] = silence_audio.dims();
142    let smaj: Vec<f32> = silence_audio
143        .into_data()
144        .convert::<f32>()
145        .to_vec()
146        .map_err(|e| anyhow::anyhow!("{e}"))?;
147    let mut sinter = vec![0.0_f32; smaj.len()];
148    for (ch, samples) in smaj.chunks_exact(sframes).enumerate() {
149        for (frame, sample) in samples.iter().enumerate() {
150            sinter[frame * sch + ch] = *sample;
151        }
152    }
153    let path = format!("{out_prefix}_silence.wav");
154    write_wav_from_f32_interleaved(&sinter, sch, sframes, meta.sample_rate_hz, Path::new(&path))?;
155    println!("wrote {path} (VAE decode of silence latent)");
156
157    Ok(())
158}