Buckets:
| """ | |
| Compute reference videos for IC-LoRA training. | |
| This script provides a command-line interface for generating reference videos to be used for IC-LoRA training. | |
| Single-file mode (input is a video file): | |
| # Compute Canny edge reference for one video — writes {name}_canny.mp4 next to the input | |
| compute_reference.py video.mp4 | |
| # Save to a specific output directory | |
| compute_reference.py video.mp4 --output-path /path/to/output/ | |
| # Compute depth reference for one video | |
| compute_reference.py video.mp4 --type depth --depth-repo-path /path/to/Video-Depth-Anything | |
| # Compute pose reference for one video | |
| compute_reference.py video.mp4 --type pose --dwpose-repo-path /path/to/DWPose | |
| Directory mode (input is a folder): | |
| Processes all .mp4 files in the folder. | |
| Without --output-path: writes {name}_{type}.mp4 next to each input. | |
| With --output-path: writes all reference videos to the specified directory. | |
| Videos are sharded across worker processes, one per visible GPU by default | |
| (each worker loads its own model instance). Override with --num-workers. | |
| # Compute Canny edge reference for all videos in a folder | |
| compute_reference.py videos_dir/ | |
| # Save all reference videos to a specific directory | |
| compute_reference.py videos_dir/ --output-path /path/to/output/ | |
| # Compute depth reference videos using Video-Depth-Anything | |
| python compute_reference.py \ | |
| /home/xinyuy/dataset_processing/OpenVE-3M/construct_video \ | |
| --depth-repo-path Video-Depth-Anything \ | |
| --encoder vitl | |
| /mnt/data/xinyuy/datasets/OpenHumanVid/good | |
| python compute_reference.py /home/xinyuy/dataset_processing/flexcombine-bench/datasets/v0/pose_original \ | |
| --type pose --dwpose-repo-path DWPose \ | |
| --output-path /home/xinyuy/dataset_processing/flexcombine-bench/datasets/v0/pose_original/pose_raw | |
| # Use a smaller depth model for faster inference | |
| compute_reference.py videos_dir/ --type depth --encoder vits --depth-repo-path /path/to/Video-Depth-Anything | |
| # Compute pose reference videos using DWPose | |
| python compute_reference.py /home/xinyuy/dataset_processing/flexcombine-bench/datasets/multi/pose+subject --type pose --dwpose-repo-path DWPose | |
| # Extract one representative frame per video for subject-to-video training | |
| compute_reference.py video.mp4 --type subject | |
| compute_reference.py videos_dir/ --type subject | |
| """ | |
| # Standard library imports | |
| import multiprocessing as mp | |
| from pathlib import Path | |
| from typing import TYPE_CHECKING | |
| # Third-party imports | |
| import cv2 | |
| import numpy as np | |
| import torch | |
| import torchvision.transforms.functional as TF # noqa: N812 | |
| import typer | |
| from rich.console import Console | |
| from rich.progress import ( | |
| BarColumn, | |
| MofNCompleteColumn, | |
| Progress, | |
| SpinnerColumn, | |
| TextColumn, | |
| TimeElapsedColumn, | |
| TimeRemainingColumn, | |
| ) | |
| from transformers.utils.logging import disable_progress_bar | |
| # Local imports | |
| from video_utils import read_video, save_video | |
| if TYPE_CHECKING: | |
| from annotator.dwpose import DWposeDetector | |
| from video_depth_anything.video_depth import VideoDepthAnything | |
| # Initialize console and disable progress bars | |
| console = Console() | |
| disable_progress_bar() | |
| def compute_subject_reference( | |
| video: torch.Tensor, | |
| frame_idx: int | None = None, | |
| ) -> tuple[torch.Tensor, np.ndarray]: | |
| """Extract a single reference frame from a video for subject-to-video conditioning. | |
| Picks a frame from the middle third of the video by default to avoid scene transitions | |
| and ensure the subject is clearly visible. | |
| Args: | |
| video: Video tensor of shape [T, C, H, W] in [0, 1] range | |
| frame_idx: Specific frame index to extract. Defaults to a random middle frame. | |
| Returns: | |
| Tuple of: | |
| - 1-frame video tensor of shape [1, C, H, W] in [0, 1] range | |
| - Raw frame as [1, H, W, C] uint8 numpy array | |
| """ | |
| T = video.shape[0] | |
| if frame_idx is None: | |
| start = max(1, T // 3) | |
| end = min(T - 1, 2 * T // 3) | |
| frame_idx = int(torch.randint(start, max(start + 1, end + 1), (1,)).item()) | |
| frame = video[frame_idx : frame_idx + 1] | |
| raw = (frame.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8) | |
| return frame, raw | |
| def compute_reference( | |
| images: torch.Tensor, | |
| ) -> tuple[torch.Tensor, np.ndarray]: | |
| """Compute Canny edge detection on a batch of images. | |
| Args: | |
| images: Batch of images tensor of shape [B, C, H, W] | |
| Returns: | |
| Tuple of: | |
| - 3-channel edge tensor of shape [B, 3, H, W] in [0, 255] range | |
| - Raw binary edge masks as [B, H, W] uint8 numpy array (0 or 255) | |
| """ | |
| # Convert to grayscale if needed | |
| if images.shape[1] == 3: | |
| images = TF.rgb_to_grayscale(images) | |
| # Ensure images are in [0, 1] range | |
| if images.max() > 1.0: | |
| images = images / 255.0 | |
| # Compute Canny edges | |
| edge_masks = [] | |
| for image in images: | |
| # Convert to numpy for OpenCV | |
| image_np = (image.squeeze().cpu().numpy() * 255).astype("uint8") | |
| # Apply Canny edge detection | |
| edges = cv2.Canny( | |
| image_np, | |
| threshold1=100, | |
| threshold2=200, | |
| ) | |
| # Convert back to tensor | |
| edge_mask = torch.from_numpy(edges).float() | |
| edge_masks.append(edge_mask) | |
| edges = torch.stack(edge_masks) # [B, H, W] | |
| raw = edges.cpu().numpy().astype(np.uint8) # [B, H, W] uint8, values 0 or 255 | |
| edges = torch.stack([edges] * 3, dim=1) # Convert to 3-channel | |
| return edges, raw | |
| _DEPTH_MODEL_CONFIGS = { | |
| "vits": {"encoder": "vits", "features": 64, "out_channels": [48, 96, 192, 384]}, | |
| "vitb": {"encoder": "vitb", "features": 128, "out_channels": [96, 192, 384, 768]}, | |
| "vitl": {"encoder": "vitl", "features": 256, "out_channels": [256, 512, 1024, 1024]}, | |
| } | |
| _DEPTH_ENCODER_TO_HF_SIZE = {"vits": "Small", "vitb": "Base", "vitl": "Large"} | |
| def load_depth_model(encoder: str = "vitl", device: str = "cuda", repo_path: Path | None = None) -> "VideoDepthAnything": | |
| """Load a Video-Depth-Anything model from a local checkpoint. | |
| Args: | |
| encoder: Encoder size, one of 'vits', 'vitb', 'vitl' | |
| device: Device to load the model on | |
| repo_path: Path to a local clone of https://github.com/DepthAnything/Video-Depth-Anything. | |
| Expects checkpoints at {repo_path}/checkpoints/video_depth_anything_{encoder}.pth | |
| Returns: | |
| Loaded VideoDepthAnything model in eval mode | |
| """ | |
| import sys | |
| if repo_path is None: | |
| raise ValueError( | |
| "--depth-repo-path is required. Clone the repo first:\n" | |
| " git clone https://github.com/DepthAnything/Video-Depth-Anything /path/to/Video-Depth-Anything" | |
| ) | |
| repo_root = str(repo_path.resolve()) | |
| if repo_root not in sys.path: | |
| sys.path.insert(0, repo_root) | |
| try: | |
| from video_depth_anything.video_depth import VideoDepthAnything | |
| except ImportError as e: | |
| raise ImportError( | |
| f"Could not import video_depth_anything from {repo_root}.\n" | |
| f"Original error: {e}\n" | |
| "Make sure --depth-repo-path points to the root of the cloned repo " | |
| "and all dependencies are installed (pip install -r requirements.txt inside the repo)." | |
| ) from e | |
| checkpoint_path = repo_path.resolve() / "checkpoints" / f"video_depth_anything_{encoder}.pth" | |
| if not checkpoint_path.exists(): | |
| raise FileNotFoundError( | |
| f"Checkpoint not found: {checkpoint_path}\n" | |
| "Download it from https://huggingface.co/depth-anything/Video-Depth-Anything-" | |
| f"{_DEPTH_ENCODER_TO_HF_SIZE[encoder]} and place it in {checkpoint_path.parent}/" | |
| ) | |
| console.print(f"Loading depth model from [bold blue]{checkpoint_path}[/]...") | |
| model = VideoDepthAnything(**_DEPTH_MODEL_CONFIGS[encoder]) | |
| model.load_state_dict(torch.load(checkpoint_path, map_location="cpu"), strict=True) | |
| model = model.to(device).eval() | |
| console.print(f"[bold green]✓[/] Depth model loaded on [bold]{device}[/]") | |
| return model | |
| def compute_depth_reference( | |
| video: torch.Tensor, | |
| fps: float, | |
| model: "VideoDepthAnything", | |
| input_size: int = 518, | |
| device: str = "cuda", | |
| ) -> tuple[torch.Tensor, np.ndarray]: | |
| """Compute depth maps for a video using Video-Depth-Anything. | |
| Args: | |
| video: Video tensor of shape [T, C, H, W] in [0, 1] range | |
| fps: Frames per second of the video | |
| model: Loaded VideoDepthAnything model | |
| input_size: Input resolution for the depth model (default 518) | |
| device: Device to run inference on | |
| Returns: | |
| Tuple of: | |
| - Depth video tensor of shape [T, 3, H, W] in [0, 1] range (normalized for display) | |
| - Raw metric depth as [T, H, W] float32 numpy array (model output, before normalization) | |
| """ | |
| # Convert [T, C, H, W] float [0,1] → [T, H, W, C] uint8 numpy | |
| frames_np = (video.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8) | |
| with torch.inference_mode(): | |
| depths, _ = model.infer_video_depth(frames_np, fps, input_size=input_size, device=device, fp32=False) | |
| raw_depths = depths.copy() # [T, H, W] float32, raw metric depth from model | |
| # depths: [T, H, W] float, normalize per-video to [0, 1] | |
| depth_min, depth_max = float(depths.min()), float(depths.max()) | |
| if depth_max > depth_min: | |
| depths = (depths - depth_min) / (depth_max - depth_min) | |
| else: | |
| depths = np.zeros_like(depths) | |
| depth_tensor = torch.from_numpy(depths).float() # [T, H, W] | |
| return depth_tensor.unsqueeze(1).expand(-1, 3, -1, -1).contiguous(), raw_depths # [T, 3, H, W] | |
| def load_pose_model(device: str = "cuda", repo_path: Path | None = None) -> "DWposeDetector": | |
| """Load a DWPose model from a local checkout of the DWPose repository. | |
| Args: | |
| device: Device to run inference on ('cuda' or 'cpu'). Note that the underlying | |
| Wholebody ONNX sessions always attempt CUDAExecutionProvider when a GPU | |
| is available; the device argument is passed for future compatibility. | |
| repo_path: Path to a local clone of https://github.com/IDEA-Research/DWPose. | |
| Expects the ControlNet directory at {repo_path}/ControlNet-v1-1-nightly/ | |
| and ONNX checkpoints at | |
| {repo_path}/ControlNet-v1-1-nightly/annotator/ckpts/yolox_l.onnx and | |
| {repo_path}/ControlNet-v1-1-nightly/annotator/ckpts/dw-ll_ucoco_384.onnx | |
| Returns: | |
| Loaded DWposeDetector ready for inference | |
| """ | |
| import os | |
| import sys | |
| if repo_path is None: | |
| raise ValueError( | |
| "--dwpose-repo-path is required. Clone the repo first:\n" | |
| " git clone https://github.com/IDEA-Research/DWPose /path/to/DWPose" | |
| ) | |
| controlnet_root = str((repo_path / "ControlNet-v1-1-nightly").resolve()) | |
| det_ckpt = Path(controlnet_root) / "annotator" / "ckpts" / "yolox_l.onnx" | |
| pose_ckpt = Path(controlnet_root) / "annotator" / "ckpts" / "dw-ll_ucoco_384.onnx" | |
| for ckpt, name in [(det_ckpt, "yolox_l.onnx"), (pose_ckpt, "dw-ll_ucoco_384.onnx")]: | |
| if not ckpt.exists(): | |
| raise FileNotFoundError( | |
| f"Checkpoint not found: {ckpt}\n" | |
| "Download ONNX checkpoints from https://huggingface.co/yzd-v/DWPose and " | |
| f"place them in {ckpt.parent}/" | |
| ) | |
| if controlnet_root not in sys.path: | |
| sys.path.insert(0, controlnet_root) | |
| try: | |
| from annotator.dwpose import DWposeDetector | |
| except ImportError as e: | |
| raise ImportError( | |
| f"Could not import annotator.dwpose from {controlnet_root}.\n" | |
| f"Original error: {e}\n" | |
| "Make sure --dwpose-repo-path points to the root of the cloned DWPose repo " | |
| "and onnxruntime is installed (pip install onnxruntime-gpu)." | |
| ) from e | |
| console.print(f"Loading DWPose model from [bold blue]{controlnet_root}[/]...") | |
| # Wholebody uses relative paths ('annotator/ckpts/...') so we must run from the ControlNet root | |
| orig_dir = os.getcwd() | |
| try: | |
| os.chdir(controlnet_root) | |
| detector = DWposeDetector() | |
| finally: | |
| os.chdir(orig_dir) | |
| console.print(f"[bold green]✓[/] DWPose model loaded on [bold]{device}[/]") | |
| return detector | |
| def compute_pose_reference( | |
| video: torch.Tensor, | |
| model: "DWposeDetector", | |
| draw_face: bool = True, | |
| ) -> tuple[torch.Tensor, list[dict]]: | |
| """Compute pose maps for a video using DWPose. | |
| Args: | |
| video: Video tensor of shape [T, C, H, W] in [0, 1] range | |
| model: Loaded DWposeDetector | |
| draw_face: Whether to draw face keypoints on the pose canvas | |
| Returns: | |
| Tuple of: | |
| - Pose video tensor of shape [T, 3, H, W] in [0, 1] range (rendered skeleton canvas) | |
| - Raw keypoints: list of T dicts, each with: | |
| 'keypoints': [num_people, 134, 2] float32 pixel coords (x, y) | |
| 'scores': [num_people, 134] float32 confidence scores | |
| """ | |
| import annotator.dwpose.util as dwpose_util | |
| from annotator.dwpose import draw_pose | |
| _orig_draw_facepose = dwpose_util.draw_facepose | |
| if not draw_face: | |
| dwpose_util.draw_facepose = lambda canvas, faces: canvas | |
| # Convert [T, C, H, W] float [0,1] → [T, H, W, C] uint8 numpy | |
| frames_np = (video.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8) | |
| pose_frames = [] | |
| raw_keypoints = [] | |
| try: | |
| for frame in frames_np: | |
| H, W, _ = frame.shape | |
| with torch.no_grad(): | |
| candidate, subset = model.pose_estimation(frame) | |
| # Save raw keypoints and scores (pixel coords, before normalization) | |
| raw_keypoints.append({ | |
| "keypoints": candidate.copy().astype(np.float32), # [num_people, 134, 2] | |
| "scores": subset.copy().astype(np.float32), # [num_people, 134] | |
| }) | |
| # Replicate DWposeDetector.__call__ to draw canvas without double inference | |
| nums, keys, locs = candidate.shape | |
| candidate_norm = candidate.copy().astype(float) | |
| candidate_norm[..., 0] /= float(W) | |
| candidate_norm[..., 1] /= float(H) | |
| body = candidate_norm[:, :18].copy().reshape(nums * 18, locs) | |
| score = subset[:, :18].copy() | |
| for i in range(len(score)): | |
| for j in range(len(score[i])): | |
| score[i][j] = int(18 * i + j) if score[i][j] > 0.3 else -1 | |
| un_visible = subset < 0.3 | |
| candidate_norm[un_visible] = -1 | |
| faces = candidate_norm[:, 24:92] | |
| hands = np.vstack([candidate_norm[:, 92:113], candidate_norm[:, 113:]]) | |
| pose = dict(bodies=dict(candidate=body, subset=score), hands=hands, faces=faces) | |
| pose_frames.append(draw_pose(pose, H, W)) | |
| finally: | |
| dwpose_util.draw_facepose = _orig_draw_facepose | |
| pose_np = np.stack(pose_frames, axis=0) # [T, H, W, 3] | |
| pose_tensor = torch.from_numpy(pose_np).float() / 255.0 | |
| return pose_tensor.permute(0, 3, 1, 2).contiguous(), raw_keypoints # [T, 3, H, W] | |
| def process_single_file( | |
| input_path: Path, | |
| reference_type: str = "canny", | |
| input_size: int = 518, | |
| batch_size: int = 100, | |
| num_frames: int | None = None, | |
| draw_face: bool = True, | |
| output_dir: Path | None = None, | |
| save_raw: bool = False, | |
| model: "VideoDepthAnything | DWposeDetector | None" = None, | |
| ) -> None: | |
| """Process a single video file and save the reference video. | |
| The output is written to {input_stem}_{reference_type}.mp4, either in output_dir | |
| (if provided) or next to the input file. | |
| Args: | |
| input_path: Path to the input video file | |
| reference_type: Type of reference to compute ('canny', 'depth', or 'pose') | |
| input_size: Input resolution for depth model inference | |
| batch_size: Batch size for Canny processing | |
| num_frames: If set, only process the first N frames of the video | |
| output_dir: Directory to save the reference video. If None, saves next to input. | |
| model: Pre-loaded depth or pose model, shared across files to avoid reloading | |
| per video. Required when reference_type is 'depth' or 'pose'. | |
| """ | |
| base_dir = output_dir if output_dir is not None else input_path.parent | |
| base_dir.mkdir(parents=True, exist_ok=True) | |
| output_path = base_dir / f"{input_path.stem}_{reference_type}.mp4" | |
| if save_raw: | |
| raw_dir = base_dir / "raw" | |
| raw_dir.mkdir(parents=True, exist_ok=True) | |
| raw_path = raw_dir / f"{input_path.stem}_{reference_type}.npy" | |
| console.print(f"Processing [bold blue]{input_path}[/] → [bold green]{output_path}[/]") | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| video, fps = read_video(input_path) | |
| if num_frames is not None: | |
| video = video[:num_frames] | |
| if reference_type == "depth": | |
| all_condition, raw_data = compute_depth_reference(video, fps, model, input_size=input_size, device=device) | |
| elif reference_type == "pose": | |
| all_condition, raw_data = compute_pose_reference(video, model, draw_face=draw_face) | |
| elif reference_type == "subject": | |
| all_condition, raw_data = compute_subject_reference(video) | |
| else: | |
| condition_frames = [] | |
| raw_frames = [] | |
| for i in range(0, len(video), batch_size): | |
| batch = video[i : i + batch_size] | |
| cond, raw = compute_reference(batch) | |
| condition_frames.append(cond) | |
| raw_frames.append(raw) | |
| all_condition = torch.cat(condition_frames, dim=0) | |
| raw_data = np.concatenate(raw_frames, axis=0) | |
| save_video(all_condition, output_path.resolve(), fps=fps) | |
| console.print(f"[bold green]✓[/] Saved reference video to [cyan]{output_path}[/]") | |
| if save_raw: | |
| np.save(raw_path, raw_data, allow_pickle=True) | |
| console.print(f"[bold green]✓[/] Saved raw data to [cyan]{raw_path}[/]") | |
| def _process_shard( | |
| gpu_id: int | None, | |
| video_files: list[Path], | |
| queue: "mp.Queue", | |
| reference_type: str, | |
| encoder: str, | |
| input_size: int, | |
| batch_size: int, | |
| depth_repo_path: Path | None, | |
| dwpose_repo_path: Path | None, | |
| num_frames: int | None, | |
| draw_face: bool, | |
| output_path: Path | None, | |
| save_raw: bool, | |
| ) -> None: | |
| """Worker process entrypoint for directory mode. | |
| Pins the process to a single GPU (via CUDA_VISIBLE_DEVICES so both torch and the | |
| DWPose ONNX sessions, which always target 'cuda:0', land on the right device), | |
| loads one model instance, and processes an assigned shard of video files. Progress | |
| and failures are reported back to the parent process through `queue`. | |
| """ | |
| import os | |
| if gpu_id is not None: | |
| os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id) | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| tag = f"[GPU {gpu_id}]" if gpu_id is not None else "[CPU]" | |
| model = None | |
| if reference_type == "depth": | |
| model = load_depth_model(encoder=encoder, device=device, repo_path=depth_repo_path) | |
| elif reference_type == "pose": | |
| model = load_pose_model(device=device, repo_path=dwpose_repo_path) | |
| for video_file in video_files: | |
| console.print(f"{tag} Processing [bold blue]{video_file.name}[/]") | |
| try: | |
| process_single_file( | |
| input_path=video_file, | |
| reference_type=reference_type, | |
| input_size=input_size, | |
| batch_size=batch_size, | |
| num_frames=num_frames, | |
| draw_face=draw_face, | |
| output_dir=output_path, | |
| save_raw=save_raw, | |
| model=model, | |
| ) | |
| queue.put(("done", video_file.name, None)) | |
| except Exception as e: # noqa: BLE001 | |
| console.print(f"{tag} [bold red]Failed[/] {video_file.name}: {e}") | |
| queue.put(("error", video_file.name, str(e))) | |
| app = typer.Typer( | |
| pretty_exceptions_enable=False, | |
| no_args_is_help=True, | |
| help="Compute reference videos for IC-LoRA training.", | |
| ) | |
| def main( | |
| input_path: Path = typer.Argument( # noqa: B008 | |
| ..., | |
| help="Path to input video file or directory containing .mp4 files", | |
| exists=True, | |
| ), | |
| output_path: Path | None = typer.Option( # noqa: B008 | |
| None, | |
| "--output-path", | |
| "-o", | |
| help="Directory to save reference videos. If omitted, saves next to each input file.", | |
| ), | |
| override: bool = typer.Option( | |
| False, | |
| "--override", | |
| help="Whether to override existing reference video files", | |
| ), | |
| reference_type: str = typer.Option( | |
| "canny", | |
| "--type", | |
| "-t", | |
| help="Type of reference to compute: 'canny' for Canny edge detection, 'depth' for Video-Depth-Anything, 'pose' for DWPose skeleton, 'subject' for a single representative frame (subject-to-video)", | |
| ), | |
| encoder: str = typer.Option( | |
| "vitl", | |
| "--encoder", | |
| "-e", | |
| help="Encoder size for depth model: 'vits' (28M), 'vitb' (113M), 'vitl' (382M, default). Only used with --type depth.", | |
| ), | |
| input_size: int = typer.Option( | |
| 518, | |
| "--input-size", | |
| help="Input resolution for depth model inference. Only used with --type depth.", | |
| ), | |
| batch_size: int = typer.Option( | |
| 100, | |
| "--batch-size", | |
| help="Batch size for Canny edge processing. Not used for depth.", | |
| ), | |
| num_workers: int | None = typer.Option( | |
| None, | |
| "--num-workers", | |
| "-w", | |
| help="Directory mode only: number of parallel worker processes. Each worker loads its " | |
| "own model instance and is pinned to one GPU. Defaults to the number of visible GPUs " | |
| "(or 1 on CPU). Keep <= GPU count for --type depth/pose to avoid oversubscribing a GPU.", | |
| ), | |
| depth_repo_path: Path | None = typer.Option( # noqa: B008 | |
| None, | |
| "--depth-repo-path", | |
| help="Path to a local clone of https://github.com/DepthAnything/Video-Depth-Anything. " | |
| "Required when using --type depth.", | |
| ), | |
| dwpose_repo_path: Path | None = typer.Option( # noqa: B008 | |
| None, | |
| "--dwpose-repo-path", | |
| help="Path to a local clone of https://github.com/IDEA-Research/DWPose. " | |
| "Expects ControlNet-v1-1-nightly/annotator/ckpts/{yolox_l,dw-ll_ucoco_384}.onnx. " | |
| "Required when using --type pose.", | |
| ), | |
| num_frames: int | None = typer.Option( | |
| None, | |
| "--num-frames", | |
| "-n", | |
| help="If set, only process the first N frames of each video.", | |
| ), | |
| draw_face: bool = typer.Option( | |
| True, | |
| "--draw-face/--no-face", | |
| help="Whether to draw face keypoints on the pose canvas. Only used with --type pose.", | |
| ), | |
| save_raw: bool = typer.Option( | |
| False, | |
| "--save-raw", | |
| help="If set, also save raw reference data as .npy files in a 'raw/' subfolder under the output directory.", | |
| ), | |
| ) -> None: | |
| """Compute reference videos for IC-LoRA training. | |
| This script generates reference videos (Canny edge maps, depth maps, or pose maps) for given videos. | |
| Single-file mode: when input_path is a file, the reference is saved as | |
| {original_name}_{type}.mp4 next to the input (or in --output-path if given). | |
| Directory mode: processes all .mp4 files in the directory. | |
| Reference videos are saved next to each input (or in --output-path if given). | |
| Examples: | |
| # Single file — Canny (writes video_canny.mp4 next to input) | |
| compute_reference.py video.mp4 | |
| # Single file — save to specific directory | |
| compute_reference.py video.mp4 --output-path /path/to/output/ | |
| # Single file — depth | |
| compute_reference.py video.mp4 --type depth --depth-repo-path /path/to/Video-Depth-Anything | |
| # Process all videos with Canny edges (default) | |
| compute_reference.py videos_dir/ | |
| # Process all videos, save to a specific directory | |
| compute_reference.py videos_dir/ --output-path /path/to/output/ | |
| # Process all videos with Video-Depth-Anything depth maps | |
| compute_reference.py videos_dir/ --type depth --depth-repo-path /path/to/Video-Depth-Anything | |
| # Use a smaller/faster depth model | |
| compute_reference.py videos_dir/ --type depth --encoder vits --depth-repo-path /path/to/Video-Depth-Anything | |
| # Process all videos with DWPose skeleton maps | |
| compute_reference.py videos_dir/ --type pose --dwpose-repo-path /path/to/DWPose | |
| # Extract a representative frame per video for subject-to-video training | |
| compute_reference.py video.mp4 --type subject | |
| compute_reference.py videos_dir/ --type subject | |
| """ | |
| if reference_type not in ("canny", "depth", "pose", "subject"): | |
| raise typer.BadParameter( | |
| f"Invalid reference type '{reference_type}'. Choose 'canny', 'depth', 'pose', or 'subject'." | |
| ) | |
| if encoder not in ("vits", "vitb", "vitl"): | |
| raise typer.BadParameter(f"Invalid encoder '{encoder}'. Choose 'vits', 'vitb', or 'vitl'.") | |
| if reference_type == "depth" and depth_repo_path is None: | |
| raise typer.BadParameter( | |
| "--depth-repo-path is required when using --type depth.\n\n" | |
| "Clone the repo first:\n" | |
| " git clone https://github.com/DepthAnything/Video-Depth-Anything /path/to/Video-Depth-Anything" | |
| ) | |
| if reference_type == "pose" and dwpose_repo_path is None: | |
| raise typer.BadParameter( | |
| "--dwpose-repo-path is required when using --type pose.\n\n" | |
| "Clone the repo first:\n" | |
| " git clone https://github.com/IDEA-Research/DWPose /path/to/DWPose" | |
| ) | |
| if input_path.is_file(): | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model = None | |
| if reference_type == "depth": | |
| model = load_depth_model(encoder=encoder, device=device, repo_path=depth_repo_path) | |
| elif reference_type == "pose": | |
| model = load_pose_model(device=device, repo_path=dwpose_repo_path) | |
| process_single_file( | |
| input_path=input_path, | |
| reference_type=reference_type, | |
| input_size=input_size, | |
| batch_size=batch_size, | |
| num_frames=num_frames, | |
| draw_face=draw_face, | |
| output_dir=output_path, | |
| save_raw=save_raw, | |
| model=model, | |
| ) | |
| return | |
| # Directory mode: process all .mp4 files, sharded across worker processes/GPUs | |
| video_files = sorted(input_path.glob("*.mp4")) | |
| if not video_files: | |
| raise typer.BadParameter(f"No .mp4 files found in {input_path}") | |
| if output_path is not None and not override: | |
| remaining = [] | |
| for f in video_files: | |
| out = output_path / f"{f.stem}_{reference_type}.mp4" | |
| if out.exists(): | |
| console.print(f"[yellow]Skipping[/] {f.name} (output already exists)") | |
| else: | |
| remaining.append(f) | |
| video_files = remaining | |
| if not video_files: | |
| console.print("[bold green]Nothing to do[/] — all outputs already exist.") | |
| return | |
| console.print(f"Found [bold]{len(video_files)}[/] video(s) to process in [bold blue]{input_path}[/]") | |
| available_gpus = torch.cuda.device_count() if torch.cuda.is_available() else 0 | |
| workers = num_workers if num_workers is not None else max(available_gpus, 1) | |
| workers = max(1, min(workers, len(video_files))) | |
| gpu_ids = [i % available_gpus for i in range(workers)] if available_gpus > 0 else [None] * workers | |
| if workers > 1: | |
| console.print(f"Parallelizing across [bold]{workers}[/] worker(s), GPUs: {gpu_ids}") | |
| shards = [video_files[i::workers] for i in range(workers)] | |
| ctx = mp.get_context("spawn") | |
| queue: mp.Queue = ctx.Queue() | |
| processes = [] | |
| for gpu_id, shard in zip(gpu_ids, shards): | |
| if not shard: | |
| continue | |
| process = ctx.Process( | |
| target=_process_shard, | |
| args=( | |
| gpu_id, | |
| shard, | |
| queue, | |
| reference_type, | |
| encoder, | |
| input_size, | |
| batch_size, | |
| depth_repo_path, | |
| dwpose_repo_path, | |
| num_frames, | |
| draw_face, | |
| output_path, | |
| save_raw, | |
| ), | |
| ) | |
| process.start() | |
| processes.append(process) | |
| errors = [] | |
| with Progress( | |
| SpinnerColumn(), | |
| TextColumn("[progress.description]{task.description}"), | |
| BarColumn(), | |
| MofNCompleteColumn(), | |
| TimeElapsedColumn(), | |
| TimeRemainingColumn(), | |
| console=console, | |
| ) as progress: | |
| task = progress.add_task("Processing videos...", total=len(video_files)) | |
| for _ in range(len(video_files)): | |
| status, name, err = queue.get() | |
| if status == "error": | |
| errors.append((name, err)) | |
| progress.advance(task) | |
| for process in processes: | |
| process.join() | |
| if errors: | |
| console.print(f"[bold red]{len(errors)} video(s) failed:[/]") | |
| for name, err in errors: | |
| console.print(f" [red]{name}[/]: {err}") | |
| if __name__ == "__main__": | |
| app() | |
Xet Storage Details
- Size:
- 29.5 kB
- Xet hash:
- be47bc7ce587a57a6e1054b8a60bae5e78a908ddbe1272a56a0a8577b90b5926
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.