1use crate::dit::{Proj, cmf_f32};
31use crate::pool::Pool;
32use cortiq_core::CmfModel;
33use std::sync::Arc;
34
35const EPS: f64 = 1e-6;
36
37pub(crate) fn rows(pool: Option<&Pool>, n: usize, f: &(dyn Fn(usize, usize) + Sync)) {
41 match pool {
42 Some(p) => p.run_rows(n, f),
43 None => f(0, n),
44 }
45}
46
47pub(crate) struct Shared(pub(crate) *mut f32);
49unsafe impl Send for Shared {}
50unsafe impl Sync for Shared {}
51impl Shared {
52 #[allow(clippy::mut_from_ref)]
54 pub(crate) unsafe fn at(&self, off: usize, len: usize) -> &mut [f32] {
55 unsafe { std::slice::from_raw_parts_mut(self.0.add(off), len) }
56 }
57}
58
59pub(crate) fn rms_plain(x: &[f32], dst: &mut [f32]) {
61 let ss = x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / x.len() as f64;
62 let inv = 1.0 / (ss + EPS).sqrt();
63 for (d, &v) in dst.iter_mut().zip(x) {
64 *d = (v as f64 * inv) as f32;
65 }
66}
67
68fn rms_w(x: &mut [f32], w: &[f32]) {
70 let ss = x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / x.len() as f64;
71 let inv = 1.0 / (ss + EPS).sqrt();
72 for (v, &g) in x.iter_mut().zip(w) {
73 *v = (*v as f64 * inv) as f32 * g;
74 }
75}
76
77fn silu(v: f32) -> f32 {
78 v / (1.0 + (-v).exp())
79}
80
81#[inline]
89pub(crate) fn gelu_tanh(v: f32) -> f32 {
90 const K: f32 = 0.797_884_56; 0.5 * v * (1.0 + (K * (v + 0.044715 * v * v * v)).tanh())
92}
93
94pub(crate) fn gelu_tanh_rows(x: &mut [f32], pool: Option<&Pool>) {
96 let dst = Shared(x.as_mut_ptr());
97 let n = x.len();
98 let grain = 4096usize;
99 let chunks = n.div_ceil(grain);
100 rows(pool, chunks, &|s, e| {
101 let (lo, hi) = (s * grain, (e * grain).min(n));
102 let r = unsafe { dst.at(lo, hi - lo) };
103 for v in r.iter_mut() {
104 *v = gelu_tanh(*v);
105 }
106 });
107}
108
109pub(crate) fn softmax(row: &mut [f32]) {
114 #[cfg(target_arch = "aarch64")]
115 {
116 crate::attention::softmax_row(row);
117 }
118 #[cfg(not(target_arch = "aarch64"))]
119 {
120 let mx = row.iter().cloned().fold(f32::MIN, f32::max);
121 let mut den = 0f32;
122 for r in row.iter_mut() {
123 *r = (*r - mx).exp();
124 den += *r;
125 }
126 if den > 0.0 {
127 let inv = 1.0 / den;
128 for r in row.iter_mut() {
129 *r *= inv;
130 }
131 }
132 }
133}
134
135fn layer_norm(x: &[f32], dst: &mut [f32]) {
137 let n = x.len() as f64;
138 let mean = x.iter().map(|&v| v as f64).sum::<f64>() / n;
139 let var = x.iter().map(|&v| (v as f64 - mean) * (v as f64 - mean)).sum::<f64>() / n;
140 let inv = 1.0 / (var + EPS).sqrt();
141 for (d, &v) in dst.iter_mut().zip(x) {
142 *d = ((v as f64 - mean) * inv) as f32;
143 }
144}
145
146pub(crate) struct Lin {
150 w: Proj,
151 b: Option<Vec<f32>>,
152 lora: Option<crate::ltxlora::LoraBranch>,
155}
156
157impl Lin {
158 pub(crate) fn load(model: &Arc<CmfModel>, name: &str, bias: bool) -> Result<Lin, String> {
159 Lin::load_lora(model, name, bias, None)
160 }
161
162 pub(crate) fn load_lora(
166 model: &Arc<CmfModel>,
167 name: &str,
168 bias: bool,
169 bank: Option<&crate::ltxlora::LoraBank>,
170 ) -> Result<Lin, String> {
171 let w = Proj::from_model(model, &format!("{name}.weight"))?;
172 let b = if bias {
173 Some(cmf_f32(model, &format!("{name}.bias"))?)
174 } else {
175 None
176 };
177 let lora = match bank {
181 Some(k) => k.branch_for(name, w.rows(), w.cols())?,
182 None => None,
183 };
184 Ok(Lin { w, b, lora })
185 }
186
187 pub(crate) fn has_lora(&self) -> bool {
192 self.lora.as_ref().is_some_and(|l| l.live())
193 }
194
195 pub(crate) fn add_lora(&self, out: &mut [f32], x: &[f32], n: usize, pool: Option<&Pool>) {
196 if let Some(l) = &self.lora {
197 l.add(x, n, out, pool);
198 }
199 }
200
201 #[cfg(target_os = "macos")]
204 fn mapped(&self) -> Option<(&Arc<CmfModel>, usize, usize, usize)> {
205 let (model, idx) = self.w.q4tp_mapped()?;
206 Some((model, idx, self.w.rows(), self.w.cols()))
207 }
208
209 pub(crate) fn add_bias(&self, out: &mut [f32], n: usize, pool: Option<&Pool>) {
211 let Some(b) = &self.b else { return };
212 let m = self.w.rows();
213 let dst = Shared(out.as_mut_ptr());
214 rows(pool, n, &|s, e| {
215 let r = unsafe { dst.at(s * m, (e - s) * m) };
216 for row in r.chunks_exact_mut(m) {
217 for (v, &bb) in row.iter_mut().zip(b) {
218 *v += bb;
219 }
220 }
221 });
222 }
223
224 pub(crate) fn apply(&self, x: &[f32], n: usize, pool: Option<&Pool>) -> Vec<f32> {
225 let m = self.w.rows();
226 let cols = self.w.cols();
227 let t_alloc = std::time::Instant::now();
228 let mut out = vec![0f32; n * m];
229 attn_prof::ALLOC.fetch_add(
230 t_alloc.elapsed().as_micros() as u64,
231 std::sync::atomic::Ordering::Relaxed,
232 );
233 let t_mm = std::time::Instant::now();
234 #[cfg(target_os = "macos")]
240 if let (Some(l), Some((model, idx, r, cl))) = (&self.lora, self.mapped()) {
241 if l.live()
245 && crate::ltxlora::route_threshold().is_none()
246 && n >= 32
247 && n * r * cl >= 128_000_000
248 && cl % 32 == 0
249 && !crate::gpu::mm_killed()
250 && crate::gpu::enabled_here()
251 && crate::gpu_metal::q4tp_matmat_lora(
252 model, idx, x, n, r, cl, &mut out, &l.side(),
253 )
254 {
255 attn_prof::MATMAT.fetch_add(
256 t_mm.elapsed().as_micros() as u64,
257 std::sync::atomic::Ordering::Relaxed,
258 );
259 self.add_bias(&mut out, n, pool);
260 return out;
261 }
262 }
263 let per_row = cols * 4;
268 let chunk = (0x1000_0000usize / per_row.max(1)).max(1);
269 let mut done = 0usize;
270 while done < n {
271 let take = chunk.min(n - done);
272 self.w.matmat(
273 &x[done * cols..(done + take) * cols],
274 take,
275 &mut out[done * m..(done + take) * m],
276 pool,
277 );
278 done += take;
279 }
280 attn_prof::MATMAT.fetch_add(
281 t_mm.elapsed().as_micros() as u64,
282 std::sync::atomic::Ordering::Relaxed,
283 );
284 if let Some(l) = &self.lora {
285 l.add(x, n, &mut out, pool);
286 }
287 if let Some(b) = &self.b {
288 let dst = Shared(out.as_mut_ptr());
289 rows(pool, n, &|s, e| {
290 let r = unsafe { dst.at(s * m, (e - s) * m) };
291 for row in r.chunks_exact_mut(m) {
292 for (v, &bb) in row.iter_mut().zip(b) {
293 *v += bb;
294 }
295 }
296 });
297 }
298 out
299 }
300}
301
302pub struct Rope {
306 cos: Vec<f32>,
307 sin: Vec<f32>,
308 heads: usize,
309 half: usize,
310}
311
312impl Rope {
313 pub fn build(positions: &[Vec<f64>], max_pos: &[f64], dim: usize, heads: usize, theta: f64) -> Rope {
317 let ndim = max_pos.len();
318 let count = dim / (2 * ndim);
319 let idx: Vec<f64> = (0..count)
321 .map(|j| {
322 let e = if count > 1 { j as f64 / (count - 1) as f64 } else { 0.0 };
323 theta.powf(e) * std::f64::consts::PI / 2.0
324 })
325 .collect();
326 let n = positions.len();
327 let half = dim / 2;
328 let pad = half - count * ndim;
329 let mut cos = vec![0f32; n * half];
330 let mut sin = vec![0f32; n * half];
331 for (t, p) in positions.iter().enumerate() {
332 let base = t * half;
333 for i in 0..pad {
334 cos[base + i] = 1.0;
335 }
336 for (j, &ind) in idx.iter().enumerate() {
337 for (d, &mp) in max_pos.iter().enumerate() {
338 let f = ind * (p[d] / mp * 2.0 - 1.0);
339 let o = base + pad + j * ndim + d;
340 cos[o] = f.cos() as f32;
341 sin[o] = f.sin() as f32;
342 }
343 }
344 }
345 Rope { cos, sin, heads, half: half / heads }
346 }
347
348 fn apply_row(&self, t: usize, row: &mut [f32]) {
350 let dh = self.half * 2;
351 let stride = self.heads * self.half;
352 for h in 0..self.heads {
353 let off = t * stride + h * self.half;
354 let (c, s) = (&self.cos[off..off + self.half], &self.sin[off..off + self.half]);
355 let v = &mut row[h * dh..(h + 1) * dh];
356 for i in 0..self.half {
357 let (a, b) = (v[i], v[i + self.half]);
358 v[i] = a * c[i] - b * s[i];
359 v[i + self.half] = b * c[i] + a * s[i];
360 }
361 }
362 }
363}
364
365
366pub(crate) mod attn_prof {
370 use std::sync::atomic::{AtomicU64, Ordering::Relaxed};
371 pub static PROJ: AtomicU64 = AtomicU64::new(0);
372 pub static NORM: AtomicU64 = AtomicU64::new(0);
373 pub static GATHER: AtomicU64 = AtomicU64::new(0);
374 pub static SCORE: AtomicU64 = AtomicU64::new(0);
375 pub static SOFT: AtomicU64 = AtomicU64::new(0);
376 pub static VALUE: AtomicU64 = AtomicU64::new(0);
377 pub static OUT: AtomicU64 = AtomicU64::new(0);
378 pub static ALLOC: AtomicU64 = AtomicU64::new(0);
379 pub static MATMAT: AtomicU64 = AtomicU64::new(0);
380
381 pub fn add(c: &AtomicU64, t: std::time::Instant) -> std::time::Instant {
382 c.fetch_add(t.elapsed().as_micros() as u64, Relaxed);
383 std::time::Instant::now()
384 }
385
386 pub fn report() -> String {
387 let s = |c: &AtomicU64| c.swap(0, Relaxed) as f64 / 1e6;
388 format!(
389 "proj {:.2}s qk-norm+rope {:.2}s gather {:.2}s scores {:.2}s softmax {:.2}s values {:.2}s out {:.2}s [linear: alloc {:.2}s matmat {:.2}s]",
390 s(&PROJ), s(&NORM), s(&GATHER), s(&SCORE), s(&SOFT), s(&VALUE), s(&OUT),
391 s(&ALLOC), s(&MATMAT)
392 )
393 }
394}
395
396pub(crate) struct Attn {
399 q: Lin,
400 k: Lin,
401 v: Lin,
402 o: Lin,
403 q_norm: Vec<f32>,
404 k_norm: Vec<f32>,
405 gate: Option<Lin>,
406 heads: usize,
407 dh: usize,
408}
409
410impl Attn {
411 pub(crate) fn load(model: &Arc<CmfModel>, p: &str, heads: usize, dh: usize) -> Result<Attn, String> {
412 Attn::load_lora(model, p, heads, dh, None)
413 }
414
415 pub(crate) fn load_lora(
416 model: &Arc<CmfModel>,
417 p: &str,
418 heads: usize,
419 dh: usize,
420 bank: Option<&crate::ltxlora::LoraBank>,
421 ) -> Result<Attn, String> {
422 Ok(Attn {
423 q: Lin::load_lora(model, &format!("{p}.to_q"), true, bank)?,
424 k: Lin::load_lora(model, &format!("{p}.to_k"), true, bank)?,
425 v: Lin::load_lora(model, &format!("{p}.to_v"), true, bank)?,
426 o: Lin::load_lora(model, &format!("{p}.to_out.0"), true, bank)?,
427 q_norm: cmf_f32(model, &format!("{p}.q_norm.weight"))?,
428 k_norm: cmf_f32(model, &format!("{p}.k_norm.weight"))?,
429 gate: match model.tensor(&format!("{p}.to_gate_logits.weight")) {
430 Some(_) => Some(Lin::load(model, &format!("{p}.to_gate_logits"), true)?),
431 None => None,
432 },
433 heads,
434 dh,
435 })
436 }
437
438 #[cfg(target_os = "macos")]
443 fn fused_qkv(
444 &self,
445 x: &[f32],
446 n: usize,
447 ctx: &[f32],
448 m: usize,
449 pool: Option<&Pool>,
450 ) -> Option<(Vec<f32>, Vec<f32>, Vec<f32>)> {
451 if !std::ptr::eq(x.as_ptr(), ctx.as_ptr()) || n != m {
452 return None;
453 }
454 if !crate::gpu::enabled_here() || crate::gpu::mm_killed() {
455 return None;
456 }
457 if self.q.has_lora() || self.k.has_lora() || self.v.has_lora() {
462 return None;
463 }
464 let (qw, kw, vw) = (self.q.mapped()?, self.k.mapped()?, self.v.mapped()?);
465 if !Arc::ptr_eq(qw.0, kw.0) || !Arc::ptr_eq(qw.0, vw.0) {
466 return None;
467 }
468 let jobs = [
469 crate::gpu_metal::MmJob { idx: qw.1, rows: qw.2, cols: qw.3 },
470 crate::gpu_metal::MmJob { idx: kw.1, rows: kw.2, cols: kw.3 },
471 crate::gpu_metal::MmJob { idx: vw.1, rows: vw.2, cols: vw.3 },
472 ];
473 if n * jobs[0].rows * jobs[0].cols < 128_000_000 || n < 32 {
474 return None;
475 }
476 let mut oq = vec![0f32; n * jobs[0].rows];
477 let mut ok = vec![0f32; n * jobs[1].rows];
478 let mut ov = vec![0f32; n * jobs[2].rows];
479 let done = {
480 let mut outs: [&mut [f32]; 3] = [&mut oq, &mut ok, &mut ov];
481 crate::gpu_metal::q4tp_matmat_many(qw.0, &jobs, x, n, &mut outs)
482 };
483 if !done {
484 return None;
485 }
486 self.q.add_lora(&mut oq, x, n, pool);
487 self.k.add_lora(&mut ok, x, n, pool);
488 self.v.add_lora(&mut ov, x, n, pool);
489 self.q.add_bias(&mut oq, n, pool);
490 self.k.add_bias(&mut ok, n, pool);
491 self.v.add_bias(&mut ov, n, pool);
492 Some((oq, ok, ov))
493 }
494
495 #[cfg(target_os = "macos")]
498 fn fused_kv(&self, ctx: &[f32], m: usize) -> Option<(Vec<f32>, Vec<f32>)> {
499 if !crate::gpu::enabled_here() || crate::gpu::mm_killed() {
500 return None;
501 }
502 if self.k.has_lora() || self.v.has_lora() {
503 return None;
504 }
505 let (kw, vw) = (self.k.mapped()?, self.v.mapped()?);
506 if !Arc::ptr_eq(kw.0, vw.0) {
507 return None;
508 }
509 let jobs = [
510 crate::gpu_metal::MmJob { idx: kw.1, rows: kw.2, cols: kw.3 },
511 crate::gpu_metal::MmJob { idx: vw.1, rows: vw.2, cols: vw.3 },
512 ];
513 if m < 32 || m * jobs[0].rows * jobs[0].cols < 128_000_000 {
514 return None;
515 }
516 let mut ok = vec![0f32; m * jobs[0].rows];
517 let mut ov = vec![0f32; m * jobs[1].rows];
518 let done = {
519 let mut outs: [&mut [f32]; 2] = [&mut ok, &mut ov];
520 crate::gpu_metal::q4tp_matmat_many(kw.0, &jobs, ctx, m, &mut outs)
521 };
522 if !done {
523 return None;
524 }
525 self.k.add_lora(&mut ok, ctx, m, None);
526 self.v.add_lora(&mut ov, ctx, m, None);
527 self.k.add_bias(&mut ok, m, None);
528 self.v.add_bias(&mut ov, m, None);
529 Some((ok, ov))
530 }
531
532 #[cfg(not(target_os = "macos"))]
533 fn fused_kv(&self, _ctx: &[f32], _m: usize) -> Option<(Vec<f32>, Vec<f32>)> {
534 None
535 }
536
537 #[cfg(not(target_os = "macos"))]
538 fn fused_qkv(
539 &self,
540 _x: &[f32],
541 _n: usize,
542 _ctx: &[f32],
543 _m: usize,
544 _pool: Option<&Pool>,
545 ) -> Option<(Vec<f32>, Vec<f32>, Vec<f32>)> {
546 None
547 }
548
549 #[allow(clippy::too_many_arguments)]
552 pub(crate) fn forward(
553 &self,
554 x: &[f32],
555 n: usize,
556 ctx: &[f32],
557 m: usize,
558 pe_q: Option<&Rope>,
559 pe_k: Option<&Rope>,
560 mask: Option<&[f32]>,
561 pool: Option<&Pool>,
562 ) -> Vec<f32> {
563 let inner = self.heads * self.dh;
564 let prof = std::env::var("CMF_LTX_PROF").is_ok();
565 let mut t = std::time::Instant::now();
566 let (mut q, mut k, v) = match self.fused_qkv(x, n, ctx, m, pool) {
572 Some(t) => t,
573 None => match self.fused_kv(ctx, m) {
577 Some((k, v)) => (self.q.apply(x, n, pool), k, v),
578 None => (
579 self.q.apply(x, n, pool),
580 self.k.apply(ctx, m, pool),
581 self.v.apply(ctx, m, pool),
582 ),
583 },
584 };
585 if prof {
586 t = attn_prof::add(&attn_prof::PROJ, t);
587 }
588
589 let qn = Shared(q.as_mut_ptr());
590 rows(pool, n, &|s, e| {
591 let r = unsafe { qn.at(s * inner, (e - s) * inner) };
592 for (i, row) in r.chunks_exact_mut(inner).enumerate() {
593 rms_w(row, &self.q_norm);
594 if let Some(pe) = pe_q {
595 pe.apply_row(s + i, row);
596 }
597 }
598 });
599 let kn = Shared(k.as_mut_ptr());
600 rows(pool, m, &|s, e| {
601 let r = unsafe { kn.at(s * inner, (e - s) * inner) };
602 for (i, row) in r.chunks_exact_mut(inner).enumerate() {
603 rms_w(row, &self.k_norm);
604 if let Some(pe) = pe_k {
605 pe.apply_row(s + i, row);
606 }
607 }
608 });
609
610 if prof {
616 t = attn_prof::add(&attn_prof::NORM, t);
617 }
618 let mut out = vec![0f32; n * inner];
619 let scale = 1.0 / (self.dh as f32).sqrt();
620 let dh = self.dh;
621 let mut qh = vec![0f32; n * dh];
622 let mut kh = vec![0f32; m * dh];
623 let mut vh = vec![0f32; m * dh];
624 let mut sc = vec![0f32; n * m];
625 let mut oh = vec![0f32; n * dh];
626 let pv_nt = std::env::var("CMF_LTX_PV_NT").as_deref() == Ok("1");
646 for h in 0..self.heads {
647 for i in 0..n {
648 qh[i * dh..(i + 1) * dh].copy_from_slice(&q[i * inner + h * dh..][..dh]);
649 }
650 for j in 0..m {
651 kh[j * dh..(j + 1) * dh].copy_from_slice(&k[j * inner + h * dh..][..dh]);
652 }
653 if pv_nt {
654 for j in 0..m {
656 let src = &v[j * inner + h * dh..][..dh];
657 for (d, s) in src.iter().enumerate() {
658 vh[d * m + j] = *s;
659 }
660 }
661 } else {
662 for j in 0..m {
663 vh[j * dh..(j + 1) * dh].copy_from_slice(&v[j * inner + h * dh..][..dh]);
664 }
665 }
666 if prof {
667 t = attn_prof::add(&attn_prof::GATHER, t);
668 }
669 crate::fcd_ops::gemm_nt(&qh, &kh, &mut sc, n, dh, m, pool);
670 if prof {
671 t = attn_prof::add(&attn_prof::SCORE, t);
672 }
673 let sp = Shared(sc.as_mut_ptr());
674 rows(pool, n, &|s, e| {
675 let r = unsafe { sp.at(s * m, (e - s) * m) };
676 for row in r.chunks_exact_mut(m) {
677 for (x, j) in row.iter_mut().zip(0..m) {
678 *x = *x * scale + mask.map_or(0.0, |mk| mk[j]);
679 }
680 softmax(row);
681 }
682 });
683 if prof {
684 t = attn_prof::add(&attn_prof::SOFT, t);
685 }
686 oh.iter_mut().for_each(|x| *x = 0.0);
687 if pv_nt {
688 crate::fcd_ops::gemm_nt(&sc, &vh, &mut oh, n, m, dh, pool);
691 } else {
692 crate::fcd_ops::gemm_dx(&sc, &vh, &mut oh, n, dh, m, pool);
693 }
694 if prof {
695 t = attn_prof::add(&attn_prof::VALUE, t);
696 }
697 for i in 0..n {
698 out[i * inner + h * dh..i * inner + (h + 1) * dh]
699 .copy_from_slice(&oh[i * dh..(i + 1) * dh]);
700 }
701 if prof {
702 t = attn_prof::add(&attn_prof::GATHER, t);
703 }
704 }
705
706 if let Some(g) = &self.gate {
707 let logits = g.apply(x, n, pool);
708 let h = self.heads;
709 let dst = Shared(out.as_mut_ptr());
710 rows(pool, n, &|s, e| {
711 let r = unsafe { dst.at(s * inner, (e - s) * inner) };
712 for (i, row) in r.chunks_exact_mut(inner).enumerate() {
713 for hh in 0..h {
714 let gate = 2.0 / (1.0 + (-logits[(s + i) * h + hh]).exp());
715 for d in row[hh * self.dh..(hh + 1) * self.dh].iter_mut() {
716 *d *= gate;
717 }
718 }
719 }
720 });
721 }
722 let r = self.o.apply(&out, n, pool);
723 if prof {
724 attn_prof::add(&attn_prof::OUT, t);
725 }
726 r
727 }
728}
729
730struct AdaLn {
735 l1: Lin,
736 l2: Lin,
737 lin: Lin,
738 dim: usize,
739}
740
741impl AdaLn {
742 fn load(model: &Arc<CmfModel>, p: &str, dim: usize) -> Result<AdaLn, String> {
743 Ok(AdaLn {
744 l1: Lin::load(model, &format!("{p}.emb.timestep_embedder.linear_1"), true)?,
745 l2: Lin::load(model, &format!("{p}.emb.timestep_embedder.linear_2"), true)?,
746 lin: Lin::load(model, &format!("{p}.linear"), true)?,
747 dim,
748 })
749 }
750
751 fn forward(&self, t: &[f32], pool: Option<&Pool>) -> (Vec<f32>, Vec<f32>) {
753 let n = t.len();
754 let half = 128usize;
757 let mut proj = vec![0f32; n * 256];
758 let ws: Vec<f64> = (0..half)
759 .map(|j| (-(10000f64).ln() * j as f64 / half as f64).exp())
760 .collect();
761 for (i, &tv) in t.iter().enumerate() {
762 for (j, &w) in ws.iter().enumerate() {
763 let a = tv as f64 * w;
764 proj[i * 256 + j] = a.cos() as f32;
765 proj[i * 256 + half + j] = a.sin() as f32;
766 }
767 }
768 let mut h = self.l1.apply(&proj, n, pool);
769 for v in h.iter_mut() {
770 *v = silu(*v);
771 }
772 let embedded = self.l2.apply(&h, n, pool);
773 let mut act = embedded.clone();
774 for v in act.iter_mut() {
775 *v = silu(*v);
776 }
777 (self.lin.apply(&act, n, pool), embedded)
778 }
779}
780
781struct TsTable {
786 vals: Vec<f32>,
787 emb: Vec<f32>,
788 idx: Vec<usize>,
789 width: usize,
790 edim: usize,
791}
792
793impl TsTable {
794 fn build(a: &AdaLn, ts: &[f32], scale: f64, pool: Option<&Pool>) -> TsTable {
795 let mut vals: Vec<f32> = Vec::new();
796 let mut idx = Vec::with_capacity(ts.len());
797 for &t in ts {
798 match vals.iter().position(|&v| v.to_bits() == t.to_bits()) {
799 Some(i) => idx.push(i),
800 None => {
801 vals.push(t);
802 idx.push(vals.len() - 1);
803 }
804 }
805 }
806 let scaled: Vec<f32> = vals.iter().map(|&v| (v as f64 * scale) as f32).collect();
807 let (v, e) = a.forward(&scaled, pool);
808 let width = v.len() / scaled.len().max(1);
809 TsTable { vals: v, emb: e, idx, width, edim: a.dim }
810 }
811
812 fn distinct(&self) -> usize {
813 self.vals.len() / self.width.max(1)
814 }
815
816 fn row(&self, r: usize) -> &[f32] {
817 &self.vals[r * self.width..(r + 1) * self.width]
818 }
819
820 fn emb_row(&self, r: usize) -> &[f32] {
821 &self.emb[r * self.edim..(r + 1) * self.edim]
822 }
823
824 fn triples(&self, table: &[f32], dim: usize, off: usize) -> Vec<[Vec<f32>; 3]> {
828 (0..self.distinct())
829 .map(|r| {
830 let v = self.row(r);
831 std::array::from_fn(|j| {
832 let o = (off + j) * dim;
833 (0..dim).map(|d| table[o + d] + v[o + d]).collect()
834 })
835 })
836 .collect()
837 }
838
839 fn pairs(&self, table: &[f32], dim: usize, off: usize) -> Vec<[Vec<f32>; 2]> {
842 (0..self.distinct())
843 .map(|r| {
844 let v = self.row(r);
845 std::array::from_fn(|j| {
846 let o = (off + j) * dim;
847 (0..dim).map(|d| table[o + d] + v[o + d]).collect()
848 })
849 })
850 .collect()
851 }
852}
853
854
855fn ada_zero_rows(
859 x: &[f32],
860 out: &mut [f32],
861 n: usize,
862 dim: usize,
863 mods: &[[Vec<f32>; 3]],
864 idx: &[usize],
865 pool: Option<&Pool>,
866) {
867 let dst = Shared(out.as_mut_ptr());
868 rows(pool, n, &|s, e| {
869 let r = unsafe { dst.at(s * dim, (e - s) * dim) };
870 for (row, i) in r.chunks_exact_mut(dim).zip(s..e) {
871 let md = &mods[idx[i]];
872 rms_plain(&x[i * dim..(i + 1) * dim], row);
873 for d in 0..dim {
874 row[d] = row[d] * (1.0 + md[1][d]) + md[0][d];
875 }
876 }
877 });
878}
879
880fn add_gated(
882 x: &mut [f32],
883 y: &[f32],
884 n: usize,
885 dim: usize,
886 mods: &[[Vec<f32>; 3]],
887 idx: &[usize],
888 pool: Option<&Pool>,
889) {
890 let dst = Shared(x.as_mut_ptr());
891 rows(pool, n, &|s, e| {
892 let r = unsafe { dst.at(s * dim, (e - s) * dim) };
893 for (row, i) in r.chunks_exact_mut(dim).zip(s..e) {
894 let g = &mods[idx[i]][2];
895 for d in 0..dim {
896 row[d] += y[i * dim + d] * g[d];
897 }
898 }
899 });
900}
901
902fn post_sa_rows(
905 x: &mut [f32],
906 y: &[f32],
907 normed: &mut [f32],
908 n: usize,
909 dim: usize,
910 mods: &[[Vec<f32>; 3]],
911 idx: &[usize],
912 pool: Option<&Pool>,
913) {
914 let a = Shared(x.as_mut_ptr());
915 let b = Shared(normed.as_mut_ptr());
916 rows(pool, n, &|s, e| {
917 let xr = unsafe { a.at(s * dim, (e - s) * dim) };
918 let nr = unsafe { b.at(s * dim, (e - s) * dim) };
919 for ((row, nrow), i) in xr.chunks_exact_mut(dim).zip(nr.chunks_exact_mut(dim)).zip(s..e) {
920 let g = &mods[idx[i]][2];
921 for d in 0..dim {
922 row[d] += y[i * dim + d] * g[d];
923 }
924 rms_plain(row, nrow);
925 }
926 });
927}
928
929fn add_scaled(x: &mut [f32], y: &[f32], n: usize, dim: usize, g: &[f32], pool: Option<&Pool>) {
931 let dst = Shared(x.as_mut_ptr());
932 rows(pool, n, &|s, e| {
933 let r = unsafe { dst.at(s * dim, (e - s) * dim) };
934 for (row, i) in r.chunks_exact_mut(dim).zip(s..e) {
935 for d in 0..dim {
936 row[d] += y[i * dim + d] * g[d];
937 }
938 }
939 });
940}
941
942fn affine_rows(
944 x: &[f32],
945 out: &mut [f32],
946 n: usize,
947 dim: usize,
948 mods: &[[Vec<f32>; 3]],
949 idx: &[usize],
950 pool: Option<&Pool>,
951) {
952 let dst = Shared(out.as_mut_ptr());
953 rows(pool, n, &|s, e| {
954 let r = unsafe { dst.at(s * dim, (e - s) * dim) };
955 for (row, i) in r.chunks_exact_mut(dim).zip(s..e) {
956 let md = &mods[idx[i]];
957 for d in 0..dim {
958 row[d] = x[i * dim + d] * (1.0 + md[1][d]) + md[0][d];
959 }
960 }
961 });
962}
963
964
965#[derive(Default)]
968struct Prof {
969 on: bool,
970 t: [f64; 6],
971}
972
973const P_ADALN: usize = 0;
974const P_SELF: usize = 1;
975const P_CROSS: usize = 2;
976const P_FUSE: usize = 3;
977const P_FF: usize = 4;
978const P_MOD: usize = 5;
979
980impl Prof {
981 fn new() -> Prof {
982 Prof { on: std::env::var("CMF_LTX_PROF").is_ok(), t: [0.0; 6] }
983 }
984 #[inline]
985 fn tick(&mut self, slot: usize, at: std::time::Instant) -> std::time::Instant {
986 if self.on {
987 self.t[slot] += at.elapsed().as_secs_f64();
988 return std::time::Instant::now();
989 }
990 at
991 }
992 fn report(&self) {
993 if !self.on {
994 return;
995 }
996 #[cfg(target_os = "macos")]
999 {
1000 use std::sync::atomic::Ordering::Relaxed;
1001 let n = crate::gpu_metal::MM_N.swap(0, Relaxed);
1002 if n > 0 {
1003 let us = |a: &std::sync::atomic::AtomicU64| a.swap(0, Relaxed) as f64 / 1e6;
1004 println!(
1005 " q4tp on device: {n} calls, upload {:.2}s kernel {:.2}s download {:.2}s",
1006 us(&crate::gpu_metal::MM_UP),
1007 us(&crate::gpu_metal::MM_GPU),
1008 us(&crate::gpu_metal::MM_DN),
1009 );
1010 }
1011 }
1012 println!(" attention: {}", attn_prof::report());
1013 let names = ["adaln", "self-attn", "cross-attn", "a<->v", "ffn", "modulate"];
1014 let total: f64 = self.t.iter().sum();
1015 let parts: Vec<String> = names
1016 .iter()
1017 .zip(&self.t)
1018 .map(|(n, v)| format!("{n} {v:.1}s ({:.0}%)", 100.0 * v / total.max(1e-9)))
1019 .collect();
1020 println!(" profile: {}", parts.join(" "));
1021 }
1022}
1023
1024struct Stream {
1027 attn1: Attn,
1028 attn2: Attn,
1029 ff_in: Lin,
1030 ff_out: Lin,
1031 sst: Vec<f32>, prompt_sst: Vec<f32>, }
1034
1035impl Stream {
1036 fn load(
1037 model: &Arc<CmfModel>,
1038 p: &str,
1039 prefix: &str,
1040 heads: usize,
1041 dh: usize,
1042 ff_bias: bool,
1043 bank: Option<&crate::ltxlora::LoraBank>,
1044 ) -> Result<Stream, String> {
1045 let a = |n: &str| format!("{p}.{prefix}{n}");
1046 Ok(Stream {
1047 attn1: Attn::load_lora(model, &a("attn1"), heads, dh, bank)?,
1048 attn2: Attn::load_lora(model, &a("attn2"), heads, dh, bank)?,
1049 ff_in: Lin::load_lora(model, &a("ff.net.0.proj"), ff_bias, bank)?,
1050 ff_out: Lin::load_lora(model, &a("ff.net.2"), ff_bias, bank)?,
1051 sst: cmf_f32(model, &a("scale_shift_table"))?,
1052 prompt_sst: cmf_f32(model, &a("prompt_scale_shift_table"))?,
1053 })
1054 }
1055
1056 fn ff(&self, x: &[f32], n: usize, pool: Option<&Pool>) -> Vec<f32> {
1057 let mut h = self.ff_in.apply(x, n, pool);
1058 gelu_tanh_rows(&mut h, pool);
1059 self.ff_out.apply(&h, n, pool)
1060 }
1061}
1062
1063struct Block {
1064 video: Stream,
1065 audio: Stream,
1066 a2v: Attn,
1067 v2a: Attn,
1068 sst_a2v_video: Vec<f32>, sst_a2v_audio: Vec<f32>, }
1071
1072pub struct StreamInput {
1077 pub latent: Vec<f32>,
1079 pub tokens: usize,
1080 pub timesteps: Vec<f32>,
1082 pub positions: Vec<Vec<f64>>,
1084 pub context: Vec<f32>,
1086 pub ctx_len: usize,
1087 pub context_mask: Vec<f32>,
1089 pub keyframes: Vec<f32>,
1091 pub sigma: f32,
1093}
1094
1095pub struct LtxDit {
1096 model: Arc<CmfModel>,
1097 blocks: Vec<Block>,
1098 patchify: Lin,
1099 a_patchify: Lin,
1100 keyframes_emb: Option<Vec<f32>>,
1101 adaln: AdaLn,
1102 a_adaln: AdaLn,
1103 prompt_adaln: AdaLn,
1104 a_prompt_adaln: AdaLn,
1105 av_v_ss: AdaLn,
1106 av_a_ss: AdaLn,
1107 av_a2v_gate: AdaLn,
1108 av_v2a_gate: AdaLn,
1109 proj_out: Lin,
1110 a_proj_out: Lin,
1111 sst_out: Vec<f32>,
1112 a_sst_out: Vec<f32>,
1113 pub heads: usize,
1114 pub dh: usize,
1115 pub a_heads: usize,
1116 pub a_dh: usize,
1117 pub max_pos: Vec<f64>,
1118 pub a_max_pos: Vec<f64>,
1119 pub cross_max_pos: f64,
1120 pub theta: f64,
1121 pub t_scale: f64,
1122 pub av_t_scale: f64,
1123 pub audio_cross_dim: usize,
1124}
1125
1126impl LtxDit {
1127 pub fn from_cmf(model: &Arc<CmfModel>) -> Result<LtxDit, String> {
1128 LtxDit::from_cmf_lora(model, None)
1129 }
1130
1131 pub fn from_cmf_lora(
1135 model: &Arc<CmfModel>,
1136 bank: Option<&crate::ltxlora::LoraBank>,
1137 ) -> Result<LtxDit, String> {
1138 let cfg_bytes = ["ltx.config_json", "dit.config_json"]
1139 .iter()
1140 .find_map(|n| model.tensor(n).map(|e| model.entry_bytes(e)))
1141 .ok_or("container carries no ltx.config_json")?;
1142 let cfg: serde_json::Value =
1143 serde_json::from_slice(cfg_bytes).map_err(|e| format!("ltx.config_json: {e}"))?;
1144 let t = cfg.get("transformer").unwrap_or(&cfg).clone();
1145 let g = |k: &str, d: f64| t.get(k).and_then(|v| v.as_f64()).unwrap_or(d);
1146 let heads = g("num_attention_heads", 32.0) as usize;
1147 let dh = g("attention_head_dim", 128.0) as usize;
1148 let a_heads = g("audio_num_attention_heads", 32.0) as usize;
1149 let a_dh = g("audio_attention_head_dim", 64.0) as usize;
1150 let n_layers = g("num_layers", 48.0) as usize;
1151 let ff_bias = t.get("ff_bias").and_then(|v| v.as_bool()).unwrap_or(true);
1152 let a_ff_bias = t.get("audio_ff_bias").and_then(|v| v.as_bool()).unwrap_or(true);
1153 let arr = |k: &str, d: Vec<f64>| -> Vec<f64> {
1154 t.get(k)
1155 .and_then(|v| v.as_array())
1156 .map(|a| a.iter().filter_map(|x| x.as_f64()).collect())
1157 .unwrap_or(d)
1158 };
1159 let max_pos = arr("positional_embedding_max_pos", vec![20.0, 2048.0, 2048.0]);
1160 let a_max_pos = arr("audio_positional_embedding_max_pos", vec![20.0]);
1161 let cross_max_pos = max_pos[0].max(a_max_pos[0]);
1162 let dim = heads * dh;
1163 let a_dim = a_heads * a_dh;
1164
1165 let mut blocks = Vec::with_capacity(n_layers);
1166 for i in 0..n_layers {
1167 let p = format!("dit.transformer_blocks.{i}");
1168 blocks.push(Block {
1169 video: Stream::load(model, &p, "", heads, dh, ff_bias, bank)?,
1170 audio: Stream::load(model, &p, "audio_", a_heads, a_dh, a_ff_bias, bank)?,
1171 a2v: Attn::load(model, &format!("{p}.audio_to_video_attn"), a_heads, a_dh)?,
1172 v2a: Attn::load(model, &format!("{p}.video_to_audio_attn"), a_heads, a_dh)?,
1173 sst_a2v_video: cmf_f32(model, &format!("{p}.scale_shift_table_a2v_ca_video"))?,
1174 sst_a2v_audio: cmf_f32(model, &format!("{p}.scale_shift_table_a2v_ca_audio"))?,
1175 });
1176 }
1177 let _ = (dim, a_dim);
1178 Ok(LtxDit {
1179 blocks,
1180 patchify: Lin::load(model, "dit.patchify_proj", true)?,
1181 a_patchify: Lin::load(model, "dit.audio_patchify_proj", true)?,
1182 keyframes_emb: match model.tensor("dit.keyframes_abs_pos_embedding") {
1183 Some(_) => Some(cmf_f32(model, "dit.keyframes_abs_pos_embedding")?),
1184 None => None,
1185 },
1186 adaln: AdaLn::load(model, "dit.adaln_single", dim)?,
1187 a_adaln: AdaLn::load(model, "dit.audio_adaln_single", a_dim)?,
1188 prompt_adaln: AdaLn::load(model, "dit.prompt_adaln_single", dim)?,
1189 a_prompt_adaln: AdaLn::load(model, "dit.audio_prompt_adaln_single", a_dim)?,
1190 av_v_ss: AdaLn::load(model, "dit.av_ca_video_scale_shift_adaln_single", dim)?,
1191 av_a_ss: AdaLn::load(model, "dit.av_ca_audio_scale_shift_adaln_single", a_dim)?,
1192 av_a2v_gate: AdaLn::load(model, "dit.av_ca_a2v_gate_adaln_single", dim)?,
1193 av_v2a_gate: AdaLn::load(model, "dit.av_ca_v2a_gate_adaln_single", a_dim)?,
1194 proj_out: Lin::load(model, "dit.proj_out", true)?,
1195 a_proj_out: Lin::load(model, "dit.audio_proj_out", true)?,
1196 sst_out: cmf_f32(model, "dit.scale_shift_table")?,
1197 a_sst_out: cmf_f32(model, "dit.audio_scale_shift_table")?,
1198 heads,
1199 dh,
1200 a_heads,
1201 a_dh,
1202 max_pos,
1203 a_max_pos,
1204 cross_max_pos,
1205 theta: g("positional_embedding_theta", 10000.0),
1206 t_scale: g("timestep_scale_multiplier", 1000.0),
1207 av_t_scale: g("av_ca_timestep_scale_multiplier", 1.0),
1208 audio_cross_dim: g("audio_cross_attention_dim", 2048.0) as usize,
1209 model: model.clone(),
1210 })
1211 }
1212
1213 pub fn blocks(&self) -> usize {
1214 self.blocks.len()
1215 }
1216
1217 pub fn container(&self) -> &Arc<CmfModel> {
1218 &self.model
1219 }
1220
1221 pub fn forward(
1223 &self,
1224 video: &StreamInput,
1225 audio: &StreamInput,
1226 pool: Option<&Pool>,
1227 ) -> (Vec<f32>, Vec<f32>) {
1228 self.forward_traced(video, audio, pool, &mut |_, _| {})
1229 }
1230
1231 pub fn forward_traced(
1232 &self,
1233 video: &StreamInput,
1234 audio: &StreamInput,
1235 pool: Option<&Pool>,
1236 trace: &mut dyn FnMut(&str, &[f32]),
1237 ) -> (Vec<f32>, Vec<f32>) {
1238 let _trust = crate::gpu::trust_gpu();
1242 let dim = self.heads * self.dh;
1243 let a_dim = self.a_heads * self.a_dh;
1244 let (n, m) = (video.tokens, audio.tokens);
1245
1246 let mut vx = self.patchify.apply(&video.latent, n, pool);
1248 if let Some(emb) = &self.keyframes_emb {
1249 for i in 0..n {
1250 if video.keyframes.get(i).copied().unwrap_or(0.0) > 0.0 {
1251 for (d, &e) in vx[i * dim..(i + 1) * dim].iter_mut().zip(emb) {
1252 *d += e;
1253 }
1254 }
1255 }
1256 }
1257 let mut ax = self.a_patchify.apply(&audio.latent, m, pool);
1258 trace("v.args.x", &vx);
1259 trace("a.args.x", &ax);
1260
1261 let vt = TsTable::build(&self.adaln, &video.timesteps, self.t_scale, pool);
1263 let at = TsTable::build(&self.a_adaln, &audio.timesteps, self.t_scale, pool);
1264 let vpt = TsTable::build(&self.prompt_adaln, &[video.sigma], self.t_scale, pool);
1265 let apt = TsTable::build(&self.a_prompt_adaln, &[audio.sigma], self.t_scale, pool);
1266 if std::env::var("CMF_LTX_PROMPTADALN").is_ok() {
1267 let r = vpt.row(0);
1268 let mx = r.iter().fold(0f32, |m, &v| m.max(v.abs()));
1269 let sum: f32 = r.iter().sum();
1270 eprintln!("prompt-adaln sigma={:.6} max|row|={mx:.6e} sum={sum:.6e}", video.sigma);
1271 }
1272 let vxs = TsTable::build(&self.av_v_ss, &video.timesteps, self.t_scale, pool);
1273 let axs = TsTable::build(&self.av_a_ss, &audio.timesteps, self.t_scale, pool);
1274 let vgt = TsTable::build(&self.av_a2v_gate, &[audio.sigma], self.av_t_scale, pool);
1277 let agt = TsTable::build(&self.av_v2a_gate, &[video.sigma], self.av_t_scale, pool);
1278
1279 let v_pe = Rope::build(&video.positions, &self.max_pos, dim, self.heads, self.theta);
1281 let a_pe = Rope::build(&audio.positions, &self.a_max_pos, a_dim, self.a_heads, self.theta);
1282 let time_only = |p: &[Vec<f64>]| p.iter().map(|r| vec![r[0]]).collect::<Vec<_>>();
1283 let v_xpe = Rope::build(
1284 &time_only(&video.positions),
1285 &[self.cross_max_pos],
1286 self.audio_cross_dim,
1287 self.heads,
1288 self.theta,
1289 );
1290 let a_xpe = Rope::build(
1291 &time_only(&audio.positions),
1292 &[self.cross_max_pos],
1293 self.audio_cross_dim,
1294 self.a_heads,
1295 self.theta,
1296 );
1297
1298 let vmask = (!video.context_mask.is_empty()).then_some(&video.context_mask[..]);
1299 let amask = (!audio.context_mask.is_empty()).then_some(&audio.context_mask[..]);
1300
1301 let mut prof = Prof::new();
1302 for (bi, blk) in self.blocks.iter().enumerate() {
1303 let mut pt = std::time::Instant::now();
1304 let v_msa = vt.triples(&blk.video.sst, dim, 0);
1305 let v_ca = vt.triples(&blk.video.sst, dim, 6);
1306 let v_mlp = vt.triples(&blk.video.sst, dim, 3);
1307 let a_msa = at.triples(&blk.audio.sst, a_dim, 0);
1308 let a_ca = at.triples(&blk.audio.sst, a_dim, 6);
1309 let a_mlp = at.triples(&blk.audio.sst, a_dim, 3);
1310
1311 pt = prof.tick(P_ADALN, pt);
1313 let mut vnorm = vec![0f32; n * dim];
1314 ada_zero_rows(&vx, &mut vnorm, n, dim, &v_msa, &vt.idx, pool);
1315 pt = prof.tick(P_MOD, pt);
1316 if bi == 0 {
1317 trace("v.b0.sa.in", &vnorm);
1318 }
1319 let vsa = blk
1320 .video
1321 .attn1
1322 .forward(&vnorm, n, &vnorm, n, Some(&v_pe), Some(&v_pe), None, pool);
1323 if bi == 0 {
1324 trace("v.b0.sa.out", &vsa);
1325 }
1326 pt = prof.tick(P_SELF, pt);
1327 let mut vnormed = vec![0f32; n * dim];
1328 post_sa_rows(&mut vx, &vsa, &mut vnormed, n, dim, &v_msa, &vt.idx, pool);
1329 let mut vq = vec![0f32; n * dim];
1330 affine_rows(&vnormed, &mut vq, n, dim, &v_ca, &vt.idx, pool);
1331 let vctx = modulate_kv(&video.context, video.ctx_len, dim, &blk.video.prompt_sst, vpt.row(0), pool);
1332 let vca = blk.video.attn2.forward(&vq, n, &vctx, video.ctx_len, None, None, vmask, pool);
1333 if bi == 0 {
1334 trace("v.b0.ca.in", &vq);
1335 trace("v.b0.ca.ctx", &vctx);
1336 trace("v.b0.ca.out", &vca);
1337 }
1338 add_gated(&mut vx, &vca, n, dim, &v_ca, &vt.idx, pool);
1339 pt = prof.tick(P_CROSS, pt);
1340
1341 let mut anorm = vec![0f32; m * a_dim];
1343 ada_zero_rows(&ax, &mut anorm, m, a_dim, &a_msa, &at.idx, pool);
1344 if bi == 0 {
1345 trace("a.b0.sa.in", &anorm);
1346 }
1347 let asa = blk
1348 .audio
1349 .attn1
1350 .forward(&anorm, m, &anorm, m, Some(&a_pe), Some(&a_pe), None, pool);
1351 if bi == 0 {
1352 trace("a.b0.sa.out", &asa);
1353 }
1354 let mut anormed = vec![0f32; m * a_dim];
1355 post_sa_rows(&mut ax, &asa, &mut anormed, m, a_dim, &a_msa, &at.idx, pool);
1356 let mut aq = vec![0f32; m * a_dim];
1357 affine_rows(&anormed, &mut aq, m, a_dim, &a_ca, &at.idx, pool);
1358 let actx = modulate_kv(&audio.context, audio.ctx_len, a_dim, &blk.audio.prompt_sst, apt.row(0), pool);
1359 let aca = blk.audio.attn2.forward(&aq, m, &actx, audio.ctx_len, None, None, amask, pool);
1360 if bi == 0 {
1361 trace("a.b0.ca.in", &aq);
1362 trace("a.b0.ca.ctx", &actx);
1363 trace("a.b0.ca.out", &aca);
1364 }
1365 add_gated(&mut ax, &aca, m, a_dim, &a_ca, &at.idx, pool);
1366 pt = prof.tick(P_CROSS, pt);
1367
1368 let vx_pre = vx.clone();
1370 let ax_pre = ax.clone();
1371 let a2v_vp = vxs.pairs(&blk.sst_a2v_video, dim, 0);
1372 let a2v_ap = axs.pairs(&blk.sst_a2v_audio, a_dim, 0);
1373 let a2v_v = ada_pair(&vx_pre, n, dim, &a2v_vp, &vxs.idx, pool);
1374 let a2v_a = ada_pair(&ax_pre, m, a_dim, &a2v_ap, &axs.idx, pool);
1375 let a2v = blk
1376 .a2v
1377 .forward(&a2v_v, n, &a2v_a, m, Some(&v_xpe), Some(&a_xpe), None, pool);
1378 if bi == 0 {
1379 trace("v.b0.a2v.in", &a2v_v);
1380 trace("v.b0.a2v.ctx", &a2v_a);
1381 trace("v.b0.a2v.out", &a2v);
1382 }
1383 let gate_a2v = gate_row(&blk.sst_a2v_video, dim, vgt.row(0));
1384 add_scaled(&mut vx, &a2v, n, dim, &gate_a2v, pool);
1385 let v2a_ap = axs.pairs(&blk.sst_a2v_audio, a_dim, 2);
1386 let v2a_vp = vxs.pairs(&blk.sst_a2v_video, dim, 2);
1387 let v2a_a = ada_pair(&ax_pre, m, a_dim, &v2a_ap, &axs.idx, pool);
1388 let v2a_v = ada_pair(&vx_pre, n, dim, &v2a_vp, &vxs.idx, pool);
1389 let v2a = blk
1390 .v2a
1391 .forward(&v2a_a, m, &v2a_v, n, Some(&a_xpe), Some(&v_xpe), None, pool);
1392 if bi == 0 {
1393 trace("a.b0.v2a.in", &v2a_a);
1394 trace("a.b0.v2a.ctx", &v2a_v);
1395 trace("a.b0.v2a.out", &v2a);
1396 }
1397 let gate_v2a = gate_row(&blk.sst_a2v_audio, a_dim, agt.row(0));
1398 add_scaled(&mut ax, &v2a, m, a_dim, &gate_v2a, pool);
1399 pt = prof.tick(P_FUSE, pt);
1400
1401 let mut vsc = vec![0f32; n * dim];
1403 ada_zero_rows(&vx, &mut vsc, n, dim, &v_mlp, &vt.idx, pool);
1404 let vff = blk.video.ff(&vsc, n, pool);
1405 if bi == 0 {
1406 trace("v.b0.ff.in", &vsc);
1407 trace("v.b0.ff.out", &vff);
1408 }
1409 add_gated(&mut vx, &vff, n, dim, &v_mlp, &vt.idx, pool);
1410 let mut asc = vec![0f32; m * a_dim];
1411 ada_zero_rows(&ax, &mut asc, m, a_dim, &a_mlp, &at.idx, pool);
1412 let aff = blk.audio.ff(&asc, m, pool);
1413 if bi == 0 {
1414 trace("a.b0.ff.in", &asc);
1415 trace("a.b0.ff.out", &aff);
1416 }
1417 add_gated(&mut ax, &aff, m, a_dim, &a_mlp, &at.idx, pool);
1418 pt = prof.tick(P_FF, pt);
1419 trace(&format!("v.block{bi}"), &vx);
1420 trace(&format!("a.block{bi}"), &ax);
1421 }
1422
1423 prof.report();
1424
1425 let vout = head(&vx, n, dim, &self.sst_out, &vt, &self.proj_out, pool);
1427 let aout = head(&ax, m, a_dim, &self.a_sst_out, &at, &self.a_proj_out, pool);
1428 trace("v.out", &vout);
1429 trace("a.out", &aout);
1430 (vout, aout)
1431 }
1432}
1433
1434fn modulate_kv(
1437 ctx: &[f32],
1438 len: usize,
1439 dim: usize,
1440 table: &[f32],
1441 extra: &[f32],
1442 pool: Option<&Pool>,
1443) -> Vec<f32> {
1444 let mut out = vec![0f32; len * dim];
1445 let shift: Vec<f32> = (0..dim).map(|d| table[d] + extra[d]).collect();
1446 let scale: Vec<f32> = (0..dim).map(|d| table[dim + d] + extra[dim + d]).collect();
1447 let dst = Shared(out.as_mut_ptr());
1451 rows(pool, len, &|s, e| {
1452 let r = unsafe { dst.at(s * dim, (e - s) * dim) };
1453 for (row, i) in r.chunks_exact_mut(dim).zip(s..e) {
1454 for d in 0..dim {
1455 row[d] = ctx[i * dim + d] * (1.0 + scale[d]) + shift[d];
1456 }
1457 }
1458 });
1459 out
1460}
1461
1462
1463fn ada_pair(
1465 x: &[f32],
1466 n: usize,
1467 dim: usize,
1468 pairs: &[[Vec<f32>; 2]],
1469 idx: &[usize],
1470 pool: Option<&Pool>,
1471) -> Vec<f32> {
1472 let mut out = vec![0f32; n * dim];
1473 let dst = Shared(out.as_mut_ptr());
1474 rows(pool, n, &|s, e| {
1475 let r = unsafe { dst.at(s * dim, (e - s) * dim) };
1476 for (row, i) in r.chunks_exact_mut(dim).zip(s..e) {
1477 let p = &pairs[idx[i]];
1478 rms_plain(&x[i * dim..(i + 1) * dim], row);
1479 for d in 0..dim {
1480 row[d] = row[d] * (1.0 + p[0][d]) + p[1][d];
1481 }
1482 }
1483 });
1484 out
1485}
1486
1487fn gate_row(table: &[f32], dim: usize, extra: &[f32]) -> Vec<f32> {
1490 (0..dim).map(|d| table[4 * dim + d] + extra[d]).collect()
1491}
1492
1493fn head(
1496 x: &[f32],
1497 n: usize,
1498 dim: usize,
1499 sst: &[f32],
1500 ts: &TsTable,
1501 proj: &Lin,
1502 pool: Option<&Pool>,
1503) -> Vec<f32> {
1504 let mut y = vec![0f32; n * dim];
1505 let dst = Shared(y.as_mut_ptr());
1506 rows(pool, n, &|s, e| {
1507 let r = unsafe { dst.at(s * dim, (e - s) * dim) };
1508 let mut ln = vec![0f32; dim];
1509 for (row, i) in r.chunks_exact_mut(dim).zip(s..e) {
1510 let emb = ts.emb_row(ts.idx[i.min(ts.idx.len() - 1)]);
1511 layer_norm(&x[i * dim..(i + 1) * dim], &mut ln);
1512 for d in 0..dim {
1513 row[d] = ln[d] * (1.0 + sst[dim + d] + emb[d]) + sst[d] + emb[d];
1514 }
1515 }
1516 });
1517 proj.apply(&y, n, pool)
1518}