EchoLVFM / space /video_io.py
EngEmmanuel's picture
Initial commit
0f5513d
Raw History Blame Contribute Delete
3.48 kB
"""Video / frame I/O helpers for the Space.
Two responsibilities:
- Encode a `(T, 3, H, W)` float tensor in [0, 1] to an mp4 (OpenCV — no
ffmpeg dependency, since the Space container won't have it).
- Load a directory of `frame_*.png` files into a `(T, 3, H, W)` tensor in
[0, 1] (used to load the bundled real frames at request time).
"""
from __future__ import annotations
from pathlib import Path
import cv2
import numpy as np
import torch
from PIL import Image
def frames_to_mp4(frames_T3HW: torch.Tensor, out_path: Path, fps: float) -> Path:
"""Write a (T, 3, H, W) float tensor in [0, 1] to mp4 at `fps`.
Uses OpenCV's `mp4v` codec — works inside the Space container (no system
ffmpeg required). Returns `out_path`.
"""
if frames_T3HW.ndim != 4 or frames_T3HW.shape[1] != 3:
raise ValueError(f"Expected (T, 3, H, W), got {tuple(frames_T3HW.shape)}")
out_path = Path(out_path)
out_path.parent.mkdir(parents=True, exist_ok=True)
T, _, H, W = frames_T3HW.shape
fourcc = cv2.VideoWriter_fourcc(*"mp4v")
writer = cv2.VideoWriter(str(out_path), fourcc, float(fps), (W, H))
if not writer.isOpened():
raise RuntimeError(f"OpenCV VideoWriter failed to open {out_path}")
try:
arr = frames_T3HW.clamp(0, 1).mul(255).round().to(torch.uint8)
arr = arr.permute(0, 2, 3, 1).cpu().numpy() # (T, H, W, 3) RGB
for frame_rgb in arr:
# OpenCV expects BGR.
writer.write(cv2.cvtColor(frame_rgb, cv2.COLOR_RGB2BGR))
finally:
writer.release()
return out_path
def stitch_side_by_side(
real_T3HW: torch.Tensor, gen_T3HW: torch.Tensor, gap_px: int = 4,
) -> torch.Tensor:
"""Concatenate two `(T, 3, H, W)` clips horizontally with a thin black
gap between them. Returned as `(T, 3, H, W_real + gap + W_gen)`.
Stitching into one mp4 means a single `<video>` element drives both
halves — no second-video load delay, no autoplay-start drift, no JS
sync needed. Both clips must share `T` and `H`."""
if real_T3HW.shape[0] != gen_T3HW.shape[0] or real_T3HW.shape[2] != gen_T3HW.shape[2]:
raise ValueError(
f"Need matching T and H, got real={tuple(real_T3HW.shape)} "
f"gen={tuple(gen_T3HW.shape)}")
T, _, H, _ = real_T3HW.shape
gap = torch.zeros(T, 3, H, gap_px, dtype=real_T3HW.dtype)
return torch.cat([real_T3HW, gap, gen_T3HW], dim=-1)
def frames_dir_to_array(frames_dir: Path, t_real: int | None = None) -> torch.Tensor:
"""Load `frame_*.png` files from a directory into (T, 3, H, W) in [0, 1].
Files are sorted by their numeric suffix (`frame_0.png`, `frame_1.png`,
...). If `t_real` is set, the first `t_real` frames are returned (and an
error is raised if there are fewer).
"""
frames_dir = Path(frames_dir)
paths = sorted(frames_dir.glob("frame_*.png"),
key=lambda p: int(p.stem.split("_")[-1]))
if not paths:
raise FileNotFoundError(f"No frame_*.png in {frames_dir}")
if t_real is not None:
if len(paths) < t_real:
raise ValueError(
f"{frames_dir} has {len(paths)} frames, need {t_real}")
paths = paths[:t_real]
tensors = []
for p in paths:
img = Image.open(p).convert("RGB")
arr = np.asarray(img, dtype=np.float32) / 255.0 # (H, W, 3)
tensors.append(torch.from_numpy(arr).permute(2, 0, 1))
return torch.stack(tensors)