pub struct ZImageDit {
pub cfg: ZConfig,
pub model: Option<Arc<CmfModel>>,
/* private fields */
}Expand description
The Z-Image transformer. Tensor names are the ORIGINAL diffusers names
under dit. (no rename map).
Fields§
§cfg: ZConfig§model: Option<Arc<CmfModel>>The container this was loaded from (None for load_dir) — device
paths need tensor indices into it.
Implementations§
Source§impl ZImageDit
impl ZImageDit
Sourcepub fn load_dir(dir: &Path) -> Result<Self, String>
pub fn load_dir(dir: &Path) -> Result<Self, String>
Load from a diffusers transformer/ directory (bf16/fp32
safetensors, read whole and widened to f32 — a dev/parity path).
Sourcepub fn from_cmf(model: &Arc<CmfModel>) -> Result<Self, String>
pub fn from_cmf(model: &Arc<CmfModel>) -> Result<Self, String>
Load from a packaged .cmf (dit.* tensors + dit.config_json).
Quantized projections stay mmap-resident.
Examples found in repository?
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
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}32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn geom(&self) -> ZGeom
pub fn geom(&self) -> ZGeom
Examples found in repository?
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
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}Sourcepub fn temb(&self, t_model: f32) -> Vec<f32>
pub fn temb(&self, t_model: f32) -> Vec<f32>
Timestep embedding for t_model = (1000 − 1000σ)/1000: sinusoid of
t·1000 (cos first, f32 args) → mlp.0 → SiLU → mlp.2. Returns [256].
Sourcepub fn mods_for_steps(&self, t_models: &[f32]) -> Vec<f32>
pub fn mods_for_steps(&self, t_models: &[f32]) -> Vec<f32>
Raw adaLN_modulation.0(temb) for every step and block:
[steps][2 + n_layers][4 · dim], chunks [scale_msa, gate_msa,
scale_mlp, gate_mlp] (no +1, no tanh). f64 accumulation.
Examples found in repository?
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
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}32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn final_scale_for_steps(&self, t_models: &[f32]) -> Vec<f32>
pub fn final_scale_for_steps(&self, t_models: &[f32]) -> Vec<f32>
1 + all_final_layer.2-1.adaLN_modulation.1(SiLU(temb)) per step:
[steps][dim].
Examples found in repository?
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
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}32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn embed_caption(&self, cap_feats: &[f32], l: usize) -> Vec<f32>
pub fn embed_caption(&self, cap_feats: &[f32], l: usize) -> Vec<f32>
Caption features [l, cap_feat_dim] → [l_p, dim]: pad rows are copies of the last row, RMSNorm(w, 1e-5) → Linear + b, rows ≥ l := cap_pad_token.
Examples found in repository?
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
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}Sourcepub fn refine_caption_cpu(&self, cap: &mut [f32], rope_cap: (&[f32], &[f32]))
pub fn refine_caption_cpu(&self, cap: &mut [f32], rope_cap: (&[f32], &[f32]))
The two unmodulated context-refiner blocks on the host, in place on
cap [l_p, dim].
Examples found in repository?
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
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}Sourcepub fn block_cpu(
&self,
blk: ZBlockId,
x: &mut [f32],
n: usize,
rope: (&[f32], &[f32]),
m: Option<&[f32]>,
)
pub fn block_cpu( &self, blk: ZBlockId, x: &mut [f32], n: usize, rope: (&[f32], &[f32]), m: Option<&[f32]>, )
One block on the host, in place on x [n, dim]. m = the block’s
raw modulation [4 · dim] (None = unmodulated: s = 0, gate = 1).
Sourcepub fn block_refs(&self) -> Option<ZBlockRefs<'_>>
pub fn block_refs(&self) -> Option<ZBlockRefs<'_>>
Device views of all blocks (requires from_cmf).
Sourcepub fn prepare(
&self,
cap_feats: &[f32],
shape: ZShape,
key: u64,
mods_all: Option<(&[f32], &[f32])>,
) -> Result<ZPrepared, String>
pub fn prepare( &self, cap_feats: &[f32], shape: ZShape, key: u64, mods_all: Option<(&[f32], &[f32])>, ) -> Result<ZPrepared, String>
Once per (prompt, resolution): caption embed → context refiner
(gpu::zimage_refine_caption, else CPU) → rope tables →
gpu::zimage_prepare (sets ZPrepared::device). mods_all =
(mods_for_steps, final_scale_for_steps) forwarded to the backend.
Examples found in repository?
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
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}32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn prepare_with(
&self,
cap_feats: &[f32],
shape: ZShape,
key: u64,
mods_all: Option<(&[f32], &[f32])>,
device: bool,
) -> Result<ZPrepared, String>
pub fn prepare_with( &self, cap_feats: &[f32], shape: ZShape, key: u64, mods_all: Option<(&[f32], &[f32])>, device: bool, ) -> Result<ZPrepared, String>
prepare with the device use explicit: device = false builds a
pure host state (CPU context refiner, no gpu::zimage_prepare) —
the reference a device test diffs step against.
Examples found in repository?
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}Sourcepub fn prepare_host(
&self,
cap_feats: &[f32],
shape: ZShape,
key: u64,
device_refine: bool,
) -> Result<ZPrepared, String>
pub fn prepare_host( &self, cap_feats: &[f32], shape: ZShape, key: u64, device_refine: bool, ) -> Result<ZPrepared, String>
The host half of prepare: caption embed → context refiner
(gpu::zimage_refine_caption when device_refine, else CPU) →
rope tables. device is false until attach_device.
Sourcepub fn preload_device(&self) -> bool
pub fn preload_device(&self) -> bool
Upload every device plane now (gpu::zimage_preload); independent
of the caption, so it can run beside the text encoder.
Sourcepub fn attach_device(
&self,
p: &mut ZPrepared,
mods_all: Option<(&[f32], &[f32])>,
) -> bool
pub fn attach_device( &self, p: &mut ZPrepared, mods_all: Option<(&[f32], &[f32])>, ) -> bool
gpu::zimage_prepare for a host state (sets p.device).
Sourcepub fn attach_device_pair(
&self,
pos: &ZPrepared,
neg: &ZPrepared,
key: u64,
mods_all: Option<(&[f32], &[f32])>,
) -> bool
pub fn attach_device_pair( &self, pos: &ZPrepared, neg: &ZPrepared, key: u64, mods_all: Option<(&[f32], &[f32])>, ) -> bool
One batch-2 device program for a CFG pair (item 0 = pos, item 1 =
neg, both at the same resolution) under key. false = the
backend has no batch 2 here; the caller steps the items one by one.
Examples found in repository?
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}More examples
32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn step_pair_device(
&self,
key: u64,
n_img: usize,
step: usize,
x_tok: &[f32],
mods: &[f32],
final_scale: &[f32],
) -> Option<(Vec<f32>, Vec<f32>)>
pub fn step_pair_device( &self, key: u64, n_img: usize, step: usize, x_tok: &[f32], mods: &[f32], final_scale: &[f32], ) -> Option<(Vec<f32>, Vec<f32>)>
Both items of a CFG pair prepared by attach_device_pair under
key, in one device forward: returns (v_pos, v_neg), or None when
the backend declined (the caller steps the items separately).
Examples found in repository?
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}More examples
32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn step(
&self,
p: &ZPrepared,
step: usize,
x_tok: &[f32],
mods: &[f32],
final_scale: &[f32],
) -> Vec<f32>
pub fn step( &self, p: &ZPrepared, step: usize, x_tok: &[f32], mods: &[f32], final_scale: &[f32], ) -> Vec<f32>
One DiT forward: gpu::zimage_step if p.device, else step_cpu.
x_tok [n_img_p, 64] (pad rows = copies of the last row), mods
[(2 + n_layers) · 4 · dim] of this step, final_scale [dim].
Returns v [n_img, 64] (before the pipeline’s negation).
Examples found in repository?
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
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}32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn tokens(&self, latent: &[f32], shape: &ZShape) -> Vec<f32>
pub fn tokens(&self, latent: &[f32], shape: &ZShape) -> Vec<f32>
The step input for a latent [16, h_lat, w_lat]: patchify, then pad to n_img_p rows by repeating the last row.
Examples found in repository?
32fn main() {
33 let a: Vec<String> = std::env::args().collect();
34 let model = Arc::new(cortiq_core::CmfModel::open(&a[1]).unwrap());
35 let (prompt, hh, ww) = (&a[2], a[3].parse::<usize>().unwrap(), a[4].parse::<usize>().unwrap());
36 let (steps, shift, i) = (a[5].parse::<usize>().unwrap(), a[6].parse::<f32>().unwrap(), a[7].parse::<usize>().unwrap());
37 let tok = Tokenizer::from_bytes(model.vocab.as_deref().unwrap()).unwrap();
38 let ids = cortiq_engine::zimagegen::prompt_ids(&tok, prompt, 512);
39 let cap = {
40 let _p = cortiq_engine::gpu::pause_gpu();
41 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&ids)
42 };
43 let dit = ZImageDit::from_cmf(&model).unwrap();
44 let sig = zimage::sigmas_torch_f32(steps, shift);
45 let t = zimage::t_model(sig[i]);
46 let mods = dit.mods_for_steps(&[t]);
47 let fs = dit.final_scale_for_steps(&[t]);
48 let shape = ZShape::new(hh, ww, ids.len());
49 let prep = dit.prepare(&cap, shape, 1, None).unwrap();
50 println!("device prepared: {}", prep.device);
51 let (c, lh, lw) = (dit.cfg.in_channels, hh / 8, ww / 8);
52 let lat = |dir: &str| -> Vec<f32> {
53 if i == 0 {
54 read_f32(&std::env::var("CMF_INIT_LATENT").unwrap())
55 } else {
56 read_f32(&format!("{dir}/lat_{i}.f32"))
57 }
58 };
59 let xa = lat(&a[8]);
60 // `ZC_NEG=<negative prompt>`: the CFG pair as one batch-2 device
61 // forward against the two items stepped one by one.
62 if let Ok(neg) = std::env::var("ZC_NEG") {
63 let nids = cortiq_engine::zimagegen::prompt_ids(&tok, &neg, 512);
64 let ncap = {
65 let _p = cortiq_engine::gpu::pause_gpu();
66 cortiq_engine::qwen3te::Qwen3Encoder::from_cmf(&model).unwrap().encode(&nids)
67 };
68 let mut np = dit.prepare(&ncap, ZShape::new(hh, ww, nids.len()), 2, None).unwrap();
69 let tok_a = dit.tokens(&xa, &shape);
70 // `ZC_TAPS=<dir>`: every block's residual stream of the three
71 // forwards (single pos, single neg, pair) into <dir>/{pos,neg,pair},
72 // then compared block by block (the pair's item rows vs the single).
73 let taps = std::env::var("ZC_TAPS").ok();
74 let set_taps = |sub: &str| {
75 if let Some(d) = &taps {
76 unsafe { std::env::set_var("CMF_ZI_TAPS", format!("{d}/{sub}")) };
77 }
78 };
79 set_taps("pos");
80 let vp = dit.step(&prep, i, &tok_a, &mods, &fs);
81 set_taps("neg");
82 let vn = dit.step(&np, i, &tok_a, &mods, &fs);
83 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
84 np.device = false;
85 let vn_cpu = dit.step(&np, i, &tok_a, &mods, &fs);
86 np.device = true;
87 // `ZC_SWAP=1`: the pair as (neg, pos) — tells a positional cause
88 // (item 0 vs item 1) from a content one (the caption length).
89 let swap = std::env::var("ZC_SWAP").as_deref() == Ok("1");
90 let (first, second) = if swap { (&np, &prep) } else { (&prep, &np) };
91 let ok = dit.attach_device_pair(first, second, 3, None);
92 println!(
93 "pair prepared: {ok} (pos L = {}, neg L = {}, n_img_p {}, order {})",
94 ids.len(),
95 nids.len(),
96 shape.n_img_p,
97 if swap { "neg,pos" } else { "pos,neg" }
98 );
99 set_taps("pair");
100 if let Some((p0, p1)) = dit.step_pair_device(3, shape.n_img, i, &tok_a, &mods, &fs) {
101 let (pp, pn) = if swap { (p1, p0) } else { (p0, p1) };
102 let nan = pp.iter().chain(&pn).filter(|v| !v.is_finite()).count();
103 println!("pair vs singles: pos {:.3e} neg {:.3e} non-finite {nan} single neg dev vs cpu {:.3e}",
104 rel(&pp, &vp), rel(&pn, &vn), rel(&vn, &vn_cpu));
105 }
106 unsafe { std::env::remove_var("CMF_ZI_TAPS") };
107 if let Some(d) = &taps {
108 // Row layout of the pair: image stage [item][n_img_p]; joint
109 // stage [item0: n_img_p + cp0][item1: n_img_p + cp1].
110 let h = dit.cfg.dim;
111 let cp = |l: usize| l.div_ceil(32) * 32;
112 let (cp_pos, cp_neg) = (cp(ids.len()), cp(nids.len()));
113 let (cp0, cp1) = if swap { (cp_neg, cp_pos) } else { (cp_pos, cp_neg) };
114 let nip = shape.n_img_p;
115 let rd = |p: String| -> Option<Vec<f32>> {
116 std::fs::read(&p).ok().map(|b| b.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
117 };
118 let mut names: Vec<String> = (0..2).map(|k| format!("nr{k}_out")).collect();
119 names.extend((0..dit.cfg.n_layers).map(|k| format!("l{k}_out")));
120 for n in names {
121 let (Some(sp), Some(sn), Some(pr)) = (
122 rd(format!("{d}/pos/step{i}/{n}.f32")),
123 rd(format!("{d}/neg/step{i}/{n}.f32")),
124 rd(format!("{d}/pair/step{i}/{n}.f32")),
125 ) else {
126 continue;
127 };
128 let (r0, r1) = if n.starts_with("nr") {
129 ((0, nip), (nip, nip))
130 } else {
131 ((0, nip + cp0), (nip + cp0, nip + cp1))
132 };
133 let item = |r: (usize, usize)| &pr[r.0 * h..(r.0 + r.1) * h];
134 let (ip, ineg) = if swap { (item(r1), item(r0)) } else { (item(r0), item(r1)) };
135 // image rows and caption rows separately
136 let split = |v: &[f32]| (v[..nip * h].to_vec(), v[nip * h..].to_vec());
137 let (pi, pc) = split(ip);
138 let (spi, spc) = split(&sp[..ip.len().min(sp.len())]);
139 println!(
140 "{n:9} pos: img {:.3e} cap {:.3e} neg: all {:.3e}",
141 rel(&pi, &spi),
142 if pc.is_empty() { 0.0 } else { rel(&pc, &spc) },
143 rel(ineg, &sn[..ineg.len().min(sn.len())])
144 );
145 }
146 }
147 return;
148 }
149 let tok_a = dit.tokens(&xa, &shape);
150 let va_dev = zimage::unpatchify(&dit.step(&prep, i, &tok_a, &mods, &fs), c, lh, lw);
151 let mut hp = zimage::ZPrepared { device: false, ..prep };
152 let va_cpu = zimage::unpatchify(&dit.step(&hp, i, &tok_a, &mods, &fs), c, lh, lw);
153 let va_tr = read_f32(&format!("{}/v_{i}.f32", a[8]));
154 println!("step {i} on A's lat: dev vs cpu {:.3e} cpu vs A's v {:.3e} dev vs A's v {:.3e}",
155 rel(&va_dev, &va_cpu), rel(&va_cpu, &va_tr), rel(&va_dev, &va_tr));
156 if let Some(b) = a.get(9) {
157 hp.device = true;
158 let xb = lat(b);
159 let vb_dev = zimage::unpatchify(&dit.step(&hp, i, &dit.tokens(&xb, &shape), &mods, &fs), c, lh, lw);
160 println!("inputs A vs B {:.3e} device outputs {:.3e} (B's own v {:.3e})",
161 rel(&xb, &xa), rel(&vb_dev, &va_dev), rel(&read_f32(&format!("{b}/v_{i}.f32")), &va_dev));
162 }
163}Sourcepub fn embed_image(
&self,
x_tok: &[f32],
n_img: usize,
n_img_p: usize,
) -> Vec<f32>
pub fn embed_image( &self, x_tok: &[f32], n_img: usize, n_img_p: usize, ) -> Vec<f32>
x_tok [n_img_p, 64] → [n_img_p, dim]: Linear + b, rows ≥ n_img := x_pad_token.
Sourcepub fn final_layer(&self, u: &[f32], n: usize, final_scale: &[f32]) -> Vec<f32>
pub fn final_layer(&self, u: &[f32], n: usize, final_scale: &[f32]) -> Vec<f32>
Final layer on the first n rows of u: LayerNorm(eps 1e-6, no
affine) · final_scale → Linear(dim → 64) + b.
Sourcepub fn step_cpu(
&self,
p: &ZPrepared,
x_tok: &[f32],
mods: &[f32],
final_scale: &[f32],
) -> Vec<f32>
pub fn step_cpu( &self, p: &ZPrepared, x_tok: &[f32], mods: &[f32], final_scale: &[f32], ) -> Vec<f32>
The host reference forward (WP2/WP3 gate against this): embed → pad rows := x_pad_token → noise refiner (img only) → concat [img, cap] → layers → LayerNorm(1e-6)·final_scale → Linear → image rows.
Examples found in repository?
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}Sourcepub fn step_cpu_taps(
&self,
p: &ZPrepared,
x_tok: &[f32],
mods: &[f32],
final_scale: &[f32],
tap: &mut dyn FnMut(&str, &[f32]),
) -> Vec<f32>
pub fn step_cpu_taps( &self, p: &ZPrepared, x_tok: &[f32], mods: &[f32], final_scale: &[f32], tap: &mut dyn FnMut(&str, &[f32]), ) -> Vec<f32>
step_cpu with a tap callback: (name, tensor) at the oracle’s tap
points (x_seq, nr{i}_out, u_in, l{i}_out, final_out).
Examples found in repository?
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}