Skip to main content

mimo_audio_dump/
mimo_audio_dump.rs

1//! MiMo audio-tower dump for the parity gates (tools/mimo_audio_ref.py cmp).
2//!
3//! ```text
4//! mimo_audio_dump decode   --wav F --out F.npy
5//! mimo_audio_dump frontend --wav F --out DIR
6//!     dec.npy [C,N], chan24k.npy [C,N'], wave24k.npy, mel.npy [M,128]
7//! mimo_audio_dump tower --src (HF_DIR | X.cmf) --out DIR (--wav F | --mel M.npy) [--codes C.npy]
8//!     feats.npy, codes_exact.npy, codes_bf16books.npy, embeds.npy (from --codes,
9//!     else from codes_bf16books), embeds_own.npy, tower.json (timings)
10//! mimo_audio_dump calib --src X.cmf --wav-dir D --out H.bin
11//!     GPTQ input Hessians of every tower linear over the clips in D, in the
12//!     `cortiq quantize-gptq --codec q4tp --hessians H.bin` cache format
13//! ```
14//!
15//! The array names match the oracle's, so `mimo_audio_ref.py cmp --ref R
16//! --eng E` lines them up.
17
18use cortiq_engine::mimo_audio::{self, MimoAudio};
19use std::path::{Path, PathBuf};
20use std::sync::Arc;
21use std::time::Instant;
22
23fn npy_write(path: &Path, descr: &str, shape: &[usize], bytes: &[u8]) {
24    let shape_s = match shape.len() {
25        1 => format!("({},)", shape[0]),
26        _ => format!(
27            "({})",
28            shape
29                .iter()
30                .map(|d| d.to_string())
31                .collect::<Vec<_>>()
32                .join(", ")
33        ),
34    };
35    let mut hdr = format!("{{'descr': '{descr}', 'fortran_order': False, 'shape': {shape_s}, }}");
36    let total = 10 + hdr.len() + 1;
37    hdr.push_str(&" ".repeat((64 - total % 64) % 64));
38    hdr.push('\n');
39    let mut out = Vec::with_capacity(10 + hdr.len() + bytes.len());
40    out.extend_from_slice(b"\x93NUMPY\x01\x00");
41    out.extend_from_slice(&(hdr.len() as u16).to_le_bytes());
42    out.extend_from_slice(hdr.as_bytes());
43    out.extend_from_slice(bytes);
44    std::fs::write(path, out).unwrap_or_else(|e| panic!("{}: {e}", path.display()));
45}
46
47fn save_f32(path: &Path, shape: &[usize], v: &[f32]) {
48    assert_eq!(
49        shape.iter().product::<usize>(),
50        v.len(),
51        "{}",
52        path.display()
53    );
54    let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
55    npy_write(path, "<f4", shape, &bytes);
56}
57
58fn save_i32(path: &Path, shape: &[usize], v: &[u32]) {
59    assert_eq!(shape.iter().product::<usize>(), v.len());
60    let bytes: Vec<u8> = v.iter().flat_map(|x| (*x as i32).to_le_bytes()).collect();
61    npy_write(path, "<i4", shape, &bytes);
62}
63
64/// `(descr, shape, payload)` of a C-order .npy.
65fn npy_read(path: &Path) -> (String, Vec<usize>, Vec<u8>) {
66    let b = std::fs::read(path).unwrap_or_else(|e| panic!("{}: {e}", path.display()));
67    assert_eq!(&b[..6], b"\x93NUMPY", "{}: not .npy", path.display());
68    let (hlen, start) = if b[6] == 1 {
69        (u16::from_le_bytes([b[8], b[9]]) as usize, 10)
70    } else {
71        (u32::from_le_bytes([b[8], b[9], b[10], b[11]]) as usize, 12)
72    };
73    let hdr = std::str::from_utf8(&b[start..start + hlen]).unwrap();
74    assert!(!hdr.contains("'fortran_order': True"), "fortran order");
75    let descr = hdr
76        .split("'descr':")
77        .nth(1)
78        .unwrap()
79        .split('\'')
80        .nth(1)
81        .unwrap()
82        .to_string();
83    let shape_s = hdr.split("'shape':").nth(1).unwrap();
84    let shape_s = &shape_s[shape_s.find('(').unwrap() + 1..shape_s.find(')').unwrap()];
85    let shape = shape_s
86        .split(',')
87        .map(str::trim)
88        .filter(|s| !s.is_empty())
89        .map(|s| s.parse().unwrap())
90        .collect();
91    (descr, shape, b[start + hlen..].to_vec())
92}
93
94fn load_f32(path: &Path) -> (Vec<usize>, Vec<f32>) {
95    let (d, shape, p) = npy_read(path);
96    assert_eq!(d, "<f4", "{}", path.display());
97    (
98        shape,
99        p.chunks_exact(4)
100            .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
101            .collect(),
102    )
103}
104
105fn load_codes(path: &Path) -> (Vec<usize>, Vec<u32>) {
106    let (d, shape, p) = npy_read(path);
107    let v = match d.as_str() {
108        "<i4" => p
109            .chunks_exact(4)
110            .map(|c| i32::from_le_bytes([c[0], c[1], c[2], c[3]]) as u32)
111            .collect(),
112        "<i8" => p
113            .chunks_exact(8)
114            .map(|c| i64::from_le_bytes(c.try_into().unwrap()) as u32)
115            .collect(),
116        other => panic!("{}: codes dtype {other}", path.display()),
117    };
118    (shape, v)
119}
120
121/// The `cortiq quantize-gptq --hessians` cache (`CMFHESS1`, the layout of
122/// cortiq-cli's `save_hessians`): identical Hessians are stored once under
123/// all their names, each as its upper triangle.
124fn save_hessians(
125    path: &Path,
126    hess: &std::collections::HashMap<String, cortiq_engine::gptq_capture::HessianAcc>,
127) {
128    use std::io::Write;
129    let mut names: Vec<&String> = hess.keys().collect();
130    names.sort();
131    let mut uniq: Vec<(Vec<&String>, &cortiq_engine::gptq_capture::HessianAcc)> = Vec::new();
132    for n in names {
133        let a = &hess[n];
134        if let Some(u) = uniq.iter_mut().find(|(_, b)| {
135            b.cols == a.cols && b.count == a.count && b.sumsq == a.sumsq && b.h == a.h
136        }) {
137            u.0.push(n);
138        } else {
139            uniq.push((vec![n], a));
140        }
141    }
142    let mut f = std::io::BufWriter::with_capacity(1 << 22, std::fs::File::create(path).unwrap());
143    f.write_all(b"CMFHESS1").unwrap();
144    f.write_all(&(uniq.len() as u64).to_le_bytes()).unwrap();
145    for (ns, a) in &uniq {
146        f.write_all(&(ns.len() as u32).to_le_bytes()).unwrap();
147        for n in ns {
148            f.write_all(&(n.len() as u32).to_le_bytes()).unwrap();
149            f.write_all(n.as_bytes()).unwrap();
150        }
151        f.write_all(&(a.cols as u64).to_le_bytes()).unwrap();
152        f.write_all(&(a.count as u64).to_le_bytes()).unwrap();
153        f.write_all(&(a.h.len() as u64).to_le_bytes()).unwrap();
154        for v in &a.sumsq {
155            f.write_all(&v.to_le_bytes()).unwrap();
156        }
157        let n = a.cols;
158        if a.h.len() == n * n {
159            for i in 0..n {
160                for v in &a.h[i * n + i..i * n + n] {
161                    f.write_all(&v.to_le_bytes()).unwrap();
162                }
163            }
164        }
165    }
166    f.flush().unwrap();
167}
168
169fn arg(args: &[String], name: &str) -> Option<String> {
170    args.iter()
171        .position(|a| a == name)
172        .and_then(|i| args.get(i + 1).cloned())
173}
174
175fn main() {
176    let args: Vec<String> = std::env::args().collect();
177    let cmd = args.get(1).map(String::as_str).unwrap_or("");
178    let out = PathBuf::from(arg(&args, "--out").expect("--out"));
179    match cmd {
180        "decode" => {
181            let wav = std::fs::read(arg(&args, "--wav").expect("--wav")).unwrap();
182            let w = mimo_audio::decode_wav(&wav).unwrap();
183            let flat: Vec<f32> = w.channels.concat();
184            save_f32(&out, &[w.channels.len(), w.frames()], &flat);
185            println!(
186                "rate {} channels {} frames {}",
187                w.sample_rate,
188                w.channels.len(),
189                w.frames()
190            );
191        }
192        "frontend" => {
193            std::fs::create_dir_all(&out).unwrap();
194            let wav = std::fs::read(arg(&args, "--wav").expect("--wav")).unwrap();
195            let t0 = Instant::now();
196            let w = mimo_audio::decode_wav(&wav).unwrap();
197            save_f32(
198                &out.join("dec.npy"),
199                &[w.channels.len(), w.frames()],
200                &w.channels.concat(),
201            );
202            let chans: Vec<Vec<f32>> = w
203                .channels
204                .iter()
205                .map(|c| mimo_audio::resample_sinc(c, w.sample_rate, mimo_audio::SAMPLE_RATE))
206                .collect();
207            save_f32(
208                &out.join("chan24k.npy"),
209                &[chans.len(), chans[0].len()],
210                &chans.concat(),
211            );
212            let mono = mimo_audio::wav_to_mono_24k(&w).unwrap();
213            save_f32(&out.join("wave24k.npy"), &[mono.len()], &mono);
214            let pool = cortiq_engine::pool::Pool::from_env();
215            let (mel, m) = mimo_audio::log_mel(&mono, pool.as_deref()).unwrap();
216            save_f32(&out.join("mel.npy"), &[m, mimo_audio::N_MELS], &mel);
217            println!(
218                "rate {} channels {} frames {} -> {} samples, {m} mel frames, K {} ({:.3}s)",
219                w.sample_rate,
220                w.channels.len(),
221                w.frames(),
222                mono.len(),
223                mimo_audio::audio_token_count(m, 4),
224                t0.elapsed().as_secs_f64()
225            );
226        }
227        "tower" => {
228            std::fs::create_dir_all(&out).unwrap();
229            let src = PathBuf::from(arg(&args, "--src").expect("--src"));
230            let t0 = Instant::now();
231            let audio = if src.extension().is_some_and(|e| e == "cmf") {
232                let model = Arc::new(cortiq_core::CmfModel::open(&src).expect("open cmf"));
233                MimoAudio::from_model(&model).expect("load towers")
234            } else {
235                MimoAudio::from_hf_dir(&src).expect("load towers")
236            };
237            let t_load = t0.elapsed().as_secs_f64();
238            let (mel, m) = if let Some(mp) = arg(&args, "--mel") {
239                let (shape, v) = load_f32(Path::new(&mp));
240                assert_eq!(shape[1], mimo_audio::N_MELS);
241                (v, shape[0])
242            } else {
243                let wav = std::fs::read(arg(&args, "--wav").expect("--wav or --mel")).unwrap();
244                audio.wav_to_mel(&wav).unwrap()
245            };
246            let t1 = Instant::now();
247            let feats = audio.features(&mel, m).unwrap();
248            let t_feats = t1.elapsed().as_secs_f64();
249            let d = audio.tokenizer.cfg.d_model;
250            let rows = feats.len() / d;
251            save_f32(&out.join("feats.npy"), &[rows, d], &feats);
252            let t2 = Instant::now();
253            let exact = audio.tokenizer.quantize(&feats, rows, false, audio.pool());
254            let t_rvq = t2.elapsed().as_secs_f64();
255            let rounded = audio.tokenizer.quantize(&feats, rows, true, audio.pool());
256            let levels = exact.len() / rows;
257            save_i32(&out.join("codes_exact.npy"), &[rows, levels], &exact);
258            save_i32(&out.join("codes_bf16books.npy"), &[rows, levels], &rounded);
259            let own = mimo_audio::AudioCodes {
260                frames: rows,
261                levels,
262                codes: if audio.bf16_codebooks {
263                    rounded.clone()
264                } else {
265                    exact.clone()
266                },
267            };
268            let t3 = Instant::now();
269            let emb_own = audio.embed_codes(&own).unwrap();
270            let t_enc = t3.elapsed().as_secs_f64();
271            save_f32(
272                &out.join("embeds_own.npy"),
273                &[emb_own.n_tokens, emb_own.dim],
274                &emb_own.rows,
275            );
276            let fixed = match arg(&args, "--codes") {
277                Some(cp) => {
278                    let (shape, v) = load_codes(Path::new(&cp));
279                    mimo_audio::AudioCodes {
280                        frames: shape[0],
281                        levels: shape[1],
282                        codes: v,
283                    }
284                }
285                None => own.clone(),
286            };
287            let emb = audio.embed_codes(&fixed).unwrap();
288            save_f32(&out.join("embeds.npy"), &[emb.n_tokens, emb.dim], &emb.rows);
289            let k = mimo_audio::audio_token_count(m, audio.encoder.cfg.group);
290            let meta = serde_json::json!({
291                "src": src.display().to_string(),
292                "mel_frames": m,
293                "segments": mimo_audio::segment_lengths(m),
294                "codes": rows,
295                "placeholder_count_K": k,
296                "embed_rows_own": emb_own.n_tokens,
297                "embed_rows_fixed": emb.n_tokens,
298                "load_s": t_load,
299                "features_s": t_feats,
300                "rvq_s": t_rvq,
301                "encoder_s": t_enc,
302                "bf16_codebooks_default": audio.bf16_codebooks,
303                "threads": cortiq_engine::pool::Pool::effective_threads(),
304            });
305            std::fs::write(
306                out.join("tower.json"),
307                serde_json::to_string_pretty(&meta).unwrap(),
308            )
309            .unwrap();
310            println!("{meta}");
311            assert_eq!(emb_own.n_tokens, k, "placeholder count != encoder rows");
312        }
313        "calib" => {
314            let src = PathBuf::from(arg(&args, "--src").expect("--src"));
315            let model = Arc::new(cortiq_core::CmfModel::open(&src).expect("open cmf"));
316            let audio = MimoAudio::from_model(&model).expect("load towers");
317            let dir = PathBuf::from(arg(&args, "--wav-dir").expect("--wav-dir"));
318            let mut wavs: Vec<PathBuf> = std::fs::read_dir(&dir)
319                .unwrap()
320                .filter_map(|e| e.ok().map(|e| e.path()))
321                .filter(|p| p.extension().is_some_and(|e| e == "wav"))
322                .collect();
323            wavs.sort();
324            let t0 = Instant::now();
325            cortiq_engine::gptq_capture::begin(true);
326            let mut frames = 0usize;
327            for w in &wavs {
328                let emb = audio.embed_wav(&std::fs::read(w).unwrap()).unwrap();
329                frames += emb.n_tokens;
330                eprintln!(
331                    "  {} -> {} rows ({:.0}s)",
332                    w.display(),
333                    emb.n_tokens,
334                    t0.elapsed().as_secs_f64()
335                );
336            }
337            let hess = cortiq_engine::gptq_capture::end();
338            save_hessians(&out, &hess);
339            println!(
340                "{} clips, {frames} LLM rows, {} linears -> {} ({:.0}s)",
341                wavs.len(),
342                hess.len(),
343                out.display(),
344                t0.elapsed().as_secs_f64()
345            );
346        }
347        _ => {
348            eprintln!(
349                "usage: mimo_audio_dump (decode|frontend|tower|calib) --out ... (see the source header)"
350            );
351            std::process::exit(2);
352        }
353    }
354}