Skip to main content

fsq_probe/
fsq_probe.rs

1//! Probe: feed controlled FSQ code sequences through codes_to_hints + VAE to
2//! test whether the FSQ/detokenizer chain produces lively latents.
3//!
4//! Usage: cargo run --release --example fsq_probe -- <model_dir>
5
6use burn::tensor::{Int, Tensor, TensorData};
7use maolan_generate::acestep::condition::AceStepCondition;
8use maolan_generate::acestep::config::AceStepConfig;
9use maolan_generate::acestep::pipeline::SilenceLatent;
10use maolan_generate::acestep::vae::{OobleckDecoder, OobleckVaeConfig};
11use std::path::Path;
12
13type B = burn::backend::NdArray<f32>;
14
15fn rms(values: &[f32]) -> f32 {
16    (values.iter().map(|v| v * v).sum::<f32>() / values.len() as f32).sqrt()
17}
18
19fn run_case(
20    name: &str,
21    codes: &[u32],
22    condition: &AceStepCondition<B>,
23    vae: &OobleckDecoder<B>,
24    silence: &SilenceLatent<B>,
25    device: &burn::tensor::Device<B>,
26) -> anyhow::Result<()> {
27    let n = codes.len();
28    let tensor = Tensor::<B, 2, Int>::from_data(TensorData::new(codes.to_vec(), [1, n]), device);
29    let hints = condition.codes_to_hints(tensor);
30    let silence_frames = silence.slice(n * 5);
31    let diff = (hints.clone() - silence_frames)
32        .abs()
33        .into_data()
34        .convert::<f32>()
35        .to_vec::<f32>()
36        .map_err(|e| anyhow::anyhow!("{e}"))?;
37    let hint_values: Vec<f32> = hints
38        .clone()
39        .into_data()
40        .convert::<f32>()
41        .to_vec()
42        .map_err(|e| anyhow::anyhow!("{e}"))?;
43    let audio = vae.decode(hints);
44    let audio_values: Vec<f32> = audio
45        .into_data()
46        .convert::<f32>()
47        .to_vec()
48        .map_err(|e| anyhow::anyhow!("{e}"))?;
49    let peak = audio_values.iter().fold(0.0_f32, |a, v| a.max(v.abs()));
50    let mean_abs_diff = diff.iter().sum::<f32>() / diff.len() as f32;
51    println!(
52        "{name}: hints rms {:.4}  mean|hints-silence| {:.4}  audio rms {:.4}  audio peak {:.4}",
53        rms(&hint_values),
54        mean_abs_diff,
55        rms(&audio_values),
56        peak
57    );
58    Ok(())
59}
60
61fn main() -> anyhow::Result<()> {
62    let model_dir = std::env::args()
63        .nth(1)
64        .unwrap_or_else(|| "/home/meka/repos/ace".to_string());
65    let model_dir = Path::new(&model_dir);
66    let device = Default::default();
67
68    let dit_config = AceStepConfig::load(&model_dir.join("dit_config.json"))?;
69    let condition = AceStepCondition::<B>::from_burnpack(
70        &dit_config,
71        &model_dir.join("acestep-condition.bpk"),
72        &device,
73    )?;
74    let vae_config = OobleckVaeConfig::load(&model_dir.join("vae_config.json"))?;
75    let vae = OobleckDecoder::<B>::from_burnpack(
76        &vae_config,
77        &model_dir.join("acestep-vae.bpk"),
78        &device,
79    )?;
80    let silence =
81        SilenceLatent::<B>::from_burnpack(&model_dir.join("silence_latent.bpk"), &device)?;
82
83    // Splitmix64 random codes.
84    let mut state = 42_u64;
85    let mut rand_codes = Vec::new();
86    for _ in 0..20 {
87        state = state
88            .wrapping_mul(6364136223846793005)
89            .wrapping_add(1442695040888963407);
90        rand_codes.push(((state >> 33) % 64000) as u32);
91    }
92
93    let lm_codes: Vec<u32> = vec![
94        61890, 51649, 53753, 53753, 56314, 53754, 55802, 4538, 4538, 4538, 5050, 5050, 5050, 5050,
95        5050, 5050, 5050, 5050, 5050, 43513,
96    ];
97
98    // Structure check: for a few distinct codes, dump the detokenizer output
99    // frames and measure how much outputs differ between codes.
100    for code in [0u32, 12345, 53754] {
101        let tensor = Tensor::<B, 2, Int>::from_data(TensorData::new(vec![code], [1, 1]), &device);
102        let hints = condition.codes_to_hints(tensor); // [1, 5, 64]
103        let values: Vec<f32> = hints
104            .into_data()
105            .convert::<f32>()
106            .to_vec()
107            .map_err(|e| anyhow::anyhow!("{e}"))?;
108        let frame_means: Vec<f32> = values
109            .as_chunks::<64>()
110            .0
111            .iter()
112            .map(|f| f.iter().sum::<f32>() / 64.0)
113            .collect();
114        let frame_rms: Vec<f32> = values.as_chunks::<64>().0.iter().map(|f| rms(f)).collect();
115        println!(
116            "code {code:>6}: frame means {:?}",
117            frame_means
118                .iter()
119                .map(|v| format!("{v:.4}"))
120                .collect::<Vec<_>>()
121        );
122        println!(
123            "           frame rms   {:?}",
124            frame_rms
125                .iter()
126                .map(|v| format!("{v:.4}"))
127                .collect::<Vec<_>>()
128        );
129    }
130
131    // Time-variance vs channel-variance of hint latents for random codes.
132    let tensor =
133        Tensor::<B, 2, Int>::from_data(TensorData::new(rand_codes.clone(), [1, 20]), &device);
134    let hints = condition.codes_to_hints(tensor); // [1, 100, 64]
135    let values: Vec<f32> = hints
136        .into_data()
137        .convert::<f32>()
138        .to_vec()
139        .map_err(|e| anyhow::anyhow!("{e}"))?;
140    let (frames, _) = values.as_chunks::<64>();
141    // variance of per-frame means (time structure) vs mean per-frame variance (channel structure)
142    let frame_means: Vec<f32> = frames
143        .iter()
144        .map(|f| f.iter().sum::<f32>() / 64.0)
145        .collect();
146    let tm = frame_means.iter().sum::<f32>() / frame_means.len() as f32;
147    let time_var =
148        frame_means.iter().map(|v| (v - tm) * (v - tm)).sum::<f32>() / frame_means.len() as f32;
149    let chan_var = frames
150        .iter()
151        .map(|f| {
152            let m = f.iter().sum::<f32>() / 64.0;
153            f.iter().map(|v| (v - m) * (v - m)).sum::<f32>() / 64.0
154        })
155        .sum::<f32>()
156        / frames.len() as f32;
157    println!("random-code hints: time var {time_var:.6}  channel var {chan_var:.6}");
158
159    run_case(
160        "random codes ",
161        &rand_codes,
162        &condition,
163        &vae,
164        &silence,
165        &device,
166    )?;
167    run_case(
168        "LM codes     ",
169        &lm_codes,
170        &condition,
171        &vae,
172        &silence,
173        &device,
174    )?;
175    run_case(
176        "constant 0   ",
177        &[0u32; 20],
178        &condition,
179        &vae,
180        &silence,
181        &device,
182    )?;
183    Ok(())
184}