Skip to main content

cortiq_engine/
ltxpipe.rs

1//! The LTX-2.5 sampler: latent geometry, the noise schedule and the Euler
2//! loops that drive [`crate::ltxdit::LtxDit`] from noise to a clean latent.
3//!
4//! The reference pipeline is two-stage — eight ancestral steps at half
5//! resolution, a latent upsample, then three deterministic steps at full
6//! resolution. Both stages run the same loop; they differ only in the sigma
7//! schedule, whether noise is re-injected, and what the latent starts from.
8//!
9//! Everything positional is derived here, because the DiT reads positions
10//! rather than shapes: video tokens carry `(seconds, pixel row, pixel
11//! column)` patch midpoints — the temporal axis divided by the frame rate so
12//! it shares a unit with audio — and the first latent frame is shifted by the
13//! causal correction, since a causal video encoder gives it one pixel frame
14//! where every later latent frame gets eight.
15
16use crate::ltxdit::{LtxDit, StreamInput};
17use crate::pool::Pool;
18
19/// Video VAE downscaling: 8 frames, 32 rows, 32 columns per latent step.
20pub const SCALE_TIME: usize = 8;
21pub const SCALE_SPACE: usize = 32;
22/// Audio latents per second: 16000 / 160 / 4.
23pub 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
28/// The distilled schedules the release ships with.
29pub 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
33/// Counter-based RNG with a Box-Muller normal — reproducible from a seed,
34/// and independent of any host library.
35pub 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        // splitmix64
43        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    /// One standard normal sample.
53    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/// Latent geometry of one render.
66#[derive(Clone, Copy, Debug)]
67pub struct Geometry {
68    pub frames: usize,
69    pub height: usize,
70    pub width: usize,
71    pub fps: f64,
72    /// Latent frames / rows / columns.
73    pub lf: usize,
74    pub lh: usize,
75    pub lw: usize,
76    /// Audio latent frames.
77    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    /// Patch midpoints in `(seconds, pixel row, pixel column)`, in the
105    /// frame-major order `patchify` produces.
106    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    /// Audio patch midpoints, in seconds.
125    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    /// Non-zero on the first latent frame, whose latent encodes a single
134    /// standalone pixel frame.
135    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
144/// A denoising stage: which schedule, and whether it re-injects noise.
145pub 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
159/// One ancestral Euler step in the rectified-flow parameterization
160/// (`alpha = 1 - sigma`): advance to `sigma_down`, then renoise back up to
161/// `sigma_next` with the variance-preserving rescale. `eta = 0` reduces it to
162/// a plain Euler step and ignores `noise`.
163fn 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
189/// The state a stage carries: patchified latents for both streams.
190pub struct Latents {
191    pub video: Vec<f32>,
192    pub audio: Vec<f32>,
193}
194
195/// What is held fixed while the rest is denoised. `mask[t] = 0` freezes
196/// token `t` at `clean[t]` and hands the transformer a timestep of zero for
197/// it — which is how one encoded image becomes the first frame of a
198/// generated shot, and how a whole encoded clip becomes the picture a
199/// soundtrack is written for.
200#[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    /// Freeze the first `frames` latent frames at `clean` (patchified).
210    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    /// Freeze the whole video stream — the picture is given, the sound is
223    /// what is being generated.
224    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    /// Freeze the whole soundtrack — the sound is given, the picture is what
233    /// is being generated.
234    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
241/// Progress callback: `(step, total, seconds for that step)`.
242pub 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    // A fresh stage starts from pure noise; a refinement stage lerps the
281    // incoming latent toward noise by the first sigma, exactly as the
282    // reference's noiser does.
283    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    // conditioning: a frozen token starts clean, stays clean, and is handed
297    // a timestep of zero so the modulation treats it as already denoised
298    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        // velocity → denoised, at the token's own timestep
362        // velocity → denoised at each token's own timestep, then the frozen
363        // tokens are put back exactly as they were
364        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
396/// Patchified video tokens `[T, 128]` back to a `[128, F, H, W]` volume.
397pub 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
414/// Patchified audio tokens `[T, 128]` back to `[8, T, 16]` (channels, time,
415/// mel bins) — the layout the audio VAE decodes.
416pub 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
429/// A `[128, F, H, W]` volume back to patchified tokens `[T, 128]` — the
430/// inverse of [`unpatchify_video`], for feeding a stage its starting latent.
431pub 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}