Skip to main content

cortiq_engine/
qwen_image21.rs

1//! Qwen-Image-2.1 denoiser (`QwenImage21Transformer2DModel`, 32
2//! single-stream blocks, 7 B) — the CPU-exact reference path and the host
3//! glue around the device step.
4//!
5//! Semantics (diffusers `transformer_qwenimage21.py`):
6//! - one joint sequence `[text/condition prefix, target image]`; text rows
7//!   come from `txt_in` (zero-centred RMSNorm → Linear → GELU-tanh →
8//!   Linear), image rows from `img_in` (64 latent channels, unpatched);
9//! - ONE shared modulation for every block: `Linear(SiLU(temb))` →
10//!   `[scale1, gate1, scale2, gate2]`; the block is
11//!   `x += tanh(g1)·attn(LN(x)·(1+s1))`, `x += tanh(g2)·mlp(LN(x)·(1+s2))`
12//!   with an affine-free LayerNorm and a bias-free SwiGLU (`out(silu(gate)·proj)`);
13//! - `causal_condition`: text and condition-image rows are modulated from
14//!   t = 0, target rows from the sampled t;
15//! - block-causal attention: `(q ≥ kv) or same_image_block` — text is
16//!   causal, every image block is bidirectional inside itself and sees
17//!   everything before it; the target sees everything;
18//! - per-head RMSNorm (weighted) on q and k, then the 3-axis complex RoPE
19//!   (θ = 10000, axes [16, 56, 56]): text advances one position on all
20//!   three axes, an image block freezes the frame axis at the running
21//!   position and centres its (h, w) grid on zero, then the position
22//!   advances by max(h, w).
23//!
24//! Because the prefix never attends to the target and never sees the
25//! timestep, its per-layer keys and values are computed ONCE per prompt
26//! ([`Qi21Dit::prefill`]) and every step recomputes only the target rows
27//! against `[cached prefix, target]` — the pipeline's KV cache.
28//!
29//! Precision: f32 activations, f64 accumulation in every norm and in the
30//! small linears of the time path; the big projections run the host GEMM
31//! (f32 accumulation) over weights dequantized to f32.
32
33use crate::dit::Proj;
34use crate::pool::Pool;
35use cortiq_core::CmfModel;
36use std::sync::Arc;
37
38/// `header.arch.arch_name` of a Qwen-Image-2.1 container.
39pub const ARCH_NAME: &str = "qwen_image21";
40
41#[derive(Clone, Debug)]
42pub struct Qi21Config {
43    pub dim: usize,
44    pub heads: usize,
45    pub head_dim: usize,
46    pub layers: usize,
47    pub in_channels: usize,
48    pub out_channels: usize,
49    pub mlp_hidden: usize,
50    pub context_in: usize,
51    pub axes: [usize; 3],
52    pub eps: f64,
53    pub causal_condition: bool,
54}
55
56impl Qi21Config {
57    pub fn from_json(v: &serde_json::Value) -> Result<Self, String> {
58        let u = |k: &str, d: usize| v[k].as_u64().map(|x| x as usize).unwrap_or(d);
59        let heads = u("num_attention_heads", 32);
60        let head_dim = u("attention_head_dim", 128);
61        let dim = heads * head_dim;
62        let axes: Vec<usize> = v["axes_dims_rope"]
63            .as_array()
64            .map(|a| a.iter().filter_map(|x| x.as_u64()).map(|x| x as usize).collect())
65            .unwrap_or_else(|| vec![16, 56, 56]);
66        if axes.len() != 3 || axes.iter().sum::<usize>() != head_dim || axes.iter().any(|a| a % 2 != 0) {
67            return Err(format!("axes_dims_rope {axes:?} does not tile head_dim {head_dim}"));
68        }
69        if u("patch_size", 1) != 1 {
70            return Err("Qwen-Image-2.1 consumes unpatched latents (patch_size 1)".into());
71        }
72        let in_channels = u("in_channels", 64);
73        Ok(Self {
74            dim,
75            heads,
76            head_dim,
77            layers: u("num_layers", 32),
78            in_channels,
79            out_channels: u("out_channels", in_channels),
80            mlp_hidden: dim * u("mlp_ratio", 3),
81            context_in: u("context_in_dim", 4096),
82            axes: [axes[0], axes[1], axes[2]],
83            eps: v["eps"].as_f64().unwrap_or(1e-6),
84            causal_condition: v["causal_condition"].as_bool().unwrap_or(true),
85        })
86    }
87}
88
89/// One run of the joint sequence.
90#[derive(Clone, Copy, Debug, PartialEq, Eq)]
91pub enum Seg {
92    Text(usize),
93    /// An image block of `h × w` latent tokens (raster order).
94    Image(usize, usize),
95}
96
97impl Seg {
98    pub fn len(&self) -> usize {
99        match *self {
100            Seg::Text(n) => n,
101            Seg::Image(h, w) => h * w,
102        }
103    }
104    pub fn is_empty(&self) -> bool {
105        self.len() == 0
106    }
107}
108
109/// The joint sequence: prefix runs, then the target image block (last).
110#[derive(Clone, Debug)]
111pub struct Qi21Layout {
112    pub segs: Vec<Seg>,
113}
114
115impl Qi21Layout {
116    /// Text-to-image: `[text(l), target(h, w)]`.
117    pub fn t2i(text: usize, h: usize, w: usize) -> Self {
118        Self {
119            segs: vec![Seg::Text(text), Seg::Image(h, w)],
120        }
121    }
122
123    pub fn total(&self) -> usize {
124        self.segs.iter().map(|s| s.len()).sum()
125    }
126
127    /// Rows before the target block.
128    pub fn prefix_len(&self) -> usize {
129        self.total() - self.target().len()
130    }
131
132    pub fn target(&self) -> Seg {
133        *self.segs.last().expect("empty layout")
134    }
135
136    pub fn validate(&self) -> Result<(), String> {
137        match self.segs.last() {
138            Some(Seg::Image(h, w)) if *h > 0 && *w > 0 => {}
139            _ => return Err("the layout must end with a non-empty target image block".into()),
140        }
141        if self.prefix_len() == 0 {
142            return Err("the prompt prefix is empty".into());
143        }
144        Ok(())
145    }
146
147    /// Per row: the image block id (−1 for text).
148    pub fn block_ids(&self) -> Vec<i32> {
149        let mut out = Vec::with_capacity(self.total());
150        let mut id = 0i32;
151        for s in &self.segs {
152            match *s {
153                Seg::Text(n) => out.extend(std::iter::repeat_n(-1, n)),
154                Seg::Image(h, w) => {
155                    out.extend(std::iter::repeat_n(id, h * w));
156                    id += 1;
157                }
158            }
159        }
160        out
161    }
162
163    /// (frame, h, w) RoPE positions of every row (diffusers `QwenImage21Rope`).
164    pub fn positions(&self) -> Vec<[i64; 3]> {
165        let mut out = Vec::with_capacity(self.total());
166        let mut pos = 0i64;
167        for s in &self.segs {
168            match *s {
169                Seg::Text(n) => {
170                    for _ in 0..n {
171                        out.push([pos, pos, pos]);
172                        pos += 1;
173                    }
174                }
175                Seg::Image(h, w) => {
176                    let (hi, wi) = (h as i64, w as i64);
177                    for r in 0..hi {
178                        for c in 0..wi {
179                            out.push([pos, r - (hi - hi / 2), c - (wi - wi / 2)]);
180                        }
181                    }
182                    pos += hi.max(wi);
183                }
184            }
185        }
186        out
187    }
188}
189
190/// cos/sin of every row's rotation, `[rows, head_dim/2]` each, in the
191/// complex-pair order the kernel applies them (frame, h, w).
192pub fn rope_tables(pos: &[[i64; 3]], axes: [usize; 3], theta: f64) -> (Vec<f32>, Vec<f32>) {
193    let half: usize = axes.iter().sum::<usize>() / 2;
194    // torch: 1 / theta^(arange(0, d, 2)/d) in f32, angle = f32(index)·freq
195    let freqs: Vec<Vec<f32>> = axes
196        .iter()
197        .map(|&d| {
198            (0..d / 2)
199                .map(|i| {
200                    let e = (2 * i) as f32 / d as f32;
201                    1.0f32 / (theta as f32).powf(e)
202                })
203                .collect()
204        })
205        .collect();
206    let mut cos = vec![0f32; pos.len() * half];
207    let mut sin = vec![0f32; pos.len() * half];
208    for (r, p) in pos.iter().enumerate() {
209        let mut j = 0;
210        for (a, fr) in freqs.iter().enumerate() {
211            for &f in fr {
212                let ang = p[a] as f32 * f;
213                cos[r * half + j] = ang.cos();
214                sin[r * half + j] = ang.sin();
215                j += 1;
216            }
217        }
218    }
219    (cos, sin)
220}
221
222struct Block {
223    /// Tensor indices of the seven projections (device path).
224    idx: Option<[usize; 7]>,
225    q: Proj,
226    k: Proj,
227    v: Proj,
228    o: Proj,
229    norm_q: Vec<f32>,
230    norm_k: Vec<f32>,
231    gate: Proj,
232    up: Proj,
233    down: Proj,
234}
235
236/// Per-layer keys and values of the prefix (post qk-norm and RoPE),
237/// `[prefix, dim]` each.
238pub struct Qi21Prefix {
239    pub layout: Qi21Layout,
240    /// Host keys/values (empty when the prefix lives on the device).
241    pub k: Vec<Vec<f32>>,
242    pub v: Vec<Vec<f32>>,
243    /// The device program holding this prefix (`gpu::qi21_*`).
244    pub device_key: Option<u64>,
245    /// cos/sin of the target rows for THIS prefix's layout: the target's
246    /// frame position follows its own prefix, so a positive and a negative
247    /// prompt of different lengths rotate their targets differently.
248    pub rope_t: (Vec<f32>, Vec<f32>),
249}
250
251impl Drop for Qi21Prefix {
252    fn drop(&mut self) {
253        if let Some(k) = self.device_key {
254            crate::gpu::qi21_release_key(k);
255        }
256    }
257}
258
259impl Qi21Prefix {
260    pub fn len(&self) -> usize {
261        self.layout.prefix_len()
262    }
263    pub fn is_empty(&self) -> bool {
264        self.len() == 0
265    }
266}
267
268pub struct Qi21Dit {
269    pub cfg: Qi21Config,
270    model: Option<Arc<CmfModel>>,
271    img_in: Proj,
272    txt_norm: Vec<f32>,
273    txt_in1: Proj,
274    txt_in2: Proj,
275    t_lin1: Proj,
276    t_lin2: Proj,
277    modulation: Proj,
278    norm_out: Proj,
279    proj_out: Proj,
280    blocks: Vec<Block>,
281    pool: Option<Arc<Pool>>,
282    /// f32 copies of `img_in` / `proj_out` for the device path.
283    img_in_f32: Vec<f32>,
284    proj_out_f32: Vec<f32>,
285}
286
287/// The device denoiser is allowed (`CMF_QI21_GPU=0` forces the host).
288pub fn gpu_allowed() -> bool {
289    std::env::var("CMF_QI21_GPU").as_deref() != Ok("0") && crate::gpu::enabled()
290}
291
292// ───────────────────────────── small helpers (copied, not shared) ─────
293
294struct SendRows(*mut f32);
295unsafe impl Send for SendRows {}
296unsafe impl Sync for SendRows {}
297impl SendRows {
298    /// SAFETY: caller guarantees disjoint `[off, off+len)` per worker.
299    #[allow(clippy::mut_from_ref)]
300    unsafe fn row(&self, off: usize, len: usize) -> &mut [f32] {
301        unsafe { std::slice::from_raw_parts_mut(self.0.add(off), len) }
302    }
303}
304
305fn pool_rows(pool: Option<&Pool>, n: usize, f: &(dyn Fn(usize, usize) + Sync)) {
306    match pool {
307        Some(p) => p.run_rows(n, f),
308        None => f(0, n),
309    }
310}
311
312fn silu(v: f32) -> f32 {
313    v / (1.0 + (-v).exp())
314}
315
316fn gelu_tanh(v: f32) -> f32 {
317    let x = v as f64;
318    (0.5 * x * (1.0 + ((2.0 / std::f64::consts::PI).sqrt() * (x + 0.044715 * x * x * x)).tanh()))
319        as f32
320}
321
322/// Affine-free LayerNorm (biased variance), f64 accumulation.
323fn layer_norm_into(x: &[f32], eps: f64, dst: &mut [f32]) {
324    let n = x.len() as f64;
325    let mean = x.iter().map(|&v| v as f64).sum::<f64>() / n;
326    let var = x.iter().map(|&v| (v as f64 - mean) * (v as f64 - mean)).sum::<f64>() / n;
327    let inv = 1.0 / (var + eps).sqrt();
328    for (d, &v) in dst.iter_mut().zip(x) {
329        *d = ((v as f64 - mean) * inv) as f32;
330    }
331}
332
333fn rms_norm_inplace(v: &mut [f32], w: &[f32], eps: f64) {
334    let ss = v.iter().map(|&x| (x as f64) * (x as f64)).sum::<f64>() / v.len() as f64;
335    let inv = 1.0 / (ss + eps).sqrt();
336    for (x, &g) in v.iter_mut().zip(w) {
337        *x = (*x as f64 * inv) as f32 * g;
338    }
339}
340
341/// Row softmax over the first `valid` entries (the rest are zeroed).
342fn softmax_prefix(row: &mut [f32], valid: usize) {
343    let (live, dead) = row.split_at_mut(valid);
344    let mx = live.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
345    let mut den = 0f64;
346    for r in live.iter_mut() {
347        *r = (*r - mx).exp();
348        den += *r as f64;
349    }
350    let inv = (1.0 / den) as f32;
351    for r in live.iter_mut() {
352        *r *= inv;
353    }
354    dead.fill(0.0);
355}
356
357/// y[n, rows] = x[n, cols] · Wᵀ on the host reference GEMM (a quantized
358/// weight is dequantized to f32 in row chunks: weight-only error).
359fn lin(p: &Proj, x: &[f32], n: usize, y: &mut [f32], pool: Option<&Pool>) {
360    use crate::zimage::host_gemm;
361    let (rows, cols) = (p.rows(), p.cols());
362    match p {
363        Proj::F32 { w, .. } => host_gemm::gemm_nt(x, w, y, n, cols, rows, pool),
364        Proj::Q(q) => {
365            const CH: usize = 768;
366            let mut wbuf = vec![0f32; CH.min(rows) * cols];
367            let mut r0 = 0;
368            while r0 < rows {
369                let rc = CH.min(rows - r0);
370                {
371                    let wp = SendRows(wbuf.as_mut_ptr());
372                    pool_rows(pool, rc, &|lo, hi| {
373                        for r in lo..hi {
374                            // SAFETY: disjoint rows.
375                            q.row_f32(r0 + r, unsafe { wp.row(r * cols, cols) });
376                        }
377                    });
378                }
379                host_gemm::gemm_nt_ld(x, &wbuf[..rc * cols], &mut y[r0..], rows, n, cols, rc, pool);
380                r0 += rc;
381            }
382        }
383    }
384}
385
386/// One row through a small projection, f64 accumulation.
387fn lin_row(p: &Proj, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
388    let (rows, cols) = (p.rows(), p.cols());
389    let mut out = vec![0f32; rows];
390    let op = SendRows(out.as_mut_ptr());
391    pool_rows(pool, rows, &|lo, hi| {
392        let mut wrow = vec![0f32; cols];
393        for o in lo..hi {
394            let w: &[f32] = match p {
395                Proj::F32 { w, .. } => &w[o * cols..(o + 1) * cols],
396                Proj::Q(q) => {
397                    q.row_f32(o, &mut wrow);
398                    &wrow
399                }
400            };
401            let s: f64 = w.iter().zip(x).map(|(&a, &c)| a as f64 * c as f64).sum();
402            // SAFETY: disjoint outputs.
403            unsafe { op.row(o, 1)[0] = s as f32 };
404        }
405    });
406    out
407}
408
409fn cfg_of(model: &CmfModel) -> Result<Qi21Config, String> {
410    let raw = model
411        .tensor_bytes("dit.config_json")
412        .map_err(|e| format!("dit.config_json: {e}"))?;
413    let v: serde_json::Value =
414        serde_json::from_slice(raw).map_err(|e| format!("dit.config_json: {e}"))?;
415    Qi21Config::from_json(&v)
416}
417
418/// The modulation of one timestep: `(mods [4·dim], final_scale [dim])`.
419pub struct Qi21Mods {
420    pub mods: Vec<f32>,
421    pub final_scale: Vec<f32>,
422}
423
424impl Qi21Dit {
425    pub fn from_cmf(model: &Arc<CmfModel>) -> Result<Self, String> {
426        let cfg = cfg_of(model)?;
427        let p = |n: &str| Proj::from_model(model, &format!("dit.{n}"));
428        let f = |n: &str| crate::dit::cmf_f32(model, &format!("dit.{n}"));
429        let mut blocks = Vec::with_capacity(cfg.layers);
430        for l in 0..cfg.layers {
431            let b = format!("transformer_blocks.{l}");
432            let names = [
433                "attn.to_q.weight",
434                "attn.to_k.weight",
435                "attn.to_v.weight",
436                "attn.to_out.0.weight",
437                "img_mlp.gate_layer.weight",
438                "img_mlp.proj.weight",
439                "img_mlp.out.weight",
440            ];
441            let mut idx = [0usize; 7];
442            let mut all = true;
443            for (i, nm) in names.iter().enumerate() {
444                match model.tensor_index(&format!("dit.{b}.{nm}")) {
445                    Some(t) => idx[i] = t,
446                    None => all = false,
447                }
448            }
449            blocks.push(Block {
450                idx: all.then_some(idx),
451                q: p(&format!("{b}.attn.to_q.weight"))?,
452                k: p(&format!("{b}.attn.to_k.weight"))?,
453                v: p(&format!("{b}.attn.to_v.weight"))?,
454                o: p(&format!("{b}.attn.to_out.0.weight"))?,
455                norm_q: f(&format!("{b}.attn.norm_q.weight"))?,
456                norm_k: f(&format!("{b}.attn.norm_k.weight"))?,
457                gate: p(&format!("{b}.img_mlp.gate_layer.weight"))?,
458                up: p(&format!("{b}.img_mlp.proj.weight"))?,
459                down: p(&format!("{b}.img_mlp.out.weight"))?,
460            });
461        }
462        let dit = Self {
463            img_in: p("img_in.weight")?,
464            txt_norm: f("txt_in.text_norm.weight")?,
465            txt_in1: p("txt_in.in_layer.weight")?,
466            txt_in2: p("txt_in.out_layer.weight")?,
467            t_lin1: p("time_text_embed.timestep_embedder.linear_1.weight")?,
468            t_lin2: p("time_text_embed.timestep_embedder.linear_2.weight")?,
469            modulation: p("modulation.1.weight")?,
470            norm_out: p("norm_out.linear.weight")?,
471            proj_out: p("proj_out.weight")?,
472            blocks,
473            pool: Pool::from_env(),
474            model: Some(model.clone()),
475            img_in_f32: f("img_in.weight")?,
476            proj_out_f32: f("proj_out.weight")?,
477            cfg,
478        };
479        dit.check_shapes()?;
480        Ok(dit)
481    }
482
483    fn check_shapes(&self) -> Result<(), String> {
484        let c = &self.cfg;
485        let want = |name: &str, p: &Proj, r: usize, k: usize| -> Result<(), String> {
486            if p.rows() != r || p.cols() != k {
487                return Err(format!(
488                    "dit.{name}: [{}, {}], the config needs [{r}, {k}]",
489                    p.rows(),
490                    p.cols()
491                ));
492            }
493            Ok(())
494        };
495        want("img_in", &self.img_in, c.dim, c.in_channels)?;
496        want("txt_in.in_layer", &self.txt_in1, c.dim, c.context_in)?;
497        want("modulation", &self.modulation, 4 * c.dim, c.dim)?;
498        want("proj_out", &self.proj_out, c.out_channels, c.dim)?;
499        let b = &self.blocks[0];
500        want("to_q", &b.q, c.dim, c.dim)?;
501        want("gate_layer", &b.gate, c.mlp_hidden, c.dim)?;
502        want("out", &b.down, c.dim, c.mlp_hidden)?;
503        Ok(())
504    }
505
506    pub fn model(&self) -> Option<&Arc<CmfModel>> {
507        self.model.as_ref()
508    }
509
510    fn pool(&self) -> Option<&Pool> {
511        self.pool.as_deref()
512    }
513
514    /// Sinusoidal timestep embedding → MLP: `temb [dim]` for the model
515    /// input `t ∈ [0, 1]` (the pipeline passes `timestep / 1000`).
516    pub fn temb(&self, t: f32) -> Vec<f32> {
517        const HALF: usize = 128;
518        let ts = 1000.0f32 * t;
519        let mut e = vec![0f32; 2 * HALF];
520        for i in 0..HALF {
521            // torch: exp(-ln(10000) · arange(half) / half) in f32
522            let f = (-(10000f32).ln() * i as f32 / HALF as f32).exp();
523            let a = ts * f;
524            e[i] = a.cos();
525            e[HALF + i] = a.sin();
526        }
527        let pool = self.pool();
528        let mut h = lin_row(&self.t_lin1, &e, pool);
529        for v in h.iter_mut() {
530            *v = silu(*v);
531        }
532        lin_row(&self.t_lin2, &h, pool)
533    }
534
535    /// Shared block modulation `[scale1, gate1, scale2, gate2]` and the
536    /// final-norm scale of one timestep.
537    pub fn mods(&self, t: f32) -> Qi21Mods {
538        let temb = self.temb(t);
539        let s: Vec<f32> = temb.iter().map(|&v| silu(v)).collect();
540        let pool = self.pool();
541        Qi21Mods {
542            mods: lin_row(&self.modulation, &s, pool),
543            final_scale: lin_row(&self.norm_out, &s, pool),
544        }
545    }
546
547    /// `txt_in`: `[n, context_in]` → `[n, dim]`.
548    pub fn embed_text(&self, feats: &[f32], n: usize) -> Vec<f32> {
549        let c = &self.cfg;
550        let pool = self.pool();
551        let mut x = feats[..n * c.context_in].to_vec();
552        for row in x.chunks_exact_mut(c.context_in) {
553            let ss = row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / c.context_in as f64;
554            let inv = 1.0 / (ss + c.eps).sqrt();
555            for (v, &w) in row.iter_mut().zip(&self.txt_norm) {
556                *v = (*v as f64 * inv) as f32 * (w + 1.0);
557            }
558        }
559        let mut h = vec![0f32; n * c.dim];
560        lin(&self.txt_in1, &x, n, &mut h, pool);
561        for v in h.iter_mut() {
562            *v = gelu_tanh(*v);
563        }
564        let mut out = vec![0f32; n * c.dim];
565        lin(&self.txt_in2, &h, n, &mut out, pool);
566        out
567    }
568
569    /// `img_in`: latent tokens `[n, in_channels]` → `[n, dim]`.
570    pub fn embed_image(&self, tok: &[f32], n: usize) -> Vec<f32> {
571        let mut out = vec![0f32; n * self.cfg.dim];
572        lin(&self.img_in, tok, n, &mut out, self.pool());
573        out
574    }
575
576    /// dst = LN(src) · (1 + scale)
577    fn norm_mod(&self, src: &[f32], scale: &[f32], dst: &mut [f32], n: usize) {
578        let hs = self.cfg.dim;
579        let eps = self.cfg.eps;
580        let sr = SendRows(dst.as_mut_ptr());
581        pool_rows(self.pool(), n, &|lo, hi| {
582            for p in lo..hi {
583                // SAFETY: disjoint rows.
584                let row = unsafe { sr.row(p * hs, hs) };
585                layer_norm_into(&src[p * hs..(p + 1) * hs], eps, row);
586                for (r, &s) in row.iter_mut().zip(scale) {
587                    *r *= 1.0 + s;
588                }
589            }
590        });
591    }
592
593    /// x += tanh(gate) ⊙ y
594    fn gated_add(&self, x: &mut [f32], y: &[f32], gate_tanh: &[f32], n: usize) {
595        let hs = self.cfg.dim;
596        let sr = SendRows(x.as_mut_ptr());
597        pool_rows(self.pool(), n, &|lo, hi| {
598            for p in lo..hi {
599                // SAFETY: disjoint rows.
600                let row = unsafe { sr.row(p * hs, hs) };
601                for ((d, &v), &g) in row.iter_mut().zip(&y[p * hs..(p + 1) * hs]).zip(gate_tanh) {
602                    *d += g * v;
603                }
604            }
605        });
606    }
607
608    /// qk RMSNorm + RoPE in place on `[n, heads·hd]`.
609    fn qk_norm_rope(&self, all: &mut [f32], w: &[f32], n: usize, cos: &[f32], sin: &[f32]) {
610        let (nh, hd) = (self.cfg.heads, self.cfg.head_dim);
611        let pairs = hd / 2;
612        let eps = self.cfg.eps;
613        let sr = SendRows(all.as_mut_ptr());
614        pool_rows(self.pool(), n, &|lo, hi| {
615            for p in lo..hi {
616                for h in 0..nh {
617                    // SAFETY: disjoint tokens.
618                    let v = unsafe { sr.row((p * nh + h) * hd, hd) };
619                    rms_norm_inplace(v, w, eps);
620                    for j in 0..pairs {
621                        let (c, s) = (cos[p * pairs + j], sin[p * pairs + j]);
622                        let (a, b) = (v[2 * j], v[2 * j + 1]);
623                        v[2 * j] = a * c - b * s;
624                        v[2 * j + 1] = a * s + b * c;
625                    }
626                }
627            }
628        });
629    }
630
631    /// Attention of `n` query rows over `m` key rows. `keys_of(i)` = how
632    /// many leading keys query `i` sees (a prefix of the key rows) plus an
633    /// extra `[lo, hi)` range it also sees (its own image block); keys not
634    /// covered are masked.
635    fn attention(
636        &self,
637        q_all: &[f32],
638        k_all: &[f32],
639        v_all: &[f32],
640        n: usize,
641        m: usize,
642        visible: &(dyn Fn(usize) -> usize + Sync),
643        out: &mut [f32],
644    ) {
645        use crate::zimage::host_gemm;
646        let (nh, hd) = (self.cfg.heads, self.cfg.head_dim);
647        let pool = self.pool();
648        let scale = 1.0 / (hd as f32).sqrt();
649        let mut qh = vec![0f32; n * hd];
650        let mut kh = vec![0f32; m * hd];
651        let mut vt = vec![0f32; hd * m];
652        let mut scores = vec![0f32; n * m];
653        let mut oh = vec![0f32; n * hd];
654        for h in 0..nh {
655            for p in 0..n {
656                for d in 0..hd {
657                    qh[p * hd + d] = q_all[(p * nh + h) * hd + d] * scale;
658                }
659            }
660            for p in 0..m {
661                kh[p * hd..(p + 1) * hd].copy_from_slice(&k_all[(p * nh + h) * hd..(p * nh + h + 1) * hd]);
662                for d in 0..hd {
663                    vt[d * m + p] = v_all[(p * nh + h) * hd + d];
664                }
665            }
666            host_gemm::gemm_nt(&qh, &kh, &mut scores, n, hd, m, pool);
667            {
668                let sp = SendRows(scores.as_mut_ptr());
669                pool_rows(pool, n, &|lo, hi| {
670                    for r in lo..hi {
671                        // SAFETY: disjoint rows.
672                        softmax_prefix(unsafe { sp.row(r * m, m) }, visible(r));
673                    }
674                });
675            }
676            host_gemm::gemm_nt(&scores, &vt, &mut oh, n, m, hd, pool);
677            for p in 0..n {
678                out[(p * nh + h) * hd..(p * nh + h + 1) * hd].copy_from_slice(&oh[p * hd..(p + 1) * hd]);
679            }
680        }
681    }
682
683    /// One block on the host, in place on `x` `[n, dim]` (the query rows).
684    /// `kv_prefix` = keys/values of earlier rows (post norm+rope) that the
685    /// queries also attend to; `visible(i)` counts the keys query `i` sees
686    /// in `[kv_prefix ++ own]` order (a leading run). Returns this block's
687    /// own keys and values.
688    #[allow(clippy::too_many_arguments)]
689    fn block_cpu(
690        &self,
691        l: usize,
692        x: &mut [f32],
693        n: usize,
694        m: &[f32],
695        rope: (&[f32], &[f32]),
696        kv_prefix: Option<(&[f32], &[f32])>,
697        visible: &(dyn Fn(usize) -> usize + Sync),
698    ) -> (Vec<f32>, Vec<f32>) {
699        let b = &self.blocks[l];
700        let hs = self.cfg.dim;
701        let pool = self.pool();
702        let (s1, g1, s2, g2) = (&m[..hs], &m[hs..2 * hs], &m[2 * hs..3 * hs], &m[3 * hs..4 * hs]);
703        let g1t: Vec<f32> = g1.iter().map(|v| v.tanh()).collect();
704        let g2t: Vec<f32> = g2.iter().map(|v| v.tanh()).collect();
705        let mut xn = vec![0f32; n * hs];
706        self.norm_mod(x, s1, &mut xn, n);
707        let mut q = vec![0f32; n * hs];
708        let mut k = vec![0f32; n * hs];
709        let mut v = vec![0f32; n * hs];
710        lin(&b.q, &xn, n, &mut q, pool);
711        lin(&b.k, &xn, n, &mut k, pool);
712        lin(&b.v, &xn, n, &mut v, pool);
713        self.qk_norm_rope(&mut q, &b.norm_q, n, rope.0, rope.1);
714        self.qk_norm_rope(&mut k, &b.norm_k, n, rope.0, rope.1);
715        let mut attn = vec![0f32; n * hs];
716        match kv_prefix {
717            Some((pk, pv)) => {
718                let lp = pk.len() / hs;
719                let mut ka = Vec::with_capacity((lp + n) * hs);
720                ka.extend_from_slice(pk);
721                ka.extend_from_slice(&k);
722                let mut va = Vec::with_capacity((lp + n) * hs);
723                va.extend_from_slice(pv);
724                va.extend_from_slice(&v);
725                self.attention(&q, &ka, &va, n, lp + n, visible, &mut attn);
726            }
727            None => self.attention(&q, &k, &v, n, n, visible, &mut attn),
728        }
729        let mut proj = vec![0f32; n * hs];
730        lin(&b.o, &attn, n, &mut proj, pool);
731        drop(attn);
732        self.gated_add(x, &proj, &g1t, n);
733        self.norm_mod(x, s2, &mut xn, n);
734        let inter = self.cfg.mlp_hidden;
735        let mut ga = vec![0f32; n * inter];
736        let mut up = vec![0f32; n * inter];
737        lin(&b.gate, &xn, n, &mut ga, pool);
738        lin(&b.up, &xn, n, &mut up, pool);
739        {
740            let sg = SendRows(ga.as_mut_ptr());
741            pool_rows(pool, n, &|lo, hi| {
742                for p in lo..hi {
743                    // SAFETY: disjoint rows.
744                    let g = unsafe { sg.row(p * inter, inter) };
745                    for (gv, &uv) in g.iter_mut().zip(&up[p * inter..(p + 1) * inter]) {
746                        *gv = silu(*gv) * uv;
747                    }
748                }
749            });
750        }
751        drop(up);
752        lin(&b.down, &ga, n, &mut proj, pool);
753        self.gated_add(x, &proj, &g2t, n);
754        (k, v)
755    }
756
757    /// Run the prefix (text rows already `embed_text`-ed, condition-image
758    /// rows `embed_image`-ed, in layout order) through every block with
759    /// the t = 0 modulation and the block-causal mask; keep each layer's
760    /// keys and values.
761    pub fn prefill(&self, mut x: Vec<f32>, layout: &Qi21Layout) -> Result<Qi21Prefix, String> {
762        layout.validate()?;
763        let lp = layout.prefix_len();
764        let hs = self.cfg.dim;
765        if x.len() != lp * hs {
766            return Err(format!("prefix rows: {} floats, the layout needs {}", x.len(), lp * hs));
767        }
768        let t0 = if self.cfg.causal_condition {
769            self.mods(0.0)
770        } else {
771            return Err("causal_condition = false has no step-independent prefix".into());
772        };
773        let pos = layout.positions();
774        let (cos, sin) = rope_tables(&pos[..lp], self.cfg.axes, 10000.0);
775        // query i sees keys [0, end_i): its own image block's end, else i+1
776        let ids = layout.block_ids();
777        let mut end = vec![0usize; lp];
778        for i in 0..lp {
779            end[i] = if ids[i] < 0 {
780                i + 1
781            } else {
782                let mut e = i + 1;
783                while e < lp && ids[e] == ids[i] {
784                    e += 1;
785                }
786                e
787            };
788        }
789        if gpu_allowed() {
790            if let Some(key) = self.prefill_device(&x, layout, &t0, (&cos, &sin), &end) {
791                return Ok(Qi21Prefix {
792                    layout: layout.clone(),
793                    k: Vec::new(),
794                    v: Vec::new(),
795                    device_key: Some(key),
796                    rope_t: self.target_rope(layout),
797                });
798            }
799        }
800        let visible = |i: usize| end[i];
801        let mut ks = Vec::with_capacity(self.cfg.layers);
802        let mut vs = Vec::with_capacity(self.cfg.layers);
803        for l in 0..self.cfg.layers {
804            let (k, v) = self.block_cpu(l, &mut x, lp, &t0.mods, (&cos, &sin), None, &visible);
805            ks.push(k);
806            vs.push(v);
807        }
808        Ok(Qi21Prefix {
809            layout: layout.clone(),
810            k: ks,
811            v: vs,
812            device_key: None,
813            rope_t: self.target_rope(layout),
814        })
815    }
816
817    fn geom(&self) -> crate::gpu::Qi21Geom {
818        crate::gpu::Qi21Geom {
819            hidden: self.cfg.dim,
820            nh: self.cfg.heads,
821            hd: self.cfg.head_dim,
822            inter: self.cfg.mlp_hidden,
823            in_ch: self.cfg.in_channels,
824            eps: self.cfg.eps as f32,
825        }
826    }
827
828    /// Build the device program and run the prefix there; `None` = the
829    /// host path runs instead.
830    fn prefill_device(
831        &self,
832        x: &[f32],
833        layout: &Qi21Layout,
834        t0: &Qi21Mods,
835        rope_p: (&[f32], &[f32]),
836        end: &[usize],
837    ) -> Option<u64> {
838        static KEY: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
839        let model = self.model.as_ref()?;
840        let refs: Vec<crate::gpu::Qi21BlockRef> = self
841            .blocks
842            .iter()
843            .map(|b| {
844                b.idx.map(|w| crate::gpu::Qi21BlockRef {
845                    w,
846                    norm_q: &b.norm_q,
847                    norm_k: &b.norm_k,
848                })
849            })
850            .collect::<Option<_>>()?;
851        let (ct, st) = self.target_rope(layout);
852        let vis: Vec<u32> = end.iter().map(|&e| e as u32).collect();
853        let key = KEY.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
854        let ok = crate::gpu::qi21_prefill(&crate::gpu::Qi21PrefillArgs {
855            model,
856            geom: self.geom(),
857            blocks: &refs,
858            img_in: &self.img_in_f32,
859            proj_out: &self.proj_out_f32,
860            key,
861            x,
862            lp: layout.prefix_len(),
863            rope_p,
864            rope_t: (&ct, &st),
865            vis: &vis,
866            mods0: &t0.mods,
867            n: layout.target().len(),
868        });
869        ok.then_some(key)
870    }
871
872    /// One denoiser call: the device program when the prefix lives there,
873    /// else the host path.
874    pub fn step(&self, prefix: &Qi21Prefix, tok: &[f32], mods: &Qi21Mods) -> Result<Vec<f32>, String> {
875        match prefix.device_key {
876            Some(key) => {
877                let n = prefix.layout.target().len();
878                let fs: Vec<f32> = mods.final_scale.iter().map(|&v| 1.0 + v).collect();
879                let mut out = vec![0f32; n * self.cfg.out_channels];
880                if crate::gpu::qi21_step(key, tok, &mods.mods, &fs, &mut out) {
881                    Ok(out)
882                } else {
883                    Err("the device denoiser step failed; rerun with CMF_QI21_GPU=0 for the host path".into())
884                }
885            }
886            None => Ok(self.step_cpu(prefix, tok, mods)),
887        }
888    }
889
890    /// Target-row RoPE tables for a prefix layout.
891    pub fn target_rope(&self, layout: &Qi21Layout) -> (Vec<f32>, Vec<f32>) {
892        let pos = layout.positions();
893        rope_tables(&pos[layout.prefix_len()..], self.cfg.axes, 10000.0)
894    }
895
896    /// One denoiser call on the host: latent tokens `[n, in_channels]` →
897    /// velocity `[n, out_channels]`, attending to the cached prefix (the
898    /// target rows rotate by the prefix's own `rope_t`).
899    pub fn step_cpu(&self, prefix: &Qi21Prefix, tok: &[f32], mods: &Qi21Mods) -> Vec<f32> {
900        let rope = (&prefix.rope_t.0[..], &prefix.rope_t.1[..]);
901        let n = prefix.layout.target().len();
902        let lp = prefix.len();
903        let hs = self.cfg.dim;
904        let mut x = self.embed_image(tok, n);
905        let all = |_: usize| lp + n;
906        for l in 0..self.cfg.layers {
907            self.block_cpu(
908                l,
909                &mut x,
910                n,
911                &mods.mods,
912                rope,
913                Some((&prefix.k[l], &prefix.v[l])),
914                &all,
915            );
916        }
917        let mut xn = vec![0f32; n * hs];
918        self.norm_mod(&x, &mods.final_scale, &mut xn, n);
919        let mut out = vec![0f32; n * self.cfg.out_channels];
920        lin(&self.proj_out, &xn, n, &mut out, self.pool());
921        out
922    }
923}
924
925#[cfg(test)]
926mod tests {
927    use super::*;
928
929    #[test]
930    fn t2i_positions_follow_the_reference_rope_layout() {
931        let l = Qi21Layout::t2i(3, 2, 3);
932        let p = l.positions();
933        assert_eq!(p[0], [0, 0, 0]);
934        assert_eq!(p[2], [2, 2, 2]);
935        // image frame = position after the text; h in [-(2-1), 1) = {-1, 0},
936        // w in [-(3-1), 1) = {-2, -1, 0}
937        assert_eq!(p[3], [3, -1, -2]);
938        assert_eq!(p[5], [3, -1, 0]);
939        assert_eq!(p[8], [3, 0, 0]);
940        assert_eq!(l.prefix_len(), 3);
941        assert_eq!(l.block_ids(), vec![-1, -1, -1, 0, 0, 0, 0, 0, 0]);
942    }
943
944    #[test]
945    fn text_after_an_image_resumes_from_the_larger_side() {
946        let l = Qi21Layout {
947            segs: vec![Seg::Text(2), Seg::Image(2, 4), Seg::Text(1), Seg::Image(1, 1)],
948        };
949        let p = l.positions();
950        // image block at frame 2, then the position advances by max(2, 4)
951        assert_eq!(p[2][0], 2);
952        assert_eq!(p[10], [6, 6, 6]);
953        assert_eq!(p[11], [7, -1, -1]);
954    }
955
956    #[test]
957    fn rope_tables_are_unit_rotations() {
958        let l = Qi21Layout::t2i(4, 2, 2);
959        let (c, s) = rope_tables(&l.positions(), [16, 56, 56], 10000.0);
960        assert_eq!(c.len(), l.total() * 64);
961        for (a, b) in c.iter().zip(&s) {
962            assert!((a * a + b * b - 1.0).abs() < 1e-5);
963        }
964        // position 0 → identity
965        assert!(c[..64].iter().all(|&v| (v - 1.0).abs() < 1e-7));
966    }
967}