1use crate::dit::Proj;
34use crate::pool::Pool;
35use cortiq_core::CmfModel;
36use std::sync::Arc;
37
38pub 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
91pub enum Seg {
92 Text(usize),
93 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#[derive(Clone, Debug)]
111pub struct Qi21Layout {
112 pub segs: Vec<Seg>,
113}
114
115impl Qi21Layout {
116 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 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 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 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
190pub 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 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 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
236pub struct Qi21Prefix {
239 pub layout: Qi21Layout,
240 pub k: Vec<Vec<f32>>,
242 pub v: Vec<Vec<f32>>,
243 pub device_key: Option<u64>,
245 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 img_in_f32: Vec<f32>,
284 proj_out_f32: Vec<f32>,
285}
286
287pub fn gpu_allowed() -> bool {
289 std::env::var("CMF_QI21_GPU").as_deref() != Ok("0") && crate::gpu::enabled()
290}
291
292struct SendRows(*mut f32);
295unsafe impl Send for SendRows {}
296unsafe impl Sync for SendRows {}
297impl SendRows {
298 #[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
322fn 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
341fn 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
357fn 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 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
386fn 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 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
418pub 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 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 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 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 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 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 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 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 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 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 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 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 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 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 #[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 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 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 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 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 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 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 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 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 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 assert!(c[..64].iter().all(|&v| (v - 1.0).abs() < 1e-7));
966 }
967}