| 1 | """Audio-video memory helpers aligned with the Echo 1.5 online runtime.""" |
| 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 | |
| 13 | from ltx_distillation.audio_voice_filter import VoiceFilterConfig, filter_voice_only |
| 14 | |
| 15 | |
| 16 | def prompt_payload_to_text(payload: Any, prompt_max_chars: Optional[int] = None) -> str: |
| 17 | if not isinstance(payload, str): |
| 18 | raise TypeError( |
| 19 | f"unsupported prompt payload type: {type(payload).__name__}; " |
| 20 | "each shot must be one prompt string" |
| 21 | ) |
| 22 | text = payload.strip() |
| 23 | return text[:prompt_max_chars] if prompt_max_chars else text |
| 24 | |
| 25 | |
| 26 | def json_to_prompts( |
| 27 | data: dict[str, Any], prompt_max_chars: Optional[int] = None |
| 28 | ) -> list[str]: |
| 29 | values = data.get("prompts", data.get("shots", [])) |
| 30 | if not isinstance(values, list): |
| 31 | return [] |
| 32 | prompts = [prompt_payload_to_text(item, prompt_max_chars) for item in values] |
| 33 | return [prompt for prompt in prompts if prompt] |
| 34 | |
| 35 | |
| 36 | def load_multishot_prompts( |
| 37 | prompts_file: str | Path, |
| 38 | prompt_max_chars: Optional[int] = None, |
| 39 | ) -> list[str]: |
| 40 | path = Path(prompts_file) |
| 41 | with path.open("r", encoding="utf-8") as handle: |
| 42 | payload = json.load(handle) |
| 43 | prompts = json_to_prompts(payload, prompt_max_chars=prompt_max_chars) |
| 44 | if not prompts: |
| 45 | raise ValueError(f"no prompts found in {path}") |
| 46 | return prompts |
| 47 | |
| 48 | |
| 49 | def normalize_audio_waveform_for_media( |
| 50 | audio_waveform: Optional[torch.Tensor], |
| 51 | ) -> Optional[torch.Tensor]: |
| 52 | if audio_waveform is None: |
| 53 | return None |
| 54 | waveform = getattr(audio_waveform, "waveform", audio_waveform) |
| 55 | waveform = torch.as_tensor(waveform).detach().cpu().float() |
| 56 | if waveform.ndim == 3: |
| 57 | if waveform.shape[0] != 1: |
| 58 | raise ValueError( |
| 59 | f"expected batch size 1, got shape={tuple(waveform.shape)}" |
| 60 | ) |
| 61 | waveform = waveform[0] |
| 62 | if waveform.ndim == 1: |
| 63 | waveform = waveform.unsqueeze(0) |
| 64 | elif ( |
| 65 | waveform.ndim == 2 |
| 66 | and waveform.shape[0] not in {1, 2} |
| 67 | and waveform.shape[1] in {1, 2} |
| 68 | ): |
| 69 | waveform = waveform.transpose(0, 1) |
| 70 | elif waveform.ndim != 2: |
| 71 | raise ValueError( |
| 72 | f"expected decoded audio with 1-3 dims, got shape={tuple(waveform.shape)}" |
| 73 | ) |
| 74 | if waveform.shape[0] == 1: |
| 75 | waveform = waveform.repeat(2, 1) |
| 76 | elif waveform.shape[0] > 2: |
| 77 | waveform = waveform[:2] |
| 78 | return waveform.contiguous() |
| 79 | |
| 80 | |
| 81 | def audio_waveform_stats(audio_waveform: Optional[torch.Tensor]) -> dict[str, Any]: |
| 82 | waveform = normalize_audio_waveform_for_media(audio_waveform) |
| 83 | if waveform is None: |
| 84 | return { |
| 85 | "present": False, |
| 86 | "shape": None, |
| 87 | "num_samples": 0, |
| 88 | "rms": 0.0, |
| 89 | "peak": 0.0, |
| 90 | } |
| 91 | waveform_f = waveform.float() |
| 92 | return { |
| 93 | "present": True, |
| 94 | "shape": list(waveform.shape), |
| 95 | "num_samples": int(waveform.shape[-1]), |
| 96 | "rms": float(waveform_f.square().mean().sqrt().item()), |
| 97 | "peak": float(waveform_f.abs().max().item()), |
| 98 | } |
| 99 | |
| 100 | |
| 101 | @dataclass |
| 102 | class MemoryEntry: |
| 103 | video_latent: torch.Tensor |
| 104 | audio_waveform: Optional[torch.Tensor] |
| 105 | audio_sample_rate: int |
| 106 | metadata: dict[str, Any] = field(default_factory=dict) |
| 107 | |
| 108 | |
| 109 | class AudioVideoMemoryBank: |
| 110 | """Stores one generated video-latent frame and full voice waveform per shot.""" |
| 111 | |
| 112 | def __init__(self, max_size: int, num_fix_frames: int = 0) -> None: |
| 113 | self.max_size = max(0, int(max_size)) |
| 114 | self.num_fix_frames = max(0, int(num_fix_frames)) |
| 115 | self.memory: list[MemoryEntry] = [] |
| 116 | |
| 117 | def _trim(self) -> None: |
| 118 | if self.max_size <= 0: |
| 119 | self.memory = [] |
| 120 | return |
| 121 | if len(self.memory) <= self.max_size: |
| 122 | return |
| 123 | fixed_count = min(self.num_fix_frames, self.max_size) |
| 124 | fixed = self.memory[:fixed_count] |
| 125 | keep_tail = self.max_size - fixed_count |
| 126 | tail = self.memory[-keep_tail:] if keep_tail else [] |
| 127 | self.memory = fixed + tail |
| 128 | |
| 129 | def save_generated_shot( |
| 130 | self, |
| 131 | video_latent: torch.Tensor, |
| 132 | audio_waveform: Optional[torch.Tensor], |
| 133 | audio_sample_rate: int, |
| 134 | *, |
| 135 | enable_audio_memory: bool, |
| 136 | voice_filter_config: VoiceFilterConfig, |
| 137 | ) -> dict[str, Any]: |
| 138 | if video_latent.ndim != 5 or video_latent.shape[0] != 1: |
| 139 | raise ValueError( |
| 140 | f"expected video latent [1, F, C, H, W], got {tuple(video_latent.shape)}" |
| 141 | ) |
| 142 | num_frames = int(video_latent.shape[1]) |
| 143 | if num_frames <= 0: |
| 144 | raise ValueError("cannot save memory from an empty video latent") |
| 145 | frame_index = random.randrange(num_frames) |
| 146 | selected_video = ( |
| 147 | video_latent[:, frame_index : frame_index + 1].detach().cpu().contiguous() |
| 148 | ) |
| 149 | |
| 150 | filtered_audio = None |
| 151 | if enable_audio_memory and audio_waveform is not None: |
| 152 | normalized = normalize_audio_waveform_for_media(audio_waveform) |
| 153 | filtered_audio = filter_voice_only( |
| 154 | normalized, |
| 155 | int(audio_sample_rate), |
| 156 | voice_filter_config, |
| 157 | ) |
| 158 | if filtered_audio is not None: |
| 159 | filtered_audio = filtered_audio.detach().cpu().contiguous() |
| 160 | |
| 161 | metadata = { |
| 162 | "selection": "random_latent_frame_full_audio", |
| 163 | "video_frame_index": frame_index, |
| 164 | "video_total_latent_frames": num_frames, |
| 165 | "audio_present": filtered_audio is not None, |
| 166 | "audio_samples": int(filtered_audio.shape[-1]) |
| 167 | if filtered_audio is not None |
| 168 | else 0, |
| 169 | "audio_sample_rate": int(audio_sample_rate), |
| 170 | "voice_filter_backend": voice_filter_config.backend |
| 171 | if enable_audio_memory |
| 172 | else "disabled", |
| 173 | } |
| 174 | self.memory.append( |
| 175 | MemoryEntry( |
| 176 | video_latent=selected_video, |
| 177 | audio_waveform=filtered_audio, |
| 178 | audio_sample_rate=int(audio_sample_rate), |
| 179 | metadata=metadata, |
| 180 | ) |
| 181 | ) |
| 182 | self._trim() |
| 183 | return metadata |
| 184 | |
| 185 | def get_memory_video(self) -> torch.Tensor: |
| 186 | if not self.memory: |
| 187 | raise RuntimeError("memory bank is empty") |
| 188 | return torch.cat( |
| 189 | [entry.video_latent for entry in self.memory], dim=1 |
| 190 | ).contiguous() |
| 191 | |
| 192 | @torch.no_grad() |
| 193 | def encode_memory_audio(self, audio_vae) -> list[Optional[torch.Tensor]]: |
| 194 | encoded: list[Optional[torch.Tensor]] = [] |
| 195 | for entry in self.memory: |
| 196 | waveform = entry.audio_waveform |
| 197 | if waveform is None or waveform.numel() <= 1 or waveform.shape[-1] <= 1: |
| 198 | encoded.append(None) |
| 199 | else: |
| 200 | encoded.append( |
| 201 | audio_vae.encode(waveform, entry.audio_sample_rate) |
| 202 | .detach() |
| 203 | .cpu() |
| 204 | .contiguous() |
| 205 | ) |
| 206 | return encoded |
| 207 | |
| 208 | def get_memory_metadata(self) -> list[dict[str, Any]]: |
| 209 | return [dict(entry.metadata) for entry in self.memory] |
| 210 | |
| 211 | def __len__(self) -> int: |
| 212 | return len(self.memory) |
| 213 | |
| 214 | |
| 215 | @torch.no_grad() |
| 216 | def build_memory_audio_pipeline_kwargs( |
| 217 | memory_bank: AudioVideoMemoryBank, |
| 218 | audio_vae, |
| 219 | *, |
| 220 | enable_audio_memory: bool, |
| 221 | memory_position_mode: str, |
| 222 | memory_position_offset: float, |
| 223 | memory_position_slot_stride: float, |
| 224 | ) -> dict[str, Any]: |
| 225 | """Encode full per-slot waveforms and assemble position-aligned audio memory.""" |
| 226 | |
| 227 | if not enable_audio_memory: |
| 228 | return {} |
| 229 | audio_slices = memory_bank.encode_memory_audio(audio_vae) |
| 230 | template = next((item for item in audio_slices if item is not None), None) |
| 231 | if template is None: |
| 232 | return {} |
| 233 | aligned = [ |
| 234 | item if item is not None else torch.zeros_like(template) |
| 235 | for item in audio_slices |
| 236 | ] |
| 237 | memory_audio = torch.cat(aligned, dim=1).contiguous() |
| 238 | segment_lengths = tuple(int(item.shape[1]) for item in aligned) |
| 239 | return { |
| 240 | "memory_audio": memory_audio, |
| 241 | "memory_audio_timestep": torch.zeros( |
| 242 | memory_audio.shape[:2], dtype=torch.float32 |
| 243 | ), |
| 244 | "memory_audio_segment_lengths": (segment_lengths,), |
| 245 | "memory_position_mode": str(memory_position_mode), |
| 246 | "memory_position_offset": float(memory_position_offset), |
| 247 | "memory_position_slot_stride": float(memory_position_slot_stride), |
| 248 | } |
| 249 |