zimage_devcheck/
zimage_devcheck.rs1use cortiq_engine::zimage::{self, ZImageDit, ZShape};
4use std::collections::HashMap;
5use std::sync::Arc;
6
7fn read_st(path: &std::path::Path) -> (HashMap<String, Vec<f32>>, serde_json::Value) {
8 let b = std::fs::read(path).unwrap();
9 let n = u64::from_le_bytes(b[..8].try_into().unwrap()) as usize;
10 let h: serde_json::Value = serde_json::from_slice(&b[8..8 + n]).unwrap();
11 let base = 8 + n;
12 let mut out = HashMap::new();
13 let mut meta = serde_json::Value::Null;
14 for (k, v) in h.as_object().unwrap() {
15 if k == "__metadata__" {
16 meta = v.clone();
17 continue;
18 }
19 if v["dtype"].as_str() != Some("F32") {
20 continue;
21 }
22 let o = v["data_offsets"].as_array().unwrap();
23 let raw = &b[base + o[0].as_u64().unwrap() as usize..base + o[1].as_u64().unwrap() as usize];
24 out.insert(k.clone(), raw.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect());
25 }
26 (out, meta)
27}
28
29fn meta_usize(meta: &serde_json::Value, k: &str) -> usize {
30 meta[k].as_str().unwrap().trim_matches('"').parse().unwrap()
31}
32
33fn rel(a: &[f32], b: &[f32]) -> f64 {
34 let n = a.len().min(b.len());
35 let (mut d, mut r) = (0f64, 0f64);
36 for i in 0..n {
37 let (x, y) = (a[i] as f64, b[i] as f64);
38 d += (x - y) * (x - y);
39 r += y * y;
40 }
41 (d / r.max(1e-300)).sqrt()
42}
43
44fn main() {
45 unsafe { std::env::set_var("CMF_GPU", "1") };
46 let a: Vec<String> = std::env::args().collect();
47 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
48 let dit = ZImageDit::from_cmf(&model).unwrap();
49 let od = std::path::Path::new(&a[2]);
50 let case = &a[3];
51 let (o, meta) = read_st(&od.join(format!("dit_{case}_fp32.safetensors")));
52 let (hh, ww, l) = (meta_usize(&meta, "H"), meta_usize(&meta, "W"), meta_usize(&meta, "L"));
53 let t = o["t_model"][0];
54 let cap = &o["cap"];
55 let x_in = &o["x_in"];
56 let shape = ZShape::new(hh, ww, l);
57 let mods = dit.mods_for_steps(&[t]);
58 let fs = dit.final_scale_for_steps(&[t]);
59 let rope = zimage::ids_and_rope(shape.grid, l, dit.cfg.rope_theta, dit.cfg.axes_dims);
60 let prep = dit.prepare(cap, shape, 1, None).unwrap();
62 println!("device prepared: {}", prep.device);
63 let mut cap_cpu = dit.embed_caption(cap, l);
64 dit.refine_caption_cpu(&mut cap_cpu, (&rope.cap.0, &rope.cap.1));
65 println!("cr1_out dev vs cpu {:.3e} cpu vs oracle {:.3e} dev vs oracle {:.3e}",
66 rel(&prep.cap, &cap_cpu), rel(&cap_cpu, &o["cr1_out"]), rel(&prep.cap, &o["cr1_out"]));
67 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
68 let x_tok = zimage::pad_rows_repeat_last(&zimage::patchify(x_in, c, lh, lw), shape.n_img, shape.n_img_p, dit.geom().patch_dim);
69 let tapdir = std::env::var("CMF_ZI_TAPS").unwrap_or("/root/zb/taps".into());
70 let v_dev = dit.step(&prep, 0, &x_tok, &mods, &fs);
71 let mods2 = dit.mods_for_steps(&[t * 0.5 + 0.25]);
74 let fs2 = dit.final_scale_for_steps(&[t * 0.5 + 0.25]);
75 let _ = dit.step(&prep, 1, &x_tok, &mods2, &fs2);
76 let v_dev2 = dit.step(&prep, 2, &x_tok, &mods, &fs);
77 println!("replay dev2 vs dev1 {:.3e}", rel(&v_dev2, &v_dev));
78 let mut taps: Vec<(String, Vec<f32>)> = Vec::new();
79 let mut prep_cpu = prep;
81 prep_cpu.device = false;
82 let v_cpu = dit.step_cpu_taps(&prep_cpu, &x_tok, &mods, &fs, &mut |n, v| taps.push((n.to_string(), v.to_vec())));
83 println!("v dev vs cpu {:.3e} cpu vs oracle(final_out) {:.3e}", rel(&v_dev, &v_cpu), rel(&v_cpu, &o["final_out"]));
84 for (n, v) in &taps {
85 let p = std::path::Path::new(&tapdir).join("step0").join(format!("{n}.f32"));
86 if let Ok(b) = std::fs::read(&p) {
87 let d: Vec<f32> = b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect();
88 let or = o.get(n.as_str()).map(|w| format!("{:.3e}", rel(v, w))).unwrap_or_default();
89 println!("{n:10} dev vs cpu {:.3e} (cpu vs oracle {or})", rel(&d, v));
90 }
91 }
92}