| 1 | """Standalone paired audio-video memory helpers for multi-shot inference.""" |
| 2 | |
| 3 | from __future__ import annotations |
| 4 | |
| 5 | import json |
| 6 | import random |
| 7 | from dataclasses import dataclass, field |
| 8 | from pathlib import Path |
| 9 | from typing import Any, Optional |
| 10 | |
| 11 | import torch |
| 12 | from PIL import Image |
| 13 | |
| 14 | from ltx_core.model.audio_vae import AudioProcessor |
| 15 | from ltx_core.types import Audio |
| 16 | from ltx_distillation.audio_memory import ( |
| 17 | latent_window_size_to_pixel_window_size, |
| 18 | mel_window_bounds_to_seconds, |
| 19 | select_audio_window_with_bounds, |
| 20 | select_video_frame_indices_from_time_range, |
| 21 | ) |
| 22 | |
| 23 | def prompt_payload_to_text(payload: Any, prompt_max_chars: Optional[int] = None) -> str: |
| 24 | if isinstance(payload, str): |
| 25 | text = payload.strip() |
| 26 | else: |
| 27 | raise TypeError( |
| 28 | f"Unsupported prompt payload type: {type(payload).__name__}. " |
| 29 | "Each shot must be a single concatenated prompt string." |
| 30 | ) |
| 31 | |
| 32 | if prompt_max_chars and len(text) > prompt_max_chars: |
| 33 | text = text[:prompt_max_chars] |
| 34 | return text |
| 35 | |
| 36 | |
| 37 | def json_to_prompts(data: dict[str, Any], prompt_max_chars: Optional[int] = None) -> list[str]: |
| 38 | if isinstance(data.get("prompts"), list): |
| 39 | prompts = [prompt_payload_to_text(item, prompt_max_chars) for item in data["prompts"]] |
| 40 | return [prompt for prompt in prompts if prompt] |
| 41 | |
| 42 | shots = data.get("shots", []) |
| 43 | if not isinstance(shots, list): |
| 44 | return [] |
| 45 | prompts = [prompt_payload_to_text(item, prompt_max_chars) for item in shots] |
| 46 | return [prompt for prompt in prompts if prompt] |
| 47 | |
| 48 | |
| 49 | def load_multishot_prompts(prompts_file: str | Path, prompt_max_chars: Optional[int] = None) -> list[str]: |
| 50 | prompts_path = Path(prompts_file) |
| 51 | with prompts_path.open("r", encoding="utf-8") as handle: |
| 52 | payload = json.load(handle) |
| 53 | prompts = json_to_prompts(payload, prompt_max_chars=prompt_max_chars) |
| 54 | if not prompts: |
| 55 | raise ValueError(f"No prompts found in multishot prompts file: {prompts_path}") |
| 56 | return prompts |
| 57 | |
| 58 | |
| 59 | def normalize_audio_waveform_for_media(audio_waveform: Optional[torch.Tensor]) -> Optional[torch.Tensor]: |
| 60 | if audio_waveform is None: |
| 61 | return None |
| 62 | |
| 63 | waveform = getattr(audio_waveform, "waveform", audio_waveform) |
| 64 | waveform = torch.as_tensor(waveform).detach().cpu().float() |
| 65 | |
| 66 | if waveform.ndim == 3: |
| 67 | if waveform.shape[0] != 1: |
| 68 | raise ValueError(f"Expected batch size 1 for decoded audio, got shape={tuple(waveform.shape)}") |
| 69 | waveform = waveform[0] |
| 70 | if waveform.ndim == 1: |
| 71 | waveform = waveform.unsqueeze(0) |
| 72 | elif waveform.ndim == 2 and waveform.shape[0] not in {1, 2} and waveform.shape[1] in {1, 2}: |
| 73 | waveform = waveform.transpose(0, 1) |
| 74 | elif waveform.ndim != 2: |
| 75 | raise ValueError(f"Expected decoded audio with 1, 2, or 3 dims, got shape={tuple(waveform.shape)}") |
| 76 | |
| 77 | if waveform.shape[0] == 1: |
| 78 | waveform = waveform.repeat(2, 1) |
| 79 | elif waveform.shape[0] > 2: |
| 80 | waveform = waveform[:2] |
| 81 | return waveform.contiguous() |
| 82 | |
| 83 | |
| 84 | def audio_waveform_stats(audio_waveform: Optional[torch.Tensor]) -> dict[str, Any]: |
| 85 | waveform = normalize_audio_waveform_for_media(audio_waveform) |
| 86 | if waveform is None: |
| 87 | return { |
| 88 | "present": False, |
| 89 | "shape": None, |
| 90 | "num_samples": 0, |
| 91 | "rms": 0.0, |
| 92 | "peak": 0.0, |
| 93 | "mean_abs": 0.0, |
| 94 | } |
| 95 | waveform_f = waveform.float() |
| 96 | return { |
| 97 | "present": True, |
| 98 | "shape": list(waveform.shape), |
| 99 | "num_samples": int(waveform.shape[-1]), |
| 100 | "rms": float(torch.sqrt(torch.mean(waveform_f.square())).item()), |
| 101 | "peak": float(torch.max(torch.abs(waveform_f)).item()), |
| 102 | "mean_abs": float(torch.mean(torch.abs(waveform_f)).item()), |
| 103 | } |
| 104 | |
| 105 | |
| 106 | def build_paired_audio_memory_kwargs( |
| 107 | memory_bank: "PairedAudioVideoMemoryBank", |
| 108 | *, |
| 109 | enable_audio_memory: bool, |
| 110 | v2a_grad_scale: float = 1.0, |
| 111 | memory_position_mode: str = "reference", |
| 112 | ) -> dict[str, Any]: |
| 113 | if not enable_audio_memory: |
| 114 | return {} |
| 115 | |
| 116 | memory_audio = memory_bank.get_memory_audio() |
| 117 | if memory_audio is None: |
| 118 | raise RuntimeError("audio memory was requested but the memory bank contains entries without audio latents") |
| 119 | |
| 120 | kwargs: dict[str, Any] = { |
| 121 | "memory_audio": memory_audio, |
| 122 | "memory_audio_timestep": torch.zeros(memory_audio.shape[:2], dtype=torch.float32), |
| 123 | "memory_audio_segment_lengths": memory_bank.get_memory_audio_segment_lengths(), |
| 124 | "v2a_grad_scale": float(v2a_grad_scale), |
| 125 | "memory_position_mode": str(memory_position_mode), |
| 126 | "paired_audio_memory": True, |
| 127 | } |
| 128 | return kwargs |
| 129 | |
| 130 | |
| 131 | @dataclass |
| 132 | class MemoryEntry: |
| 133 | frame: Image.Image | list[Image.Image] |
| 134 | audio_latent: Optional[torch.Tensor] = None |
| 135 | metadata: dict[str, Any] = field(default_factory=dict) |
| 136 | |
| 137 | |
| 138 | class PairedAudioVideoMemoryBank: |
| 139 | def __init__(self, max_size: int, save_mode: str, num_fix_frames: int = 0) -> None: |
| 140 | self.max_size = int(max_size) |
| 141 | self.save_mode = str(save_mode) |
| 142 | self.num_fix_frames = max(0, int(num_fix_frames)) |
| 143 | self.memory: list[MemoryEntry] = [] |
| 144 | |
| 145 | @staticmethod |
| 146 | def _prepare_audio_latent(audio_latent: Optional[torch.Tensor]) -> Optional[torch.Tensor]: |
| 147 | if audio_latent is None: |
| 148 | return None |
| 149 | if audio_latent.dim() != 3: |
| 150 | raise ValueError(f"Expected audio_latent shape [B, T, C], got shape={tuple(audio_latent.shape)}") |
| 151 | return audio_latent.detach().cpu().contiguous() |
| 152 | |
| 153 | @staticmethod |
| 154 | def _normalize_waveform_channels(waveform: torch.Tensor, target_channels: int = 2) -> torch.Tensor: |
| 155 | waveform = torch.as_tensor(waveform).detach().cpu().float() |
| 156 | if waveform.ndim == 3: |
| 157 | if waveform.shape[0] != 1: |
| 158 | raise ValueError(f"Expected batch size 1 for waveform, got shape={tuple(waveform.shape)}") |
| 159 | waveform = waveform[0] |
| 160 | if waveform.ndim == 1: |
| 161 | waveform = waveform.unsqueeze(0) |
| 162 | if waveform.ndim != 2: |
| 163 | raise ValueError(f"Expected waveform [C, T], got shape={tuple(waveform.shape)}") |
| 164 | if waveform.shape[0] == target_channels: |
| 165 | return waveform.contiguous() |
| 166 | if waveform.shape[0] == 1: |
| 167 | return waveform.repeat(target_channels, 1).contiguous() |
| 168 | if waveform.shape[0] > target_channels: |
| 169 | return waveform[:target_channels].contiguous() |
| 170 | pad = target_channels - waveform.shape[0] |
| 171 | return torch.cat([waveform, waveform[-1:].repeat(pad, 1)], dim=0).contiguous() |
| 172 | |
| 173 | @staticmethod |
| 174 | def _select_audio_window(audio_latent: torch.Tensor, window_size: int) -> tuple[torch.Tensor, dict[str, Any]]: |
| 175 | total_frames = int(audio_latent.shape[1]) |
| 176 | window_size = max(1, int(window_size)) |
| 177 | window_len = min(total_frames, window_size) |
| 178 | window_start = max((total_frames - window_len) // 2, 0) |
| 179 | window_end = window_start + window_len |
| 180 | metadata = { |
| 181 | "audio_window_start": int(window_start), |
| 182 | "audio_window_end": int(window_end), |
| 183 | "audio_window_length": int(window_len), |
| 184 | "audio_total_frames": int(total_frames), |
| 185 | } |
| 186 | return audio_latent[:, window_start:window_end].contiguous(), metadata |
| 187 | |
| 188 | @staticmethod |
| 189 | def _select_audio_window_from_waveform( |
| 190 | audio_latent: torch.Tensor, |
| 191 | *, |
| 192 | audio_waveform: torch.Tensor, |
| 193 | audio_sample_rate: int, |
| 194 | window_size: int, |
| 195 | selection_mode: str, |
| 196 | mel_bins: int, |
| 197 | mel_hop_length: int, |
| 198 | n_fft: int, |
| 199 | downsample_factor: int, |
| 200 | is_causal: bool, |
| 201 | ) -> tuple[torch.Tensor, dict[str, Any]]: |
| 202 | waveform = PairedAudioVideoMemoryBank._normalize_waveform_channels(audio_waveform, target_channels=2) |
| 203 | processor = AudioProcessor( |
| 204 | target_sample_rate=int(audio_sample_rate), |
| 205 | mel_bins=int(mel_bins), |
| 206 | mel_hop_length=int(mel_hop_length), |
| 207 | n_fft=int(n_fft), |
| 208 | ) |
| 209 | mel_spectrogram = processor.waveform_to_mel( |
| 210 | Audio(waveform=waveform.unsqueeze(0), sampling_rate=int(audio_sample_rate)) |
| 211 | ) |
| 212 | pixel_window_size = latent_window_size_to_pixel_window_size( |
| 213 | int(window_size), |
| 214 | downsample_factor=int(downsample_factor), |
| 215 | is_causal=bool(is_causal), |
| 216 | ) |
| 217 | _, window_start_indices, window_end_indices = select_audio_window_with_bounds( |
| 218 | mel_spectrogram, |
| 219 | pixel_window_size, |
| 220 | mode=str(selection_mode).lower(), |
| 221 | ) |
| 222 | mel_start = int(window_start_indices[0].item()) |
| 223 | mel_end = int(window_end_indices[0].item()) |
| 224 | start_time_sec, end_time_sec = mel_window_bounds_to_seconds( |
| 225 | mel_start, |
| 226 | mel_end, |
| 227 | hop_length=int(mel_hop_length), |
| 228 | sample_rate=int(audio_sample_rate), |
| 229 | ) |
| 230 | |
| 231 | total_frames = int(audio_latent.shape[1]) |
| 232 | window_len = min(total_frames, max(1, int(window_size))) |
| 233 | duration_sec = max(float(waveform.shape[-1]) / float(audio_sample_rate), 1e-6) |
| 234 | center_time_sec = max(0.0, min(0.5 * (start_time_sec + end_time_sec), duration_sec)) |
| 235 | center_latent = int(round(center_time_sec / duration_sec * float(max(total_frames - 1, 0)))) |
| 236 | window_start = max(0, min(center_latent - window_len // 2, max(total_frames - window_len, 0))) |
| 237 | window_end = window_start + window_len |
| 238 | metadata = { |
| 239 | "audio_window_selection_mode": str(selection_mode).lower(), |
| 240 | "audio_window_start": int(window_start), |
| 241 | "audio_window_end": int(window_end), |
| 242 | "audio_window_length": int(window_len), |
| 243 | "audio_total_frames": int(total_frames), |
| 244 | "mel_window_start": int(mel_start), |
| 245 | "mel_window_end": int(mel_end), |
| 246 | "audio_window_start_time_sec": float(start_time_sec), |
| 247 | "audio_window_end_time_sec": float(end_time_sec), |
| 248 | } |
| 249 | return audio_latent[:, window_start:window_end].contiguous(), metadata |
| 250 | |
| 251 | @staticmethod |
| 252 | def _select_video_clip_around_frame( |
| 253 | frames: list[Image.Image], |
| 254 | *, |
| 255 | center_frame: int, |
| 256 | video_clip_num_frames: int, |
| 257 | ) -> tuple[list[Image.Image], dict[str, Any]]: |
| 258 | video_clip_num_frames = max(1, int(video_clip_num_frames)) |
| 259 | center_frame = max(0, min(int(center_frame), len(frames) - 1)) |
| 260 | left_context = (video_clip_num_frames - 1) // 2 |
| 261 | clip_start = max(0, min(center_frame - left_context, max(len(frames) - video_clip_num_frames, 0))) |
| 262 | clip_end = min(clip_start + video_clip_num_frames, len(frames)) |
| 263 | clip = list(frames[clip_start:clip_end]) |
| 264 | if clip and len(clip) < video_clip_num_frames: |
| 265 | clip.extend([clip[-1]] * (video_clip_num_frames - len(clip))) |
| 266 | metadata = { |
| 267 | "video_clip_start": int(clip_start), |
| 268 | "video_clip_end": int(clip_end), |
| 269 | "video_clip_length": int(len(clip)), |
| 270 | "video_clip_center_frame": int(center_frame), |
| 271 | "video_total_frames": int(len(frames)), |
| 272 | } |
| 273 | return clip, metadata |
| 274 | |
| 275 | @staticmethod |
| 276 | def _select_video_clip_for_audio_window( |
| 277 | frames: list[Image.Image], |
| 278 | *, |
| 279 | audio_window_start: int, |
| 280 | audio_window_end: int, |
| 281 | audio_total_frames: int, |
| 282 | video_clip_num_frames: int, |
| 283 | ) -> tuple[list[Image.Image], dict[str, Any]]: |
| 284 | audio_total_frames = max(1, int(audio_total_frames)) |
| 285 | window_center = (float(audio_window_start) + float(audio_window_end - 1)) * 0.5 |
| 286 | center_ratio = window_center / float(max(audio_total_frames - 1, 1)) |
| 287 | center_frame = int(round(center_ratio * float(max(len(frames) - 1, 0)))) |
| 288 | return PairedAudioVideoMemoryBank._select_video_clip_around_frame( |
| 289 | frames, |
| 290 | center_frame=center_frame, |
| 291 | video_clip_num_frames=video_clip_num_frames, |
| 292 | ) |
| 293 | |
| 294 | def _trim(self) -> None: |
| 295 | if self.max_size <= 0 or len(self.memory) <= self.max_size: |
| 296 | return |
| 297 | fixed = self.memory[: self.num_fix_frames] |
| 298 | tail = self.memory[self.num_fix_frames :] |
| 299 | keep_tail = max(0, self.max_size - len(fixed)) |
| 300 | self.memory = fixed + tail[-keep_tail:] |
| 301 | |
| 302 | def save_memory_slot( |
| 303 | self, |
| 304 | frames: list[Image.Image], |
| 305 | audio_latent: torch.Tensor, |
| 306 | *, |
| 307 | audio_window_size: int, |
| 308 | video_clip_num_frames: int, |
| 309 | audio_waveform: Optional[torch.Tensor] = None, |
| 310 | audio_sample_rate: int = 16000, |
| 311 | video_fps: float = 24.0, |
| 312 | audio_window_selection_mode: str = "center", |
| 313 | video_frame_selection_mode: str = "center", |
| 314 | audio_memory_mel_bins: int = 128, |
| 315 | audio_memory_mel_hop_length: int = 160, |
| 316 | audio_memory_n_fft: int = 1024, |
| 317 | audio_memory_downsample_factor: int = 4, |
| 318 | audio_memory_is_causal: bool = True, |
| 319 | ) -> dict[str, Any]: |
| 320 | audio_latent = self._prepare_audio_latent(audio_latent) |
| 321 | if audio_latent is None: |
| 322 | raise ValueError("paired audio memory slot requires audio_latent") |
| 323 | |
| 324 | selection_mode = str(audio_window_selection_mode).lower() |
| 325 | if audio_waveform is not None and selection_mode != "center": |
| 326 | try: |
| 327 | window_latent, audio_metadata = self._select_audio_window_from_waveform( |
| 328 | audio_latent, |
| 329 | audio_waveform=audio_waveform, |
| 330 | audio_sample_rate=audio_sample_rate, |
| 331 | window_size=audio_window_size, |
| 332 | selection_mode=selection_mode, |
| 333 | mel_bins=audio_memory_mel_bins, |
| 334 | mel_hop_length=audio_memory_mel_hop_length, |
| 335 | n_fft=audio_memory_n_fft, |
| 336 | downsample_factor=audio_memory_downsample_factor, |
| 337 | is_causal=audio_memory_is_causal, |
| 338 | ) |
| 339 | selected_frame = select_video_frame_indices_from_time_range( |
| 340 | num_frames=len(frames), |
| 341 | fps=float(video_fps), |
| 342 | start_time_sec=float(audio_metadata["audio_window_start_time_sec"]), |
| 343 | end_time_sec=float(audio_metadata["audio_window_end_time_sec"]), |
| 344 | count=1, |
| 345 | mode=str(video_frame_selection_mode).lower(), |
| 346 | )[0] |
| 347 | video_clip, video_metadata = self._select_video_clip_around_frame( |
| 348 | frames, |
| 349 | center_frame=int(selected_frame), |
| 350 | video_clip_num_frames=video_clip_num_frames, |
| 351 | ) |
| 352 | except Exception as exc: |
| 353 | window_latent, audio_metadata = self._select_audio_window(audio_latent, audio_window_size) |
| 354 | audio_metadata["audio_window_selection_mode"] = "center" |
| 355 | audio_metadata["selection_fallback"] = f"{selection_mode}: {exc}" |
| 356 | video_clip, video_metadata = self._select_video_clip_for_audio_window( |
| 357 | frames, |
| 358 | audio_window_start=int(audio_metadata["audio_window_start"]), |
| 359 | audio_window_end=int(audio_metadata["audio_window_end"]), |
| 360 | audio_total_frames=int(audio_metadata["audio_total_frames"]), |
| 361 | video_clip_num_frames=video_clip_num_frames, |
| 362 | ) |
| 363 | else: |
| 364 | window_latent, audio_metadata = self._select_audio_window(audio_latent, audio_window_size) |
| 365 | audio_metadata["audio_window_selection_mode"] = "center" |
| 366 | video_clip, video_metadata = self._select_video_clip_for_audio_window( |
| 367 | frames, |
| 368 | audio_window_start=int(audio_metadata["audio_window_start"]), |
| 369 | audio_window_end=int(audio_metadata["audio_window_end"]), |
| 370 | audio_total_frames=int(audio_metadata["audio_total_frames"]), |
| 371 | video_clip_num_frames=video_clip_num_frames, |
| 372 | ) |
| 373 | |
| 374 | metadata = {"selection_mode": "paired_audio_window", **audio_metadata, **video_metadata} |
| 375 | entry = MemoryEntry(frame=video_clip, audio_latent=window_latent, metadata=metadata) |
| 376 | fixed = self.memory[: self.num_fix_frames] |
| 377 | free = self.memory[self.num_fix_frames :] |
| 378 | free.append(entry) |
| 379 | self.memory = fixed + free |
| 380 | self._trim() |
| 381 | return metadata |
| 382 | |
| 383 | def get_memory_frames(self) -> list[Image.Image | list[Image.Image]]: |
| 384 | return [entry.frame for entry in self.memory] |
| 385 | |
| 386 | def get_memory_metadata(self) -> list[dict[str, Any]]: |
| 387 | return [dict(entry.metadata) for entry in self.memory] |
| 388 | |
| 389 | def get_memory_audio(self) -> Optional[torch.Tensor]: |
| 390 | audio_latents = [entry.audio_latent for entry in self.memory] |
| 391 | if not audio_latents or any(audio_latent is None for audio_latent in audio_latents): |
| 392 | return None |
| 393 | first = audio_latents[0] |
| 394 | assert first is not None |
| 395 | batch_size = first.shape[0] |
| 396 | channels = first.shape[2] |
| 397 | for audio_latent in audio_latents: |
| 398 | assert audio_latent is not None |
| 399 | if audio_latent.shape[0] != batch_size or audio_latent.shape[2] != channels: |
| 400 | raise ValueError( |
| 401 | "All memory audio latents must share batch and channel dimensions, " |
| 402 | f"got first={tuple(first.shape)} current={tuple(audio_latent.shape)}" |
| 403 | ) |
| 404 | return torch.cat(audio_latents, dim=1).contiguous() |
| 405 | |
| 406 | def get_memory_audio_segment_lengths(self) -> tuple[tuple[int, ...], ...]: |
| 407 | audio_latents = [entry.audio_latent for entry in self.memory] |
| 408 | if not audio_latents or any(audio_latent is None for audio_latent in audio_latents): |
| 409 | return () |
| 410 | return (tuple(int(audio_latent.shape[1]) for audio_latent in audio_latents if audio_latent is not None),) |
| 411 | |
| 412 | def __len__(self) -> int: |
| 413 | return len(self.memory) |
| 414 | |
| 415 | |
| 416 | def video_uint8_to_pil_frames(video_uint8: torch.Tensor) -> list[Image.Image]: |
| 417 | if video_uint8.ndim != 4: |
| 418 | raise ValueError(f"Expected [F, H, W, C] uint8 video, got shape={tuple(video_uint8.shape)}") |
| 419 | if video_uint8.shape[-1] != 3: |
| 420 | raise ValueError(f"Expected RGB video with trailing channel dim 3, got shape={tuple(video_uint8.shape)}") |
| 421 | video_uint8 = video_uint8.detach().cpu().contiguous() |
| 422 | return [Image.fromarray(frame.numpy()) for frame in video_uint8] |
| 423 |