Buckets:
| 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.