Files
AI-Profile-Router/platform/ltx2-studio/rolling_extend.py
T

144 lines
4.8 KiB
Python

"""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}")