acestep_phases/
acestep_phases.rs1use 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 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 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 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}