1use 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
64fn 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
121fn 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}