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/// Progress callback: `(step, total, seconds for that step)`.
196pub type Progress<'a> = &'a mut dyn FnMut(usize, usize, f64);
197
198#[allow(clippy::too_many_arguments)]
199pub fn run_stage(
200    dit: &LtxDit,
201    geo: &Geometry,
202    stage: &Stage,
203    video_ctx: &[f32],
204    audio_ctx: &[f32],
205    ctx_len: usize,
206    init: Option<Latents>,
207    rng: &mut Rng,
208    pool: Option<&Pool>,
209    progress: Progress<'_>,
210) -> Latents {
211    let vt = geo.video_tokens();
212    let at = geo.af;
213    let vch = 128usize;
214    let ach = 128usize;
215    let s0 = stage.sigmas[0];
216
217    // A fresh stage starts from pure noise; a refinement stage lerps the
218    // incoming latent toward noise by the first sigma, exactly as the
219    // reference's noiser does.
220    let mut v = vec![0f32; vt * vch];
221    let mut a = vec![0f32; at * ach];
222    rng.fill_normal(&mut v);
223    rng.fill_normal(&mut a);
224    if let Some(prev) = init {
225        for (x, &p) in v.iter_mut().zip(&prev.video) {
226            *x = p + (*x - p) * s0;
227        }
228        for (x, &p) in a.iter_mut().zip(&prev.audio) {
229            *x = p + (*x - p) * s0;
230        }
231    }
232
233    let vpos = geo.video_positions();
234    let apos = geo.audio_positions();
235    let kf = geo.keyframes_mask();
236    let steps = stage.sigmas.len() - 1;
237    let eta = if stage.ancestral { 1.0 } else { 0.0 };
238
239    for i in 0..steps {
240        let t0 = std::time::Instant::now();
241        let sigma = stage.sigmas[i];
242        let sigma_next = stage.sigmas[i + 1];
243        let vin = StreamInput {
244            latent: v.clone(),
245            tokens: vt,
246            timesteps: vec![sigma; vt],
247            positions: vpos.clone(),
248            context: video_ctx.to_vec(),
249            ctx_len,
250            context_mask: Vec::new(),
251            keyframes: kf.clone(),
252            sigma,
253        };
254        let ain = StreamInput {
255            latent: a.clone(),
256            tokens: at,
257            timesteps: vec![sigma; at],
258            positions: apos.clone(),
259            context: audio_ctx.to_vec(),
260            ctx_len,
261            context_mask: Vec::new(),
262            keyframes: Vec::new(),
263            sigma,
264        };
265        let (vv, av) = dit.forward(&vin, &ain, pool);
266        // velocity → denoised, at the token's own timestep
267        let vd: Vec<f32> = v.iter().zip(&vv).map(|(&x, &g)| x - g * sigma).collect();
268        let ad: Vec<f32> = a.iter().zip(&av).map(|(&x, &g)| x - g * sigma).collect();
269        let (vn, an) = if eta > 0.0 && sigma_next > 0.0 {
270            let mut vn = vec![0f32; v.len()];
271            let mut an = vec![0f32; a.len()];
272            rng.fill_normal(&mut vn);
273            rng.fill_normal(&mut an);
274            (Some(vn), Some(an))
275        } else {
276            (None, None)
277        };
278        euler_step(&mut v, &vd, sigma, sigma_next, eta, vn.as_deref());
279        euler_step(&mut a, &ad, sigma, sigma_next, eta, an.as_deref());
280        progress(i + 1, steps, t0.elapsed().as_secs_f64());
281    }
282    Latents { video: v, audio: a }
283}
284
285/// Patchified video tokens `[T, 128]` back to a `[128, F, H, W]` volume.
286pub fn unpatchify_video(tokens: &[f32], geo: &Geometry) -> Vec<f32> {
287    let (lf, lh, lw) = (geo.lf, geo.lh, geo.lw);
288    let c = 128usize;
289    let mut out = vec![0f32; c * lf * lh * lw];
290    for f in 0..lf {
291        for h in 0..lh {
292            for w in 0..lw {
293                let t = (f * lh + h) * lw + w;
294                for ch in 0..c {
295                    out[((ch * lf + f) * lh + h) * lw + w] = tokens[t * c + ch];
296                }
297            }
298        }
299    }
300    out
301}
302
303/// Patchified audio tokens `[T, 128]` back to `[8, T, 16]` (channels, time,
304/// mel bins) — the layout the audio VAE decodes.
305pub fn unpatchify_audio(tokens: &[f32], frames: usize) -> Vec<f32> {
306    let (c, mel) = (8usize, 16usize);
307    let mut out = vec![0f32; c * frames * mel];
308    for t in 0..frames {
309        for ch in 0..c {
310            for m in 0..mel {
311                out[(ch * frames + t) * mel + m] = tokens[t * c * mel + ch * mel + m];
312            }
313        }
314    }
315    out
316}
317
318/// A `[128, F, H, W]` volume back to patchified tokens `[T, 128]` — the
319/// inverse of [`unpatchify_video`], for feeding a stage its starting latent.
320pub fn patchify_video(vol: &[f32], geo: &Geometry) -> Vec<f32> {
321    let (lf, lh, lw) = (geo.lf, geo.lh, geo.lw);
322    let c = 128usize;
323    let mut out = vec![0f32; c * lf * lh * lw];
324    for f in 0..lf {
325        for h in 0..lh {
326            for w in 0..lw {
327                let t = (f * lh + h) * lw + w;
328                for ch in 0..c {
329                    out[t * c + ch] = vol[((ch * lf + f) * lh + h) * lw + w];
330                }
331            }
332        }
333    }
334    out
335}