1use crate::ltxdit::{LtxDit, StreamInput};
17use crate::pool::Pool;
18
19pub const SCALE_TIME: usize = 8;
21pub const SCALE_SPACE: usize = 32;
22pub const AUDIO_LATENTS_PER_SEC: f64 = 25.0;
24const AUDIO_HOP: f64 = 160.0;
25const AUDIO_RATE: f64 = 16000.0;
26const AUDIO_DOWNSAMPLE: f64 = 4.0;
27
28pub const STAGE1_SIGMAS: [f32; 9] =
30 [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0];
31pub const STAGE2_SIGMAS: [f32; 4] = [0.909375, 0.725, 0.421875, 0.0];
32
33pub struct Rng(u64);
36
37impl Rng {
38 pub fn new(seed: u64) -> Rng {
39 Rng(seed.wrapping_mul(0x9E37_79B9_7F4A_7C15) | 1)
40 }
41 fn next_u64(&mut self) -> u64 {
42 self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
44 let mut z = self.0;
45 z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
46 z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
47 z ^ (z >> 31)
48 }
49 fn unit(&mut self) -> f64 {
50 (self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
51 }
52 pub fn normal(&mut self) -> f32 {
54 let u1 = self.unit().max(1e-12);
55 let u2 = self.unit();
56 ((-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()) as f32
57 }
58 pub fn fill_normal(&mut self, dst: &mut [f32]) {
59 for v in dst.iter_mut() {
60 *v = self.normal();
61 }
62 }
63}
64
65#[derive(Clone, Copy, Debug)]
67pub struct Geometry {
68 pub frames: usize,
69 pub height: usize,
70 pub width: usize,
71 pub fps: f64,
72 pub lf: usize,
74 pub lh: usize,
75 pub lw: usize,
76 pub af: usize,
78}
79
80impl Geometry {
81 pub fn new(frames: usize, height: usize, width: usize, fps: f64) -> Geometry {
82 let lf = (frames - 1) / SCALE_TIME + 1;
83 let duration = frames as f64 / fps;
84 Geometry {
85 frames,
86 height,
87 width,
88 fps,
89 lf,
90 lh: height / SCALE_SPACE,
91 lw: width / SCALE_SPACE,
92 af: (duration * AUDIO_LATENTS_PER_SEC).round() as usize,
93 }
94 }
95
96 pub fn video_tokens(&self) -> usize {
97 self.lf * self.lh * self.lw
98 }
99
100 pub fn tokens_per_frame(&self) -> usize {
101 self.lh * self.lw
102 }
103
104 pub fn video_positions(&self) -> Vec<Vec<f64>> {
107 let causal = |v: f64| (v + 1.0 - SCALE_TIME as f64).max(0.0);
108 let mut out = Vec::with_capacity(self.video_tokens());
109 for f in 0..self.lf {
110 let t0 = causal((f * SCALE_TIME) as f64) / self.fps;
111 let t1 = causal(((f + 1) * SCALE_TIME) as f64) / self.fps;
112 let t = (t0 + t1) / 2.0;
113 for h in 0..self.lh {
114 let y = ((h * SCALE_SPACE) as f64 + ((h + 1) * SCALE_SPACE) as f64) / 2.0;
115 for w in 0..self.lw {
116 let x = ((w * SCALE_SPACE) as f64 + ((w + 1) * SCALE_SPACE) as f64) / 2.0;
117 out.push(vec![t, y, x]);
118 }
119 }
120 }
121 out
122 }
123
124 pub fn audio_positions(&self) -> Vec<Vec<f64>> {
126 let sec = |i: usize| {
127 let mel = (i as f64 * AUDIO_DOWNSAMPLE + 1.0 - AUDIO_DOWNSAMPLE).max(0.0);
128 mel * AUDIO_HOP / AUDIO_RATE
129 };
130 (0..self.af).map(|i| vec![(sec(i) + sec(i + 1)) / 2.0]).collect()
131 }
132
133 pub fn keyframes_mask(&self) -> Vec<f32> {
136 let mut m = vec![0f32; self.video_tokens()];
137 for v in m.iter_mut().take(self.tokens_per_frame()) {
138 *v = 1.0;
139 }
140 m
141 }
142}
143
144pub struct Stage {
146 pub sigmas: Vec<f32>,
147 pub ancestral: bool,
148}
149
150impl Stage {
151 pub fn stage1() -> Stage {
152 Stage { sigmas: STAGE1_SIGMAS.to_vec(), ancestral: true }
153 }
154 pub fn stage2() -> Stage {
155 Stage { sigmas: STAGE2_SIGMAS.to_vec(), ancestral: false }
156 }
157}
158
159fn euler_step(x: &mut [f32], denoised: &[f32], sigma: f32, sigma_next: f32, eta: f32, noise: Option<&[f32]>) {
164 if sigma_next == 0.0 {
165 x.copy_from_slice(denoised);
166 return;
167 }
168 let down_ratio = 1.0 + (sigma_next / sigma - 1.0) * eta;
169 let sigma_down = sigma_next * down_ratio;
170 let r = sigma_down / sigma;
171 for (v, &d) in x.iter_mut().zip(denoised) {
172 *v = r * *v + (1.0 - r) * d;
173 }
174 if eta > 0.0 {
175 let alpha_next = 1.0 - sigma_next;
176 let alpha_down = 1.0 - sigma_down;
177 let coeff = (sigma_next * sigma_next
178 - sigma_down * sigma_down * alpha_next * alpha_next / (alpha_down * alpha_down))
179 .max(0.0)
180 .sqrt();
181 let scale = alpha_next / alpha_down;
182 let n = noise.expect("ancestral step needs noise");
183 for (v, &e) in x.iter_mut().zip(n) {
184 *v = scale * *v + e * coeff;
185 }
186 }
187}
188
189pub struct Latents {
191 pub video: Vec<f32>,
192 pub audio: Vec<f32>,
193}
194
195#[derive(Clone, Default)]
201pub struct Conditioning {
202 pub video_mask: Vec<f32>,
203 pub video_clean: Vec<f32>,
204 pub audio_mask: Vec<f32>,
205 pub audio_clean: Vec<f32>,
206}
207
208impl Conditioning {
209 pub fn video_prefix(geo: &Geometry, clean: &[f32], frames: usize) -> Conditioning {
211 let per = geo.tokens_per_frame();
212 let mut mask = vec![1f32; geo.video_tokens()];
213 let mut full = vec![0f32; geo.video_tokens() * 128];
214 let n = (frames * per).min(geo.video_tokens());
215 for (t, m) in mask.iter_mut().enumerate().take(n) {
216 *m = 0.0;
217 full[t * 128..(t + 1) * 128].copy_from_slice(&clean[t * 128..(t + 1) * 128]);
218 }
219 Conditioning { video_mask: mask, video_clean: full, ..Default::default() }
220 }
221
222 pub fn video_all(geo: &Geometry, clean: &[f32]) -> Conditioning {
225 Conditioning {
226 video_mask: vec![0f32; geo.video_tokens()],
227 video_clean: clean.to_vec(),
228 ..Default::default()
229 }
230 }
231
232 pub fn with_audio_all(mut self, geo: &Geometry, clean: &[f32]) -> Conditioning {
235 self.audio_mask = vec![0f32; geo.af];
236 self.audio_clean = clean.to_vec();
237 self
238 }
239}
240
241pub type Progress<'a> = &'a mut dyn FnMut(usize, usize, f64);
243
244#[allow(clippy::too_many_arguments)]
245pub fn run_stage(
246 dit: &LtxDit,
247 geo: &Geometry,
248 stage: &Stage,
249 video_ctx: &[f32],
250 audio_ctx: &[f32],
251 ctx_len: usize,
252 init: Option<Latents>,
253 rng: &mut Rng,
254 pool: Option<&Pool>,
255 progress: Progress<'_>,
256) -> Latents {
257 run_stage_cond(dit, geo, stage, video_ctx, audio_ctx, ctx_len, init, None, rng, pool, progress)
258}
259
260#[allow(clippy::too_many_arguments)]
261pub fn run_stage_cond(
262 dit: &LtxDit,
263 geo: &Geometry,
264 stage: &Stage,
265 video_ctx: &[f32],
266 audio_ctx: &[f32],
267 ctx_len: usize,
268 init: Option<Latents>,
269 cond: Option<&Conditioning>,
270 rng: &mut Rng,
271 pool: Option<&Pool>,
272 progress: Progress<'_>,
273) -> Latents {
274 let vt = geo.video_tokens();
275 let at = geo.af;
276 let vch = 128usize;
277 let ach = 128usize;
278 let s0 = stage.sigmas[0];
279
280 let mut v = vec![0f32; vt * vch];
284 let mut a = vec![0f32; at * ach];
285 rng.fill_normal(&mut v);
286 rng.fill_normal(&mut a);
287 if let Some(prev) = init {
288 for (x, &p) in v.iter_mut().zip(&prev.video) {
289 *x = p + (*x - p) * s0;
290 }
291 for (x, &p) in a.iter_mut().zip(&prev.audio) {
292 *x = p + (*x - p) * s0;
293 }
294 }
295
296 let vmask: Vec<f32> = cond
299 .map(|c| c.video_mask.clone())
300 .filter(|m| m.len() == vt)
301 .unwrap_or_else(|| vec![1f32; vt]);
302 let amask: Vec<f32> = cond
303 .map(|c| c.audio_mask.clone())
304 .filter(|m| m.len() == at)
305 .unwrap_or_else(|| vec![1f32; at]);
306 let vclean: Vec<f32> = cond
307 .map(|c| c.video_clean.clone())
308 .filter(|c| c.len() == v.len())
309 .unwrap_or_else(|| vec![0f32; v.len()]);
310 let aclean: Vec<f32> = cond
311 .map(|c| c.audio_clean.clone())
312 .filter(|c| c.len() == a.len())
313 .unwrap_or_else(|| vec![0f32; a.len()]);
314 let blend = |x: &mut [f32], clean: &[f32], mask: &[f32], ch: usize| {
315 for (t, &m) in mask.iter().enumerate() {
316 if m >= 1.0 {
317 continue;
318 }
319 for d in 0..ch {
320 let i = t * ch + d;
321 x[i] = clean[i] + (x[i] - clean[i]) * m;
322 }
323 }
324 };
325 blend(&mut v, &vclean, &vmask, vch);
326 blend(&mut a, &aclean, &amask, ach);
327
328 let vpos = geo.video_positions();
329 let apos = geo.audio_positions();
330 let kf = geo.keyframes_mask();
331 let steps = stage.sigmas.len() - 1;
332 let eta = if stage.ancestral { 1.0 } else { 0.0 };
333
334 for i in 0..steps {
335 let t0 = std::time::Instant::now();
336 let sigma = stage.sigmas[i];
337 let sigma_next = stage.sigmas[i + 1];
338 let vin = StreamInput {
339 latent: v.clone(),
340 tokens: vt,
341 timesteps: vmask.iter().map(|m| sigma * m).collect(),
342 positions: vpos.clone(),
343 context: video_ctx.to_vec(),
344 ctx_len,
345 context_mask: Vec::new(),
346 keyframes: kf.clone(),
347 sigma,
348 };
349 let ain = StreamInput {
350 latent: a.clone(),
351 tokens: at,
352 timesteps: amask.iter().map(|m| sigma * m).collect(),
353 positions: apos.clone(),
354 context: audio_ctx.to_vec(),
355 ctx_len,
356 context_mask: Vec::new(),
357 keyframes: Vec::new(),
358 sigma,
359 };
360 let (vv, av) = dit.forward(&vin, &ain, pool);
361 let mut vd: Vec<f32> = v
365 .iter()
366 .zip(&vv)
367 .enumerate()
368 .map(|(i, (&x, &g))| x - g * sigma * vmask[i / vch])
369 .collect();
370 let mut ad: Vec<f32> = a
371 .iter()
372 .zip(&av)
373 .enumerate()
374 .map(|(i, (&x, &g))| x - g * sigma * amask[i / ach])
375 .collect();
376 blend(&mut vd, &vclean, &vmask, vch);
377 blend(&mut ad, &aclean, &amask, ach);
378 let (vn, an) = if eta > 0.0 && sigma_next > 0.0 {
379 let mut vn = vec![0f32; v.len()];
380 let mut an = vec![0f32; a.len()];
381 rng.fill_normal(&mut vn);
382 rng.fill_normal(&mut an);
383 (Some(vn), Some(an))
384 } else {
385 (None, None)
386 };
387 euler_step(&mut v, &vd, sigma, sigma_next, eta, vn.as_deref());
388 euler_step(&mut a, &ad, sigma, sigma_next, eta, an.as_deref());
389 blend(&mut v, &vclean, &vmask, vch);
390 blend(&mut a, &aclean, &amask, ach);
391 progress(i + 1, steps, t0.elapsed().as_secs_f64());
392 }
393 Latents { video: v, audio: a }
394}
395
396pub fn unpatchify_video(tokens: &[f32], geo: &Geometry) -> Vec<f32> {
398 let (lf, lh, lw) = (geo.lf, geo.lh, geo.lw);
399 let c = 128usize;
400 let mut out = vec![0f32; c * lf * lh * lw];
401 for f in 0..lf {
402 for h in 0..lh {
403 for w in 0..lw {
404 let t = (f * lh + h) * lw + w;
405 for ch in 0..c {
406 out[((ch * lf + f) * lh + h) * lw + w] = tokens[t * c + ch];
407 }
408 }
409 }
410 }
411 out
412}
413
414pub fn unpatchify_audio(tokens: &[f32], frames: usize) -> Vec<f32> {
417 let (c, mel) = (8usize, 16usize);
418 let mut out = vec![0f32; c * frames * mel];
419 for t in 0..frames {
420 for ch in 0..c {
421 for m in 0..mel {
422 out[(ch * frames + t) * mel + m] = tokens[t * c * mel + ch * mel + m];
423 }
424 }
425 }
426 out
427}
428
429pub fn patchify_video(vol: &[f32], geo: &Geometry) -> Vec<f32> {
432 let (lf, lh, lw) = (geo.lf, geo.lh, geo.lw);
433 let c = 128usize;
434 let mut out = vec![0f32; c * lf * lh * lw];
435 for f in 0..lf {
436 for h in 0..lh {
437 for w in 0..lw {
438 let t = (f * lh + h) * lw + w;
439 for ch in 0..c {
440 out[t * c + ch] = vol[((ch * lf + f) * lh + h) * lw + w];
441 }
442 }
443 }
444 }
445 out
446}