stzhao's picture
download
raw
19.2 kB
from __future__ import annotations
import argparse
from pathlib import Path
from typing import Callable
import numpy as np
import torch
from tqdm import tqdm
try:
from .forward import (
BACKBONE_SPECS,
ResolvedRAEConfig,
decode_latents,
encode_frames,
frames_to_tensor,
get_device,
load_rae,
load_video_data_entries,
load_video_frames,
resolve_rae_config,
save_video,
split_name_to_tag,
validate_existing_path,
validate_video_root,
)
except ImportError:
from forward import ( # type: ignore
BACKBONE_SPECS,
ResolvedRAEConfig,
decode_latents,
encode_frames,
frames_to_tensor,
get_device,
load_rae,
load_video_data_entries,
load_video_frames,
resolve_rae_config,
save_video,
split_name_to_tag,
validate_existing_path,
validate_video_root,
)
REPO_ROOT = Path(__file__).resolve().parents[2]
DEFAULT_RAE_ROOT = Path("/mnt/posttrain/zhaoshitian/models/RAE-collections")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Encode videos into RAE latents, interpolate between adjacent latent frames, and decode the interpolated sequence back to videos."
)
parser.add_argument(
"--config",
type=Path,
default=None,
help="Optional stage-1 or stage-2 YAML. If omitted, the script builds a stage-1 RAE config from --rae-backbone.",
)
parser.add_argument(
"--rae-backbone",
type=str,
choices=sorted(BACKBONE_SPECS),
default="dinov2",
help="RAE encoder backbone used when --config is not supplied.",
)
parser.add_argument(
"--rae-root",
type=Path,
default=DEFAULT_RAE_ROOT,
help="Root directory containing RAE decoder and latent normalization weights.",
)
parser.add_argument(
"--encoder-path",
type=Path,
default=None,
help="Optional local encoder path override for the selected backbone.",
)
parser.add_argument(
"--video-path",
type=Path,
default=None,
help="Optional single video to process. When set, split-based batch processing is skipped.",
)
parser.add_argument(
"--video-data-dir-path",
type=Path,
help="Root directory containing video data and data.json. Videos are resolved relative to this directory.",
)
parser.add_argument(
"--dataset-name",
type=str,
default="ucf101",
help="Dataset name used in default output directory paths (e.g. 'bridgedatav2').",
)
parser.add_argument(
"--data-split",
type=str,
default=None,
help="If set, only load entries whose 'split' field matches (e.g. 'train', 'raw').",
)
parser.add_argument(
"--data-category",
type=str,
default=None,
help="If set, only load entries whose 'category' field matches (e.g. 'toykitchen1').",
)
parser.add_argument(
"--output-dir",
type=Path,
default=None,
help="Directory to store interpolated reconstruction videos.",
)
parser.add_argument(
"--save-latents",
action="store_true",
help="Also save the interpolated latent sequences as `.npz` files.",
)
parser.add_argument(
"--latent-output-dir",
type=Path,
default=None,
help="Directory to store interpolated latent `.npz` files when --save-latents is enabled.",
)
parser.add_argument(
"--image-size",
type=int,
default=None,
help="Center-crop size before encoding. Defaults to the selected encoder's input size.",
)
parser.add_argument("--encode-batch-size", type=int, default=16, help="Batch size used for latent extraction.")
parser.add_argument("--decode-batch-size", type=int, default=16, help="Batch size used for latent decoding.")
parser.add_argument("--device", type=str, default=None, help="Torch device. Defaults to cuda when available.")
parser.add_argument("--max-videos", type=int, default=None, help="Optional cap for debugging a subset of videos.")
parser.add_argument("--overwrite", action="store_true", help="Recompute outputs even when they already exist.")
parser.add_argument(
"--num-interp",
type=int,
required=True,
help="Number of interpolated latent frames inserted between every pair of adjacent source frames.",
)
parser.add_argument(
"--interp-mode",
type=str,
choices=("auto", "linear", "slerp"),
default="auto",
help="Interpolation geometry used between adjacent latent frames.",
)
return parser.parse_args()
def validate_args(args: argparse.Namespace) -> None:
if args.num_interp < 1:
raise ValueError("--num-interp must be at least 1.")
if args.latent_output_dir is not None and not args.save_latents:
raise ValueError("--latent-output-dir requires --save-latents.")
if args.video_path is not None and args.max_videos is not None:
raise ValueError("--max-videos cannot be used together with --video-path.")
def resolve_interpolation_mode(backbone: str, interp_mode: str) -> str:
if interp_mode != "auto":
return interp_mode
if backbone == "siglip2":
return "slerp"
return "linear"
def default_video_output_root(dataset_name: str, split_name: str, output_suffix: str, interp_mode: str, num_interp: int) -> Path:
name = f"{dataset_name}_{split_name_to_tag(split_name)}_{output_suffix}_{interp_mode}_x{num_interp}"
return REPO_ROOT / "video_interpolations" / name
def default_latent_output_root(dataset_name: str, split_name: str, output_suffix: str, interp_mode: str, num_interp: int) -> Path:
name = f"{dataset_name}_{split_name_to_tag(split_name)}_{output_suffix}_{interp_mode}_x{num_interp}"
return REPO_ROOT / "video_features_interpolated" / name
def default_single_video_output_root(output_suffix: str, interp_mode: str, num_interp: int) -> Path:
name = f"single_video_{output_suffix}_{interp_mode}_x{num_interp}"
return REPO_ROOT / "video_interpolations" / name
def default_single_latent_output_root(output_suffix: str, interp_mode: str, num_interp: int) -> Path:
name = f"single_video_{output_suffix}_{interp_mode}_x{num_interp}"
return REPO_ROOT / "video_features_interpolated" / name
def resolve_runtime_args(args: argparse.Namespace, resolved: ResolvedRAEConfig) -> str:
if args.video_path is not None:
args.video_path = args.video_path.expanduser()
args.video_root = args.video_data_dir_path.expanduser()
args.rae_root = args.rae_root.expanduser()
if args.image_size is None:
args.image_size = resolved.encoder_input_size
applied_interp_mode = resolve_interpolation_mode(resolved.backbone, args.interp_mode)
if args.output_dir is None:
if args.video_path is not None:
args.output_dir = default_single_video_output_root(resolved.output_suffix, applied_interp_mode, args.num_interp)
else:
args.output_dir = default_video_output_root(args.dataset_name, args.data_split or "data", resolved.output_suffix, applied_interp_mode, args.num_interp)
else:
args.output_dir = args.output_dir.expanduser()
if args.save_latents:
if args.latent_output_dir is None:
if args.video_path is not None:
args.latent_output_dir = default_single_latent_output_root(
resolved.output_suffix,
applied_interp_mode,
args.num_interp,
)
else:
args.latent_output_dir = default_latent_output_root(
args.dataset_name,
args.data_split or "data",
resolved.output_suffix,
applied_interp_mode,
args.num_interp,
)
else:
args.latent_output_dir = args.latent_output_dir.expanduser()
return applied_interp_mode
def print_run_configuration(
args: argparse.Namespace,
resolved: ResolvedRAEConfig,
device: torch.device,
applied_interp_mode: str,
) -> None:
print(f"RAE source: {resolved.source}")
print(f"RAE backbone: {resolved.backbone}")
print(f"Encoder path/id: {resolved.encoder_reference}")
if resolved.decoder_path is not None:
print(f"Decoder checkpoint: {resolved.decoder_path}")
if resolved.stats_path is not None:
print(f"Normalization stats: {resolved.stats_path}")
print(f"Encoder input size: {resolved.encoder_input_size}")
print(f"Frame crop size: {args.image_size}")
print(f"Interpolation mode: {applied_interp_mode}")
print(f"Interpolated frames per gap: {args.num_interp}")
if args.video_path is not None:
print(f"Input video: {args.video_path}")
else:
print(f"Input split: {args.data_split or 'all'}")
print(f"Output videos: {args.output_dir}")
if args.save_latents:
print(f"Output interpolated latents: {args.latent_output_dir}")
print(f"Device: {device}")
def build_output_paths(
video_root: Path,
latent_root: Path | None,
relative_video_path: Path,
) -> tuple[Path, Path | None]:
video_stem = relative_video_path.stem
parent_dir = relative_video_path.parent
if str(parent_dir) == ".":
video_path = video_root / f"{video_stem}.mp4"
else:
video_path = video_root / parent_dir / f"{video_stem}.mp4"
if latent_root is None:
return video_path, None
if str(parent_dir) == ".":
latent_path = latent_root / f"{video_stem}_patch_tokens.npz"
else:
latent_path = latent_root / parent_dir / f"{video_stem}_patch_tokens.npz"
return video_path, latent_path
def collect_video_entries(args: argparse.Namespace) -> list[tuple[Path, Path]]:
if args.video_path is not None:
validate_existing_path(args.video_path, "Input video")
try:
relative_video_path = args.video_path.relative_to(args.video_root)
except ValueError:
relative_video_path = Path(args.video_path.name)
return [(relative_video_path, args.video_path)]
validate_video_root(args.video_root)
entries = load_video_data_entries(args.video_data_dir_path, split=args.data_split, category=args.data_category)
if args.max_videos is not None:
entries = entries[: args.max_videos]
missing = [entry for entry in entries if not (args.video_root / entry).exists()]
if missing:
preview = ", ".join(missing[:5])
raise FileNotFoundError(f"Missing {len(missing)} videos under {args.video_root}. First missing entries: {preview}")
return [(Path(entry), args.video_root / entry) for entry in entries]
def linear_interpolate_pair(left: torch.Tensor, right: torch.Tensor, alpha: float) -> torch.Tensor:
return torch.lerp(left, right, alpha)
def slerp_interpolate_pair(
left: torch.Tensor,
right: torch.Tensor,
alpha: float,
*,
eps: float = 1e-6,
dot_threshold: float = 0.9995,
) -> torch.Tensor:
if left.shape != right.shape:
raise ValueError(f"Slerp expects matching latent shapes, got {tuple(left.shape)} and {tuple(right.shape)}")
if left.dim() != 3:
raise ValueError(f"Slerp expects a single latent frame shaped [C, H, W], got {tuple(left.shape)}")
channels = left.shape[0]
left_vecs = left.movedim(0, -1).reshape(-1, channels)
right_vecs = right.movedim(0, -1).reshape(-1, channels)
result = torch.lerp(left_vecs, right_vecs, alpha)
left_norm = left_vecs.norm(dim=-1, keepdim=True)
right_norm = right_vecs.norm(dim=-1, keepdim=True)
valid_mask = (left_norm.squeeze(-1) > eps) & (right_norm.squeeze(-1) > eps)
if valid_mask.any():
left_unit = left_vecs[valid_mask] / left_norm[valid_mask]
right_unit = right_vecs[valid_mask] / right_norm[valid_mask]
dot = (left_unit * right_unit).sum(dim=-1, keepdim=True).clamp(-1.0, 1.0)
theta = torch.acos(dot.clamp(-1.0 + eps, 1.0 - eps))
sin_theta = torch.sin(theta)
target_norm = torch.lerp(left_norm[valid_mask], right_norm[valid_mask], alpha)
linear_vec = torch.lerp(left_vecs[valid_mask], right_vecs[valid_mask], alpha)
linear_norm = linear_vec.norm(dim=-1, keepdim=True)
linear_dir = linear_vec / linear_norm.clamp_min(eps)
linear_dir = torch.where(linear_norm > eps, linear_dir, left_unit)
coeff_left = torch.sin((1.0 - alpha) * theta) / sin_theta.clamp_min(eps)
coeff_right = torch.sin(alpha * theta) / sin_theta.clamp_min(eps)
slerp_dir = coeff_left * left_unit + coeff_right * right_unit
slerp_norm = slerp_dir.norm(dim=-1, keepdim=True)
slerp_dir = slerp_dir / slerp_norm.clamp_min(eps)
slerp_dir = torch.where(slerp_norm > eps, slerp_dir, left_unit)
fallback_mask = (sin_theta.abs() < eps) | (dot.abs() > dot_threshold)
direction = torch.where(fallback_mask, linear_dir, slerp_dir)
result[valid_mask] = direction * target_norm
return result.reshape(*left.movedim(0, -1).shape).movedim(-1, 0)
def get_interpolator(interp_mode: str) -> Callable[[torch.Tensor, torch.Tensor, float], torch.Tensor]:
if interp_mode == "linear":
return linear_interpolate_pair
if interp_mode == "slerp":
return slerp_interpolate_pair
raise ValueError(f"Unsupported interpolation mode: {interp_mode}")
def interpolate_timesteps(source_timesteps: np.ndarray, num_interp: int) -> np.ndarray:
if source_timesteps.size <= 1:
return source_timesteps.astype(np.float32, copy=True)
output = [np.float32(source_timesteps[0])]
for idx in range(source_timesteps.size - 1):
left = source_timesteps[idx]
right = source_timesteps[idx + 1]
for step in range(1, num_interp + 1):
alpha = step / float(num_interp + 1)
output.append(np.float32(left + alpha * (right - left)))
output.append(np.float32(right))
return np.asarray(output, dtype=np.float32)
def interpolate_latent_sequence(
latents: torch.Tensor,
num_interp: int,
interp_mode: str,
) -> tuple[torch.Tensor, np.ndarray, np.ndarray]:
if latents.dim() != 4:
raise ValueError(f"Expected latents shaped [T, C, H, W], got {tuple(latents.shape)}")
if latents.shape[0] == 0:
raise ValueError("Cannot interpolate an empty latent sequence.")
source_timesteps = np.linspace(0.0, 1.0, latents.shape[0], dtype=np.float32)
if latents.shape[0] == 1:
return latents.clone(), source_timesteps.copy(), source_timesteps
interpolator = get_interpolator(interp_mode)
output_latents = [latents[0]]
for idx in range(latents.shape[0] - 1):
left = latents[idx]
right = latents[idx + 1]
for step in range(1, num_interp + 1):
alpha = step / float(num_interp + 1)
output_latents.append(interpolator(left, right, alpha))
output_latents.append(right)
interpolated_timesteps = interpolate_timesteps(source_timesteps, num_interp)
return torch.stack(output_latents, dim=0), interpolated_timesteps, source_timesteps
def save_interpolated_latents(
save_path: Path,
latents: torch.Tensor,
timesteps: np.ndarray,
source_timesteps: np.ndarray,
num_interp: int,
interp_mode: str,
) -> None:
save_path.parent.mkdir(parents=True, exist_ok=True)
np.savez(
save_path,
features=latents.numpy().astype(np.float32),
timesteps=timesteps.astype(np.float32),
source_timesteps=source_timesteps.astype(np.float32),
num_interp=np.int32(num_interp),
interp_mode=np.array(interp_mode),
)
def target_video_fps(source_fps: float, source_frame_count: int, num_interp: int) -> float:
if source_frame_count <= 1:
return float(source_fps)
return float(source_fps) * float(num_interp + 1)
def process_video(
rae,
video_path: Path,
video_output_path: Path,
latent_output_path: Path | None,
args: argparse.Namespace,
device: torch.device,
applied_interp_mode: str,
) -> None:
need_video = args.overwrite or not video_output_path.exists()
need_latents = args.save_latents and latent_output_path is not None and (args.overwrite or not latent_output_path.exists())
if not need_video and not need_latents:
return
frames_np, source_fps = load_video_frames(video_path, args.image_size)
frame_tensor = frames_to_tensor(frames_np)
source_latents = encode_frames(rae, frame_tensor, args.encode_batch_size, device)
interpolated_latents, interpolated_timesteps, source_timesteps = interpolate_latent_sequence(
source_latents,
args.num_interp,
applied_interp_mode,
)
if need_latents and latent_output_path is not None:
save_interpolated_latents(
latent_output_path,
interpolated_latents,
interpolated_timesteps,
source_timesteps,
args.num_interp,
applied_interp_mode,
)
if need_video:
decoded_frames = decode_latents(rae, interpolated_latents, args.decode_batch_size, device)
save_video(
decoded_frames,
video_output_path,
fps=target_video_fps(source_fps, source_latents.shape[0], args.num_interp),
)
def main() -> None:
args = parse_args()
validate_args(args)
device = get_device(args.device)
resolved_rae = resolve_rae_config(args)
applied_interp_mode = resolve_runtime_args(args, resolved_rae)
video_entries = collect_video_entries(args)
args.output_dir.mkdir(parents=True, exist_ok=True)
if args.save_latents and args.latent_output_dir is not None:
args.latent_output_dir.mkdir(parents=True, exist_ok=True)
print_run_configuration(args, resolved_rae, device, applied_interp_mode)
rae = load_rae(resolved_rae.config, device)
progress_label = "Interpolating video latents" if args.video_path is not None else "Interpolating UCF101 latents"
progress = tqdm(video_entries, desc=progress_label)
for relative_video_path, video_path in progress:
video_output_path, latent_output_path = build_output_paths(
args.output_dir,
args.latent_output_dir if args.save_latents else None,
relative_video_path,
)
process_video(
rae,
video_path,
video_output_path,
latent_output_path,
args,
device,
applied_interp_mode,
)
progress.set_postfix_str(relative_video_path.name)
print(f"Saved interpolated videos to {args.output_dir}")
if args.save_latents and args.latent_output_dir is not None:
print(f"Saved interpolated latents to {args.latent_output_dir}")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
19.2 kB
·
Xet hash:
bcf2c7cfc33ee179e8239b34ad9cdf8c6e5864e8f72243f4b68cedaca98bcee2

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.