Add rolling LTX video extension
This commit is contained in:
@@ -0,0 +1,143 @@
|
||||
"""Bounded-context local video extension and loss-minimizing assembly helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
import av
|
||||
from imageio_ffmpeg import get_ffmpeg_exe
|
||||
|
||||
from api_types import ExtendMode
|
||||
|
||||
# LTX video lengths must be 8n+1 frames. Keep this helper independent from the
|
||||
# handlers package so it can also be exercised by maintenance and migration
|
||||
# scripts without importing every request handler.
|
||||
TIME_FACTOR = 8
|
||||
MIN_FRAMES = 9
|
||||
|
||||
|
||||
def correct_frame_count(frames: int) -> int:
|
||||
return max(MIN_FRAMES, (frames // TIME_FACTOR) * TIME_FACTOR + 1)
|
||||
|
||||
_MAX_WINDOW_SECONDS = 18.0
|
||||
_MAX_CONTEXT_SECONDS = 8.0
|
||||
_MIN_CONTEXT_SECONDS = 2.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RollingExtendPlan:
|
||||
rolling: bool
|
||||
context_frames: int
|
||||
total_generation_frames: int
|
||||
|
||||
|
||||
def plan_rolling_extend(source_frames: int, extend_frames: int, fps: float) -> RollingExtendPlan:
|
||||
"""Keep the GPU job bounded while retaining as much source context as possible."""
|
||||
max_total = ((round(_MAX_WINDOW_SECONDS * fps)) // TIME_FACTOR) * TIME_FACTOR + 1
|
||||
if source_frames + extend_frames <= max_total:
|
||||
return RollingExtendPlan(False, source_frames, source_frames + extend_frames)
|
||||
|
||||
max_context = correct_frame_count(round(_MAX_CONTEXT_SECONDS * fps) + 1)
|
||||
available_context = max_total - extend_frames
|
||||
if available_context < round(_MIN_CONTEXT_SECONDS * fps) + 1:
|
||||
raise ValueError(
|
||||
"The requested extension is too long for the rolling VRAM window. "
|
||||
"Choose a shorter duration."
|
||||
)
|
||||
context_frames = min(source_frames, max_context, correct_frame_count(available_context))
|
||||
return RollingExtendPlan(True, context_frames, context_frames + extend_frames)
|
||||
|
||||
|
||||
def extract_context(
|
||||
source: Path,
|
||||
destination: Path,
|
||||
*,
|
||||
mode: ExtendMode,
|
||||
fps: float,
|
||||
source_frames: int,
|
||||
context_frames: int,
|
||||
) -> None:
|
||||
start_frame = 0 if mode == "start" else source_frames - context_frames
|
||||
start_seconds = max(0.0, start_frame / fps)
|
||||
duration_seconds = context_frames / fps
|
||||
_run_ffmpeg(
|
||||
"-i", str(source),
|
||||
"-ss", f"{start_seconds:.9f}",
|
||||
"-t", f"{duration_seconds:.9f}",
|
||||
"-map", "0:v:0", "-map", "0:a:0?",
|
||||
"-c:v", "libx264", "-preset", "fast", "-crf", "16",
|
||||
"-pix_fmt", "yuv420p", "-c:a", "aac", "-b:a", "192k",
|
||||
"-movflags", "+faststart", "-y", str(destination),
|
||||
)
|
||||
|
||||
|
||||
def assemble_rolling_result(
|
||||
source: Path,
|
||||
generated_window: Path,
|
||||
destination: Path,
|
||||
*,
|
||||
mode: ExtendMode,
|
||||
fps: float,
|
||||
source_frames: int,
|
||||
context_frames: int,
|
||||
extend_frames: int,
|
||||
) -> None:
|
||||
"""Join only the newly generated window edge to the untouched logical source."""
|
||||
if mode == "end":
|
||||
window_video = f"trim=start_frame={context_frames}"
|
||||
window_audio = f"atrim=start={context_frames / fps:.9f}"
|
||||
video_inputs = "[v0][v1]"
|
||||
audio_inputs = "[a0][a1]"
|
||||
else:
|
||||
window_video = f"trim=end_frame={extend_frames}"
|
||||
window_audio = f"atrim=end={extend_frames / fps:.9f}"
|
||||
video_inputs = "[v1][v0]"
|
||||
audio_inputs = "[a1][a0]"
|
||||
|
||||
have_audio = _has_audio(source) and _has_audio(generated_window)
|
||||
filters = [
|
||||
f"[0:v]trim=end_frame={source_frames},fps={fps:.9f},format=yuv420p,setpts=PTS-STARTPTS[v0]",
|
||||
f"[1:v]{window_video},fps={fps:.9f},format=yuv420p,setpts=PTS-STARTPTS[v1]",
|
||||
f"{video_inputs}concat=n=2:v=1:a=0[v]",
|
||||
]
|
||||
args = ["-i", str(source), "-i", str(generated_window)]
|
||||
if have_audio:
|
||||
filters.extend(
|
||||
[
|
||||
f"[0:a]atrim=end={source_frames / fps:.9f},aresample=48000,asetpts=PTS-STARTPTS[a0]",
|
||||
f"[1:a]{window_audio},aresample=48000,asetpts=PTS-STARTPTS[a1]",
|
||||
f"{audio_inputs}concat=n=2:v=0:a=1[a]",
|
||||
]
|
||||
)
|
||||
args.extend(["-filter_complex", ";".join(filters), "-map", "[v]"])
|
||||
if have_audio:
|
||||
args.extend(["-map", "[a]", "-c:a", "aac", "-b:a", "192k"])
|
||||
else:
|
||||
args.append("-an")
|
||||
args.extend(
|
||||
[
|
||||
"-c:v", "libx264", "-preset", "fast", "-crf", "16",
|
||||
"-pix_fmt", "yuv420p", "-r", f"{fps:.9f}",
|
||||
"-movflags", "+faststart", "-y", str(destination),
|
||||
]
|
||||
)
|
||||
_run_ffmpeg(*args)
|
||||
|
||||
|
||||
def _has_audio(path: Path) -> bool:
|
||||
with av.open(str(path)) as container:
|
||||
return any(stream.type == "audio" for stream in container.streams)
|
||||
|
||||
|
||||
def _run_ffmpeg(*args: str) -> None:
|
||||
process = subprocess.run(
|
||||
[get_ffmpeg_exe(), "-hide_banner", "-loglevel", "error", *args],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False,
|
||||
)
|
||||
if process.returncode != 0:
|
||||
detail = process.stderr.strip()[-2000:] or f"ffmpeg exited with {process.returncode}"
|
||||
raise RuntimeError(f"Rolling extend media assembly failed: {detail}")
|
||||
Reference in New Issue
Block a user