Skip to main content

zimage_metal_check/
zimage_metal_check.rs

1//! One Z-Image DiT forward on the device against the CPU path on the same
2//! container and against the diffusers fp32 oracle (plan WP3 parity gate).
3//!
4//! zimage_metal_check <container.cmf> <oracle dir> <case, e.g. r512_p0_t8_i0> [reps] [neg case]
5//!
6//! Prints: refined caption dev vs CPU; `v` device vs CPU (same caption),
7//! device vs oracle fp32, CPU vs oracle fp32, oracle bf16 vs fp32 (the
8//! diffusers-bf16 floor, when the bf16 dump is present); the in-process
9//! step times. `ZC_NEG=1` with a CFG oracle case: the batch-2 pair
10//! (positive = the oracle caption, negative = the same caption) against
11//! two single forwards. `ZC_NOCPU=1` skips the CPU forward.
12use cortiq_engine::zimage::{self, ZImageDit, ZShape};
13use std::collections::HashMap;
14use std::sync::Arc;
15
16fn read_st(path: &std::path::Path) -> (HashMap<String, Vec<f32>>, serde_json::Value) {
17    let b = std::fs::read(path).unwrap_or_else(|e| panic!("{}: {e}", path.display()));
18    let n = u64::from_le_bytes(b[..8].try_into().unwrap()) as usize;
19    let h: serde_json::Value = serde_json::from_slice(&b[8..8 + n]).unwrap();
20    let base = 8 + n;
21    let mut out = HashMap::new();
22    let mut meta = serde_json::Value::Null;
23    for (k, v) in h.as_object().unwrap() {
24        if k == "__metadata__" {
25            meta = v.clone();
26            continue;
27        }
28        let o = v["data_offsets"].as_array().unwrap();
29        let raw = &b[base + o[0].as_u64().unwrap() as usize..base + o[1].as_u64().unwrap() as usize];
30        let data: Vec<f32> = match v["dtype"].as_str() {
31            Some("F32") => raw.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect(),
32            Some("BF16") => raw
33                .chunks_exact(2)
34                .map(|c| f32::from_bits((u16::from_le_bytes([c[0], c[1]]) as u32) << 16))
35                .collect(),
36            _ => continue,
37        };
38        out.insert(k.clone(), data);
39    }
40    (out, meta)
41}
42
43fn meta_usize(meta: &serde_json::Value, k: &str) -> usize {
44    meta[k].as_str().unwrap().trim_matches('"').parse().unwrap()
45}
46
47fn rel(a: &[f32], b: &[f32]) -> f64 {
48    let n = a.len().min(b.len());
49    let (mut d, mut r) = (0f64, 0f64);
50    for i in 0..n {
51        let (x, y) = (a[i] as f64, b[i] as f64);
52        d += (x - y) * (x - y);
53        r += y * y;
54    }
55    (d / r.max(1e-300)).sqrt()
56}
57
58fn main() {
59    let a: Vec<String> = std::env::args().collect();
60    let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
61    let dit = ZImageDit::from_cmf(&model).unwrap();
62    let od = std::path::Path::new(&a[2]);
63    let case = &a[3];
64    let reps: usize = a.get(4).and_then(|v| v.parse().ok()).unwrap_or(2);
65    let (o, meta) = read_st(&od.join(format!("dit_{case}_fp32.safetensors")));
66    let bf = od.join(format!("dit_{case}_bf16.safetensors"));
67    let (hh, ww, l) = (meta_usize(&meta, "H"), meta_usize(&meta, "W"), meta_usize(&meta, "L"));
68    let t = o["t_model"][0];
69    let cap = &o["cap"];
70    let x_in = &o["x_in"];
71    let shape = ZShape::new(hh, ww, l);
72    let mods = dit.mods_for_steps(&[t]);
73    let fs = dit.final_scale_for_steps(&[t]);
74    let rope = zimage::ids_and_rope(shape.grid, l, dit.cfg.rope_theta, dit.cfg.axes_dims);
75    let t0 = std::time::Instant::now();
76    let prep = dit.prepare(cap, shape, 1, None).unwrap();
77    println!("device prepared: {} ({:.3}s incl. caption refine)", prep.device, t0.elapsed().as_secs_f64());
78    let mut cap_cpu = dit.embed_caption(cap, l);
79    dit.refine_caption_cpu(&mut cap_cpu, (&rope.cap.0, &rope.cap.1));
80    if let Some(cr) = o.get("cr1_out") {
81        println!(
82            "caption  dev vs cpu {:.3e}   cpu vs oracle {:.3e}   dev vs oracle {:.3e}",
83            rel(&prep.cap, &cap_cpu),
84            rel(&cap_cpu, cr),
85            rel(&prep.cap, cr)
86        );
87    }
88    let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
89    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);
90    let mut v_dev = Vec::new();
91    for r in 0..reps {
92        let ts = std::time::Instant::now();
93        let v = dit.step(&prep, 0, &x_tok, &mods, &fs);
94        println!("device step {r}: {:.3}s", ts.elapsed().as_secs_f64());
95        if r > 0 {
96            println!("replay   {:.3e}", rel(&v, &v_dev));
97        }
98        v_dev = v;
99    }
100    let orc = &o["final_out"];
101    println!("v        dev vs oracle fp32 {:.3e}", rel(&v_dev, orc));
102    if bf.exists() {
103        let (ob, _) = read_st(&bf);
104        if let Some(vb) = ob.get("final_out") {
105            println!("v        oracle bf16 vs fp32 {:.3e} (the diffusers-bf16 floor)", rel(vb, orc));
106        }
107    }
108    if std::env::var("ZC_NOCPU").as_deref() != Ok("1") {
109        let mut prep_cpu = dit.prepare_with(cap, shape, 2, None, false).unwrap();
110        prep_cpu.cap = prep.cap.clone();
111        let ts = std::time::Instant::now();
112        let v_cpu = dit.step_cpu(&prep_cpu, &x_tok, &mods, &fs);
113        println!("cpu step: {:.3}s", ts.elapsed().as_secs_f64());
114        println!(
115            "v        dev vs cpu {:.3e}   cpu vs oracle fp32 {:.3e}",
116            rel(&v_dev, &v_cpu),
117            rel(&v_cpu, orc)
118        );
119    }
120    if std::env::var("ZC_NEG").as_deref() == Ok("1") {
121        // batch-2 pair: item 1 = the same caption at a different padded
122        // length is not available here, so use the same caption; the pair
123        // must reproduce the single forward exactly per item
124        let mut p2 = dit.prepare(cap, shape, 7, None).unwrap();
125        p2.device = false;
126        let pair_key = 9;
127        let okp = dit.attach_device_pair(&prep, &p2, pair_key, None);
128        println!("pair prepared: {okp}");
129        if okp {
130            let ts = std::time::Instant::now();
131            let r = dit.step_pair_device(pair_key, shape.n_img, 0, &x_tok, &mods, &fs).unwrap();
132            println!("pair step: {:.3}s", ts.elapsed().as_secs_f64());
133            println!("pair     item0 vs single {:.3e}   item1 vs single {:.3e}", rel(&r.0, &v_dev), rel(&r.1, &v_dev));
134        }
135    }
136}