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
195pub 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 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 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
285pub 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
303pub 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
318pub 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}