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}More examples
examples/zimage_metal_check.rs (line 89)
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}