返回 JoyAI-Echo
memory_multishot.py
根目录 / ltx-distillation / src / ltx_distillation / inference / memory_multishot.py
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
423 lines PYTHON