Skip to main content

patchify

Function patchify 

Source
pub fn patchify(latent: &[f32], c: usize, h: usize, w: usize) -> Vec<f32>
Expand description

latent [c, h, w] → tokens [(h/2)·(w/2), 4c], feature (dy·2+dx)·c + ch.

Examples found in repository?
examples/zimage_devcheck.rs (line 68)
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    // caption: device-refined (prepare) vs CPU-refined
61    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    // Replays must be stateless: the same inputs again, and a different
72    // step's mods in between.
73    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    // the CPU path must see the SAME (device-refined) caption to isolate the DiT
80    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}