1use crate::ltxdit::{Attn, Lin, Rope, gelu_tanh, rms_plain, rows};
26use crate::pool::Pool;
27use crate::qtensor::QTensor;
28use cortiq_core::CmfModel;
29use std::sync::Arc;
30
31const EPS: f64 = 1e-6;
32
33fn rms(x: &[f32], w: &[f32], dst: &mut [f32]) {
36 let ss = x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / x.len() as f64;
37 let inv = 1.0 / (ss + EPS).sqrt();
38 for ((d, &v), &g) in dst.iter_mut().zip(x).zip(w) {
39 *d = (v as f64 * inv) as f32 * g;
40 }
41}
42
43fn rms_nw(x: &mut [f32]) {
45 let ss = x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / x.len() as f64;
46 let inv = 1.0 / (ss + EPS).sqrt();
47 for v in x.iter_mut() {
48 *v = (*v as f64 * inv) as f32;
49 }
50}
51
52struct Rot {
55 cos: Vec<f32>,
56 sin: Vec<f32>,
57 half: usize,
58}
59
60impl Rot {
61 fn build(seq: usize, head_dim: usize, base: f64, rotary: f64) -> Rot {
66 let half = head_dim / 2;
67 let rope_angles = (rotary * head_dim as f64 / 2.0) as usize;
68 let inv: Vec<f64> = (0..half)
69 .map(|j| {
70 if j < rope_angles {
71 1.0 / base.powf((2 * j) as f64 / head_dim as f64)
72 } else {
73 0.0
74 }
75 })
76 .collect();
77 let mut cos = vec![0f32; seq * half];
78 let mut sin = vec![0f32; seq * half];
79 for p in 0..seq {
80 for (j, &f) in inv.iter().enumerate() {
81 let a = p as f64 * f;
82 cos[p * half + j] = a.cos() as f32;
83 sin[p * half + j] = a.sin() as f32;
84 }
85 }
86 Rot { cos, sin, half }
87 }
88
89 fn apply(&self, t: usize, row: &mut [f32]) {
90 let h = self.half;
91 let (c, s) = (&self.cos[t * h..(t + 1) * h], &self.sin[t * h..(t + 1) * h]);
92 for i in 0..h {
93 let (a, b) = (row[i], row[i + h]);
94 row[i] = a * c[i] - b * s[i];
95 row[i + h] = b * c[i] + a * s[i];
96 }
97 }
98}
99
100struct GemmaLayer {
101 q: Lin,
102 k: Lin,
103 v: Option<Lin>,
104 o: Lin,
105 q_norm: Vec<f32>,
106 k_norm: Vec<f32>,
107 in_norm: Vec<f32>,
108 post_attn_norm: Vec<f32>,
109 pre_ff_norm: Vec<f32>,
110 post_ff_norm: Vec<f32>,
111 gate: Lin,
112 up: Lin,
113 down: Lin,
114 scalar: f32,
115 head_dim: usize,
116 q_heads: usize,
117 kv_heads: usize,
118 sliding: bool,
119}
120
121pub struct LtxTextEncoder {
123 embed: QTensor,
124 mapped_bytes: u64,
125 layers: Vec<GemmaLayer>,
126 norm: Vec<f32>,
127 video_agg: Lin,
128 audio_agg: Lin,
129 v_conn: Connector,
130 a_conn: Connector,
131 hidden: usize,
132 embed_scale: f32,
133 pub max_len: usize,
134 pub bos: u32,
135 pub pad: u32,
136}
137
138struct Connector {
139 blocks: Vec<(Attn, Lin, Lin)>,
140 registers: Vec<f32>,
141 dim: usize,
142 heads: usize,
143 dh: usize,
144 max_pos: f64,
145}
146
147fn vecf(model: &Arc<CmfModel>, name: &str) -> Result<Vec<f32>, String> {
148 crate::dit::cmf_f32(model, name)
149}
150
151impl LtxTextEncoder {
152 pub fn from_cmf(model: &Arc<CmfModel>) -> Result<LtxTextEncoder, String> {
153 let cfg_bytes = ["te.gemma_config_json", "ltx.gemma_config_json"]
154 .iter()
155 .find_map(|n| model.tensor(n).map(|e| model.entry_bytes(e)));
156 let cfg: serde_json::Value = match cfg_bytes {
157 Some(b) => serde_json::from_slice(b).map_err(|e| format!("gemma config: {e}"))?,
158 None => serde_json::Value::Null,
159 };
160 let tc = cfg
161 .get("text_config")
162 .cloned()
163 .unwrap_or(serde_json::Value::Null);
164 let g = |k: &str, d: f64| tc.get(k).and_then(|v| v.as_f64()).unwrap_or(d);
165 let hidden = g("hidden_size", 3840.0) as usize;
166 let n_layers = g("num_hidden_layers", 48.0) as usize;
167 let head_dim = g("head_dim", 256.0) as usize;
168 let global_head_dim = g("global_head_dim", 512.0) as usize;
169 let types: Vec<String> = tc
170 .get("layer_types")
171 .and_then(|v| v.as_array())
172 .map(|a| {
173 a.iter()
174 .map(|s| s.as_str().unwrap_or("sliding_attention").to_string())
175 .collect()
176 })
177 .unwrap_or_else(|| {
178 (0..n_layers)
180 .map(|i| {
181 if i % 6 == 5 {
182 "full_attention".into()
183 } else {
184 "sliding_attention".into()
185 }
186 })
187 .collect()
188 });
189
190 let mut layers = Vec::with_capacity(n_layers);
191 for i in 0..n_layers {
192 let p = format!("te.model.layers.{i}");
193 let sliding = types[i] != "full_attention";
194 let hd = if sliding { head_dim } else { global_head_dim };
195 let q = Lin::load(model, &format!("{p}.self_attn.q_proj"), false)?;
196 let k = Lin::load(model, &format!("{p}.self_attn.k_proj"), false)?;
197 let v = match model.tensor(&format!("{p}.self_attn.v_proj.weight")) {
198 Some(_) => Some(Lin::load(model, &format!("{p}.self_attn.v_proj"), false)?),
199 None => None,
200 };
201 let q_rows = model
202 .tensor(&format!("{p}.self_attn.q_proj.weight"))
203 .ok_or_else(|| format!("missing {p}.self_attn.q_proj.weight"))?
204 .shape[0];
205 let k_rows = model
206 .tensor(&format!("{p}.self_attn.k_proj.weight"))
207 .ok_or("missing k_proj")?
208 .shape[0];
209 layers.push(GemmaLayer {
210 q,
211 k,
212 v,
213 o: Lin::load(model, &format!("{p}.self_attn.o_proj"), false)?,
214 q_norm: vecf(model, &format!("{p}.self_attn.q_norm.weight"))?,
215 k_norm: vecf(model, &format!("{p}.self_attn.k_norm.weight"))?,
216 in_norm: vecf(model, &format!("{p}.input_layernorm.weight"))?,
217 post_attn_norm: vecf(model, &format!("{p}.post_attention_layernorm.weight"))?,
218 pre_ff_norm: vecf(model, &format!("{p}.pre_feedforward_layernorm.weight"))?,
219 post_ff_norm: vecf(model, &format!("{p}.post_feedforward_layernorm.weight"))?,
220 gate: Lin::load(model, &format!("{p}.mlp.gate_proj"), false)?,
221 up: Lin::load(model, &format!("{p}.mlp.up_proj"), false)?,
222 down: Lin::load(model, &format!("{p}.mlp.down_proj"), false)?,
223 scalar: vecf(model, &format!("{p}.layer_scalar"))?[0],
224 head_dim: hd,
225 q_heads: q_rows / hd,
226 kv_heads: k_rows / hd,
227 sliding,
228 });
229 }
230
231 let conn =
232 |prefix: &str, dim: usize, heads: usize, dh: usize| -> Result<Connector, String> {
233 let mut blocks = Vec::new();
234 let mut i = 0usize;
235 while model
236 .tensor(&format!(
237 "{prefix}.transformer_1d_blocks.{i}.attn1.to_q.weight"
238 ))
239 .is_some()
240 {
241 let p = format!("{prefix}.transformer_1d_blocks.{i}");
242 blocks.push((
243 Attn::load(model, &format!("{p}.attn1"), heads, dh)?,
244 Lin::load(model, &format!("{p}.ff.net.0.proj"), true)?,
245 Lin::load(model, &format!("{p}.ff.net.2"), true)?,
246 ));
247 i += 1;
248 }
249 Ok(Connector {
250 blocks,
251 registers: vecf(model, &format!("{prefix}.learnable_registers"))?,
252 dim,
253 heads,
254 dh,
255 max_pos: 4096.0,
256 })
257 };
258
259 Ok(LtxTextEncoder {
260 embed: QTensor::from_model(model, "te.model.embed_tokens.weight")?,
261 mapped_bytes: model.primary_bytes().len() as u64,
262 layers,
263 norm: vecf(model, "te.model.norm.weight")?,
264 video_agg: Lin::load(
265 model,
266 "te.text_embedding_projection.video_aggregate_embed",
267 true,
268 )?,
269 audio_agg: Lin::load(
270 model,
271 "te.text_embedding_projection.audio_aggregate_embed",
272 true,
273 )?,
274 v_conn: conn("dit.video_embeddings_connector", 4096, 32, 128)?,
275 a_conn: conn("dit.audio_embeddings_connector", 2048, 32, 64)?,
276 hidden,
277 embed_scale: (hidden as f64).sqrt() as f32,
278 max_len: 1024,
279 bos: tc.get("bos_token_id").and_then(|v| v.as_u64()).unwrap_or(2) as u32,
280 pad: tc.get("pad_token_id").and_then(|v| v.as_u64()).unwrap_or(0) as u32,
281 })
282 }
283
284 pub fn pad_ids(&self, ids: &[u32]) -> (Vec<u32>, Vec<f32>) {
287 let mut body = Vec::with_capacity(self.max_len);
288 if ids.first() != Some(&self.bos) {
289 body.push(self.bos);
290 }
291 body.extend_from_slice(ids);
292 body.truncate(self.max_len);
293 let padlen = self.max_len - body.len();
294 let mut out = vec![self.pad; padlen];
295 out.extend_from_slice(&body);
296 let mut mask = vec![0f32; padlen];
297 mask.extend(std::iter::repeat_n(1f32, body.len()));
298 (out, mask)
299 }
300
301 pub fn hidden_states(&self, ids: &[u32], mask: &[f32], pool: Option<&Pool>) -> Vec<Vec<f32>> {
304 let _pause = self.crowds_the_device().then(crate::gpu::pause_gpu);
311 self.hidden_states_inner(ids, mask, pool)
312 }
313
314 #[cfg(target_os = "macos")]
315 fn crowds_the_device(&self) -> bool {
316 let one_window = crate::gpu_metal::max_buffer_bytes();
322 let wired = crate::gpu_metal::working_set_bytes();
323 (one_window > 0 && self.mapped_bytes > one_window)
324 || (wired > 0 && self.mapped_bytes > wired)
325 }
326
327 #[cfg(not(target_os = "macos"))]
328 fn crowds_the_device(&self) -> bool {
329 false
330 }
331
332 fn hidden_states_inner(&self, ids: &[u32], mask: &[f32], pool: Option<&Pool>) -> Vec<Vec<f32>> {
333 let t = ids.len();
334 let d = self.hidden;
335 let mut x = vec![0f32; t * d];
336 for (i, &id) in ids.iter().enumerate() {
337 self.embed.row_f32(id as usize, &mut x[i * d..(i + 1) * d]);
338 for v in x[i * d..(i + 1) * d].iter_mut() {
339 *v *= self.embed_scale;
340 }
341 }
342 let mut out = vec![x.clone()];
343 let rot_slide = Rot::build(t, self.layers[0].head_dim, 10000.0, 1.0);
345 let full = self.layers.iter().find(|l| !l.sliding);
346 let rot_full = full.map(|l| Rot::build(t, l.head_dim, 1_000_000.0, 0.25));
347
348 for layer in &self.layers {
349 let rot = if layer.sliding {
350 &rot_slide
351 } else {
352 rot_full.as_ref().unwrap()
353 };
354 let mut h = vec![0f32; t * d];
355 for i in 0..t {
356 rms(
357 &x[i * d..(i + 1) * d],
358 &layer.in_norm,
359 &mut h[i * d..(i + 1) * d],
360 );
361 }
362 let attn = self.attention(layer, &h, t, mask, rot, pool);
363 for i in 0..t {
364 let mut n = vec![0f32; d];
365 rms(&attn[i * d..(i + 1) * d], &layer.post_attn_norm, &mut n);
366 for (v, &a) in x[i * d..(i + 1) * d].iter_mut().zip(&n) {
367 *v += a;
368 }
369 }
370 let mut h2 = vec![0f32; t * d];
371 for i in 0..t {
372 rms(
373 &x[i * d..(i + 1) * d],
374 &layer.pre_ff_norm,
375 &mut h2[i * d..(i + 1) * d],
376 );
377 }
378 let mut g = layer.gate.apply(&h2, t, pool);
379 let u = layer.up.apply(&h2, t, pool);
380 crate::ltxdit::gelu_tanh_rows(&mut g, pool);
381 for (a, &b) in g.iter_mut().zip(&u) {
382 *a *= b;
383 }
384 let ff = layer.down.apply(&g, t, pool);
385 for i in 0..t {
386 let mut n = vec![0f32; d];
387 rms(&ff[i * d..(i + 1) * d], &layer.post_ff_norm, &mut n);
388 for (v, &a) in x[i * d..(i + 1) * d].iter_mut().zip(&n) {
389 *v += a;
390 }
391 }
392 for v in x.iter_mut() {
393 *v *= layer.scalar;
394 }
395 out.push(x.clone());
396 }
397 if let Some(last) = out.last_mut() {
403 let mut n = vec![0f32; t * d];
404 for i in 0..t {
405 rms(
406 &last[i * d..(i + 1) * d],
407 &self.norm,
408 &mut n[i * d..(i + 1) * d],
409 );
410 }
411 *last = n;
412 }
413 out
414 }
415
416 fn attention(
417 &self,
418 l: &GemmaLayer,
419 h: &[f32],
420 t: usize,
421 mask: &[f32],
422 rot: &Rot,
423 pool: Option<&Pool>,
424 ) -> Vec<f32> {
425 let hd = l.head_dim;
426 let qi = l.q_heads * hd;
427 let ki = l.kv_heads * hd;
428 let mut q = l.q.apply(h, t, pool);
429 let mut k = l.k.apply(h, t, pool);
430 let mut v = match &l.v {
431 Some(p) => p.apply(h, t, pool),
432 None => k.clone(),
433 };
434 for i in 0..t {
435 for hh in 0..l.q_heads {
436 let r = &mut q[i * qi + hh * hd..i * qi + (hh + 1) * hd];
437 let mut n = vec![0f32; hd];
438 rms(r, &l.q_norm, &mut n);
439 r.copy_from_slice(&n);
440 rot.apply(i, r);
441 }
442 for hh in 0..l.kv_heads {
443 let r = &mut k[i * ki + hh * hd..i * ki + (hh + 1) * hd];
444 let mut n = vec![0f32; hd];
445 rms(r, &l.k_norm, &mut n);
446 r.copy_from_slice(&n);
447 rot.apply(i, r);
448 rms_nw(&mut v[i * ki + hh * hd..i * ki + (hh + 1) * hd]);
449 }
450 }
451 let window = if l.sliding { 1024usize } else { usize::MAX };
455 let mut bias = vec![0f32; t * t];
456 for i in 0..t {
457 for j in 0..t {
458 let blocked = j > i || (window != usize::MAX && i - j >= window) || mask[j] == 0.0;
459 bias[i * t + j] = if blocked { f32::NEG_INFINITY } else { 0.0 };
460 }
461 }
462 let dead: Vec<bool> = (0..t).map(|i| mask[i] == 0.0).collect();
468 let mut out = vec![0f32; t * qi];
469 let mut qh = vec![0f32; t * hd];
470 let mut kh = vec![0f32; t * hd];
471 let mut vh = vec![0f32; t * hd];
472 let mut sc = vec![0f32; t * t];
473 let mut oh = vec![0f32; t * hd];
474 for hh in 0..l.q_heads {
475 let kv = hh * l.kv_heads / l.q_heads;
476 for i in 0..t {
477 qh[i * hd..(i + 1) * hd].copy_from_slice(&q[i * qi + hh * hd..][..hd]);
478 kh[i * hd..(i + 1) * hd].copy_from_slice(&k[i * ki + kv * hd..][..hd]);
479 vh[i * hd..(i + 1) * hd].copy_from_slice(&v[i * ki + kv * hd..][..hd]);
480 }
481 crate::fcd_ops::gemm_nt(&qh, &kh, &mut sc, t, hd, t, pool);
482 let sp = crate::ltxdit::Shared(sc.as_mut_ptr());
483 rows(pool, t, &|s, e| {
484 let r = unsafe { sp.at(s * t, (e - s) * t) };
485 for (row, i) in r.chunks_exact_mut(t).zip(s..e) {
486 if dead[i] {
487 row.iter_mut().for_each(|x| *x = 0.0);
488 continue;
489 }
490 for (x, b) in row.iter_mut().zip(&bias[i * t..(i + 1) * t]) {
491 *x += *b;
492 }
493 crate::ltxdit::softmax(row);
494 }
495 });
496 oh.iter_mut().for_each(|x| *x = 0.0);
497 crate::fcd_ops::gemm_dx(&sc, &vh, &mut oh, t, hd, t, pool);
498 for i in 0..t {
499 out[i * qi + hh * hd..i * qi + (hh + 1) * hd]
500 .copy_from_slice(&oh[i * hd..(i + 1) * hd]);
501 }
502 }
503 l.o.apply(&out, t, pool)
504 }
505
506 pub fn encode_ids(
508 &self,
509 ids: &[u32],
510 mask: &[f32],
511 pool: Option<&Pool>,
512 ) -> (Vec<f32>, Vec<f32>, usize) {
513 let _pause = self.crowds_the_device().then(crate::gpu::pause_gpu);
518 let hs = self.hidden_states_inner(ids, mask, pool);
519 let (t, d, l) = (ids.len(), self.hidden, hs.len());
520 let mut feats = vec![0f32; t * d * l];
523 for (li, layer) in hs.iter().enumerate() {
524 for i in 0..t {
525 let row = &layer[i * d..(i + 1) * d];
526 let var = row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / d as f64;
527 let inv = 1.0 / (var + 1e-6).sqrt();
528 let keep = mask[i] != 0.0;
529 for j in 0..d {
530 feats[i * d * l + j * l + li] = if keep {
531 (row[j] as f64 * inv) as f32
532 } else {
533 0.0
534 };
535 }
536 }
537 }
538 let project = |agg: &Lin, out_dim: usize| -> Vec<f32> {
543 let scale = ((out_dim as f64) / (d as f64)).sqrt() as f32;
544 let scaled: Vec<f32> = feats.iter().map(|&v| v * scale).collect();
545 crate::gpu::cpu_scope(|| agg.apply(&scaled, t, pool))
546 };
547 let vfeat = project(&self.video_agg, 4096);
548 let afeat = project(&self.audio_agg, 2048);
549 let order: Vec<usize> = (0..t)
551 .filter(|&i| mask[i] != 0.0)
552 .chain((0..t).filter(|&i| mask[i] == 0.0))
553 .collect();
554 let valid = mask.iter().filter(|&&m| m != 0.0).count();
555 let reorder = |x: &[f32], dim: usize| -> Vec<f32> {
556 let mut o = vec![0f32; t * dim];
557 for (new, &old) in order.iter().enumerate() {
558 o[new * dim..(new + 1) * dim].copy_from_slice(&x[old * dim..(old + 1) * dim]);
559 }
560 o
561 };
562 let v = self.v_conn.run(&reorder(&vfeat, 4096), t, valid, pool);
563 let a = self.a_conn.run(&reorder(&afeat, 2048), t, valid, pool);
564 (v, a, t)
565 }
566}
567
568impl Connector {
569 fn run(&self, x: &[f32], t: usize, valid: usize, pool: Option<&Pool>) -> Vec<f32> {
574 let d = self.dim;
575 let regs = self.registers.len() / d;
576 let mut h = x.to_vec();
577 for i in valid..t {
578 let r = i % regs;
579 h[i * d..(i + 1) * d].copy_from_slice(&self.registers[r * d..(r + 1) * d]);
580 }
581 let pos: Vec<Vec<f64>> = (0..t).map(|i| vec![i as f64]).collect();
582 let pe = Rope::build(&pos, &[self.max_pos], d, self.heads, 10000.0);
583 for (attn, ff_in, ff_out) in &self.blocks {
584 let mut n = vec![0f32; t * d];
585 for i in 0..t {
586 rms_plain(&h[i * d..(i + 1) * d], &mut n[i * d..(i + 1) * d]);
587 }
588 let a = attn.forward(&n, t, &n, t, Some(&pe), Some(&pe), None, pool);
589 for (v, &y) in h.iter_mut().zip(&a) {
590 *v += y;
591 }
592 let mut n2 = vec![0f32; t * d];
593 for i in 0..t {
594 rms_plain(&h[i * d..(i + 1) * d], &mut n2[i * d..(i + 1) * d]);
595 }
596 let mut g = ff_in.apply(&n2, t, pool);
597 crate::ltxdit::gelu_tanh_rows(&mut g, pool);
598 let f = ff_out.apply(&g, t, pool);
599 for (v, &y) in h.iter_mut().zip(&f) {
600 *v += y;
601 }
602 }
603 let mut out = vec![0f32; t * d];
604 for i in 0..t {
605 rms_plain(&h[i * d..(i + 1) * d], &mut out[i * d..(i + 1) * d]);
606 }
607 let _ = self.dh;
608 out
609 }
610}