返回 JoyAI-Echo
memory_multishot.py
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
249 lines PYTHON