Instructions to use Kry4ta1/Effecteraser-VOR-Inference with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use Kry4ta1/Effecteraser-VOR-Inference with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("Kry4ta1/Effecteraser-VOR-Inference", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download src/videox_fun/data/remove_dataset.py from Kry4ta1/Effecteraser-VOR-Inference: direct link, hf CLI and curl.
- Browser
- Download file 20.3 kB
-
https://huggingface.co/Kry4ta1/Effecteraser-VOR-Inference/resolve/main/src/videox_fun/data/remove_dataset.py
- Command line
-
hf download hf://Kry4ta1/Effecteraser-VOR-Inference/src/videox_fun/data/remove_dataset.py
-
curl -L -o remove_dataset.py https://huggingface.co/Kry4ta1/Effecteraser-VOR-Inference/resolve/main/src/videox_fun/data/remove_dataset.py
20.3 kB
| import random | |
| from pathlib import Path | |
| from typing import Dict, List, Optional, Sequence, Tuple, Union | |
| import cv2 | |
| import torch | |
| import torch.nn.functional as F | |
| from torch.utils.data import Dataset | |
| VIDEO_EXTENSIONS = {".mp4", ".mov", ".avi", ".mkv", ".webm"} | |
| class ShortVideoError(ValueError): | |
| pass | |
| def orientation_aware_size( | |
| source_width: int, | |
| source_height: int, | |
| landscape_height: int, | |
| landscape_width: int, | |
| ) -> Tuple[int, int]: | |
| if min(source_width, source_height, landscape_height, landscape_width) <= 0: | |
| raise ValueError( | |
| "Source and target dimensions must be positive, got " | |
| f"source={source_height}x{source_width}, " | |
| f"target={landscape_height}x{landscape_width}" | |
| ) | |
| if source_height > source_width: | |
| return landscape_width, landscape_height | |
| return landscape_height, landscape_width | |
| def load_scene_names(scene_list_path: Optional[str]) -> Optional[List[str]]: | |
| if scene_list_path is None: | |
| return None | |
| path = Path(scene_list_path).expanduser().resolve() | |
| names = [] | |
| if path.is_file(): | |
| for line in path.read_text(encoding="utf-8").splitlines(): | |
| value = line.strip() | |
| if value and not value.startswith("#"): | |
| names.append(Path(value).stem) | |
| elif path.is_dir(): | |
| names = [ | |
| child.stem | |
| for child in sorted(path.iterdir()) | |
| if not child.name.startswith(".") | |
| ] | |
| else: | |
| raise FileNotFoundError(f"Scene list file/directory does not exist: {path}") | |
| names = list(dict.fromkeys(name for name in names if name)) | |
| if not names: | |
| raise ValueError(f"Scene list contains no valid names: {path}") | |
| return names | |
| def discover_video_triplets( | |
| dataset_root: str, | |
| include_substrings: Optional[Sequence[str]] = None, | |
| exclude_substring: Optional[str] = None, | |
| allow_empty_after_filter: bool = False, | |
| ) -> List[Tuple[Path, Path, Path]]: | |
| root = Path(dataset_root).expanduser().resolve() | |
| folder_names = ("BG", "MASK", "FG_BG") | |
| folders = {name: root / name for name in folder_names} | |
| missing_folders = [str(path) for path in folders.values() if not path.is_dir()] | |
| if missing_folders: | |
| raise FileNotFoundError(f"Missing dataset folders: {missing_folders}") | |
| indexed: Dict[str, Dict[str, Path]] = {} | |
| for folder_name, folder in folders.items(): | |
| files = { | |
| path.relative_to(folder).as_posix(): path | |
| for path in folder.rglob("*") | |
| if ( | |
| path.is_file() | |
| and path.suffix.lower() in VIDEO_EXTENSIONS | |
| and ( | |
| not include_substrings | |
| or any(value in path.name for value in include_substrings) | |
| ) | |
| and ( | |
| not exclude_substring | |
| or exclude_substring not in path.name | |
| ) | |
| ) | |
| } | |
| if not files and not allow_empty_after_filter: | |
| raise ValueError(f"No supported videos found in {folder}") | |
| indexed[folder_name] = files | |
| expected = set(indexed["BG"]) | |
| errors = [] | |
| for folder_name in ("MASK", "FG_BG"): | |
| actual = set(indexed[folder_name]) | |
| missing = sorted(expected - actual) | |
| extra = sorted(actual - expected) | |
| if missing or extra: | |
| errors.append(f"{folder_name}: missing={missing[:10]}, extra={extra[:10]}") | |
| if errors: | |
| raise ValueError("BG/MASK/FG_BG filenames do not match exactly. " + " | ".join(errors)) | |
| triplets = [ | |
| (indexed["BG"][name], indexed["MASK"][name], indexed["FG_BG"][name]) | |
| for name in sorted(expected) | |
| ] | |
| if not triplets and not allow_empty_after_filter: | |
| raise ValueError( | |
| f"No video triplets remain after filtering {root}: " | |
| f"include={list(include_substrings) if include_substrings else None}, " | |
| f"exclude={exclude_substring!r}" | |
| ) | |
| return triplets | |
| def muse_temporal_union(mask: torch.Tensor, temporal_ratio: int = 4) -> torch.Tensor: | |
| if mask.ndim != 5 or mask.shape[1] < 1: | |
| raise ValueError(f"Expected mask features [B,C,T,H,W], got {tuple(mask.shape)}") | |
| if temporal_ratio < 1: | |
| raise ValueError(f"temporal_ratio must be positive, got {temporal_ratio}") | |
| if mask.shape[2] == 1 or temporal_ratio == 1: | |
| return mask | |
| first = mask[:, :, :1] | |
| remaining = mask[:, :, 1:] | |
| pad = (-remaining.shape[2]) % temporal_ratio | |
| if pad: | |
| remaining = F.pad(remaining, (0, 0, 0, 0, 0, pad), value=0) | |
| b, c, t, h, w = remaining.shape | |
| remaining = remaining.reshape(b, c, t // temporal_ratio, temporal_ratio, h, w) | |
| return torch.cat([first, remaining.amax(dim=3)], dim=2) | |
| def degrade_mask( | |
| mask: torch.Tensor, | |
| frame_drop_probability: float = 0.7, | |
| frame_drop_rate_min: float = 0.2, | |
| frame_drop_rate_max: float = 0.99, | |
| morphology_probability: float = 0.5, | |
| morphology_kernel_sizes: Sequence[int] = (3, 5, 7), | |
| bbox_probability: float = 0.25, | |
| generator: Optional[torch.Generator] = None, | |
| return_metadata: bool = False, | |
| ) -> Union[torch.Tensor, Tuple[torch.Tensor, List[Dict[str, object]]]]: | |
| if mask.ndim != 5 or mask.shape[1] != 1: | |
| raise ValueError(f"Expected mask [B,1,T,H,W], got {tuple(mask.shape)}") | |
| if not 0 <= frame_drop_rate_min <= frame_drop_rate_max <= 1: | |
| raise ValueError( | |
| "Frame drop rates must satisfy 0 <= min <= max <= 1, got " | |
| f"{frame_drop_rate_min}, {frame_drop_rate_max}" | |
| ) | |
| if not morphology_kernel_sizes or any( | |
| kernel < 1 or kernel % 2 == 0 for kernel in morphology_kernel_sizes | |
| ): | |
| raise ValueError( | |
| "Morphology kernel sizes must be non-empty positive odd integers, got " | |
| f"{tuple(morphology_kernel_sizes)}" | |
| ) | |
| degraded = (mask > 0.5).to(mask.dtype) | |
| device = degraded.device | |
| metadata = [] | |
| def draw() -> float: | |
| return torch.rand((), generator=generator, device=device).item() | |
| for batch_index in range(degraded.shape[0]): | |
| sample = degraded[batch_index : batch_index + 1] | |
| sample_metadata = { | |
| "sample_index": batch_index, | |
| "operations": [], | |
| "input_mask_ratio": sample.float().mean().item(), | |
| } | |
| if draw() < frame_drop_probability: | |
| drop_rate = frame_drop_rate_min + ( | |
| frame_drop_rate_max - frame_drop_rate_min | |
| ) * draw() | |
| keep = ( | |
| torch.rand( | |
| (1, 1, sample.shape[2], 1, 1), | |
| generator=generator, | |
| device=device, | |
| ) | |
| >= drop_rate | |
| ).to(sample.dtype) | |
| if keep.sum() == 0: | |
| keep[:, :, int(draw() * sample.shape[2]) % sample.shape[2]] = 1 | |
| sample = sample * keep | |
| sample_metadata["operations"].append("frame_dropout") | |
| sample_metadata["sampled_frame_drop_rate"] = drop_rate | |
| sample_metadata["actual_frame_drop_rate"] = 1.0 - keep.float().mean().item() | |
| if draw() < morphology_probability: | |
| kernel_index = int(draw() * len(morphology_kernel_sizes)) % len( | |
| morphology_kernel_sizes | |
| ) | |
| kernel = int(morphology_kernel_sizes[kernel_index]) | |
| padding = kernel // 2 | |
| if draw() < 0.5: | |
| sample = F.max_pool3d(sample, (1, kernel, kernel), stride=1, padding=(0, padding, padding)) | |
| morphology_type = "dilation" | |
| else: | |
| sample = -F.max_pool3d( | |
| -sample, | |
| (1, kernel, kernel), | |
| stride=1, | |
| padding=(0, padding, padding), | |
| ) | |
| morphology_type = "erosion" | |
| sample_metadata["operations"].append(morphology_type) | |
| sample_metadata["morphology_kernel"] = kernel | |
| if draw() < bbox_probability: | |
| boxed = torch.zeros_like(sample) | |
| for frame_index in range(sample.shape[2]): | |
| coordinates = torch.nonzero(sample[0, 0, frame_index] > 0.5, as_tuple=False) | |
| if coordinates.numel() == 0: | |
| continue | |
| top_left = coordinates.amin(dim=0) | |
| bottom_right = coordinates.amax(dim=0) | |
| boxed[ | |
| 0, | |
| 0, | |
| frame_index, | |
| top_left[0] : bottom_right[0] + 1, | |
| top_left[1] : bottom_right[1] + 1, | |
| ] = 1 | |
| sample = boxed | |
| sample_metadata["operations"].append("bbox_fit") | |
| degraded[batch_index : batch_index + 1] = sample | |
| sample_metadata["output_mask_ratio"] = sample.float().mean().item() | |
| metadata.append(sample_metadata) | |
| if return_metadata: | |
| return degraded, metadata | |
| return degraded | |
| def derive_side_effect_mask( | |
| foreground_background: torch.Tensor, | |
| background: torch.Tensor, | |
| object_mask: torch.Tensor, | |
| difference_threshold: float = 0.05, | |
| ) -> torch.Tensor: | |
| if foreground_background.shape != background.shape: | |
| raise ValueError( | |
| "FG_BG and BG must have identical shapes, got " | |
| f"{tuple(foreground_background.shape)} and {tuple(background.shape)}" | |
| ) | |
| if foreground_background.ndim != 5 or object_mask.ndim != 5: | |
| raise ValueError("Expected video and mask tensors in [B,C,T,H,W] layout") | |
| difference = (foreground_background.float() - background.float()).abs().mean(dim=1, keepdim=True) | |
| side_effect = difference > difference_threshold | |
| return (side_effect & (object_mask > 0.5).logical_not()).to(background.dtype) | |
| class RemoveTripletDataset(Dataset): | |
| def __init__( | |
| self, | |
| dataset_root: str, | |
| num_frames: int = 81, | |
| frame_stride: int = 1, | |
| height: int = 480, | |
| width: int = 832, | |
| seed: int = 0, | |
| extra_dataset_root: Optional[str] = None, | |
| extra_dataset_root2: Optional[str] = None, | |
| extra_dataset_roots: Optional[Sequence[str]] = None, | |
| extra_filter_key: Optional[str] = None, | |
| extra_scene_list_path: Optional[str] = None, | |
| ): | |
| if num_frames < 1 or (num_frames - 1) % 4 != 0: | |
| raise ValueError(f"num_frames must be 4n+1 for Wan VAE, got {num_frames}") | |
| if frame_stride < 1: | |
| raise ValueError(f"frame_stride must be positive, got {frame_stride}") | |
| if height % 16 or width % 16: | |
| raise ValueError(f"height and width must be divisible by 16, got {height}x{width}") | |
| self.root = Path(dataset_root).expanduser().resolve() | |
| self.triplets = [] | |
| self.triplet_roots = [] | |
| self.triplet_sources = [] | |
| self.source_summaries = [] | |
| def add_source( | |
| source_name: str, | |
| root_value: str, | |
| include_substrings: Optional[Sequence[str]] = None, | |
| ) -> None: | |
| root = Path(root_value).expanduser().resolve() | |
| source_triplets = discover_video_triplets( | |
| str(root), | |
| include_substrings=include_substrings, | |
| allow_empty_after_filter=True, | |
| ) | |
| self.triplets.extend(source_triplets) | |
| self.triplet_roots.extend([root] * len(source_triplets)) | |
| self.triplet_sources.extend([source_name] * len(source_triplets)) | |
| self.source_summaries.append( | |
| { | |
| "name": source_name, | |
| "root": str(root), | |
| "count": len(source_triplets), | |
| "include": list(include_substrings) if include_substrings else None, | |
| "exclude": None, | |
| } | |
| ) | |
| if include_substrings and not source_triplets: | |
| raise ValueError( | |
| f"Source '{source_name}' at {root} has no triplets whose filename " | |
| f"contains any configured scene/filter value" | |
| ) | |
| self.extra_filter_key = extra_filter_key | |
| self.extra_scene_list_path = extra_scene_list_path | |
| scene_names = load_scene_names(extra_scene_list_path) | |
| extra1_include = ( | |
| scene_names | |
| if scene_names is not None | |
| else [extra_filter_key] | |
| if extra_filter_key | |
| else None | |
| ) | |
| add_source("main", str(self.root)) | |
| if extra_dataset_root is not None: | |
| add_source("extra1", extra_dataset_root, include_substrings=extra1_include) | |
| if extra_dataset_root2 is not None: | |
| add_source("extra2", extra_dataset_root2) | |
| for extra_index, extra_root in enumerate(extra_dataset_roots or [], start=3): | |
| add_source(f"extra{extra_index}", extra_root) | |
| if not self.triplets: | |
| raise ValueError("No video triplets were found in the provided dataset roots") | |
| self.num_frames = num_frames | |
| self.frame_stride = frame_stride | |
| self.height = height | |
| self.width = width | |
| self.seed = seed | |
| self._warned_trimmed_indices = set() | |
| self._invalid_indices = set() | |
| def __len__(self) -> int: | |
| return len(self.triplets) | |
| def _probe_video(path: Path) -> Tuple[int, int, int]: | |
| capture = cv2.VideoCapture(str(path)) | |
| if not capture.isOpened(): | |
| raise RuntimeError(f"Failed to open video: {path}") | |
| try: | |
| frame_count = int(capture.get(cv2.CAP_PROP_FRAME_COUNT)) | |
| width = int(capture.get(cv2.CAP_PROP_FRAME_WIDTH)) | |
| height = int(capture.get(cv2.CAP_PROP_FRAME_HEIGHT)) | |
| finally: | |
| capture.release() | |
| if frame_count <= 0 or width <= 0 or height <= 0: | |
| raise RuntimeError( | |
| f"Invalid video metadata for {path}: frames={frame_count}, size={width}x{height}" | |
| ) | |
| return frame_count, width, height | |
| def _sample_indices( | |
| self, | |
| frame_count: int, | |
| rng: Optional[random.Random] = None, | |
| ) -> List[int]: | |
| required = (self.num_frames - 1) * self.frame_stride + 1 | |
| stride = self.frame_stride if frame_count >= required else 1 | |
| required = (self.num_frames - 1) * stride + 1 | |
| if frame_count < required: | |
| raise ValueError( | |
| f"Video has {frame_count} frames, but at least {self.num_frames} are required" | |
| ) | |
| random_source = rng if rng is not None else random | |
| start = random_source.randint(0, frame_count - required) | |
| return [start + index * stride for index in range(self.num_frames)] | |
| def _decode(path: Path, indices: List[int], size: Tuple[int, int], is_mask: bool) -> torch.Tensor: | |
| target_height, target_width = size | |
| capture = cv2.VideoCapture(str(path)) | |
| if not capture.isOpened(): | |
| raise RuntimeError(f"Failed to open video: {path}") | |
| frames = [] | |
| next_frame_index = None | |
| try: | |
| for frame_index in indices: | |
| if frame_index != next_frame_index: | |
| capture.set(cv2.CAP_PROP_POS_FRAMES, frame_index) | |
| ok, frame = capture.read() | |
| next_frame_index = frame_index + 1 | |
| if not ok: | |
| raise RuntimeError(f"{path}: failed to decode frame {frame_index}") | |
| if is_mask: | |
| frame = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY) | |
| frame = cv2.resize( | |
| frame, | |
| (target_width, target_height), | |
| interpolation=cv2.INTER_NEAREST, | |
| ) | |
| tensor = torch.from_numpy(frame.copy()).unsqueeze(0).float().div_(255.0) | |
| tensor = (tensor > 0.5).float() | |
| else: | |
| frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) | |
| frame = cv2.resize( | |
| frame, | |
| (target_width, target_height), | |
| interpolation=cv2.INTER_LINEAR, | |
| ) | |
| tensor = torch.from_numpy(frame.copy()).permute(2, 0, 1).float() | |
| tensor = tensor.div_(127.5).sub_(1.0) | |
| frames.append(tensor) | |
| finally: | |
| capture.release() | |
| return torch.stack(frames, dim=1).contiguous() | |
| def __getitem__( | |
| self, | |
| index: Union[int, Tuple[int, int]], | |
| ) -> Dict[str, torch.Tensor]: | |
| epoch = None | |
| if isinstance(index, tuple): | |
| if len(index) != 2: | |
| raise ValueError(f"Expected (sample_index, epoch), got {index}") | |
| index, epoch = (int(index[0]), int(index[1])) | |
| else: | |
| index = int(index) | |
| for offset in range(len(self.triplets)): | |
| candidate_index = (index + offset) % len(self.triplets) | |
| if candidate_index in self._invalid_indices: | |
| continue | |
| try: | |
| return self._load_item(candidate_index, epoch) | |
| except ShortVideoError as error: | |
| self._invalid_indices.add(candidate_index) | |
| print( | |
| f"[Dataset:SkipShort] index={candidate_index}, {error}", | |
| flush=True, | |
| ) | |
| raise RuntimeError( | |
| f"No triplet contains at least {self.num_frames} aligned frames" | |
| ) | |
| def _load_item( | |
| self, | |
| index: int, | |
| epoch: Optional[int], | |
| ) -> Dict[str, torch.Tensor]: | |
| background_path, mask_path, foreground_background_path = self.triplets[index] | |
| source_root = self.triplet_roots[index] | |
| source_name = self.triplet_sources[index] | |
| metadata = [ | |
| self._probe_video(background_path), | |
| self._probe_video(mask_path), | |
| self._probe_video(foreground_background_path), | |
| ] | |
| frame_counts = [item[0] for item in metadata] | |
| usable_frame_count = min(frame_counts) | |
| if usable_frame_count < self.num_frames: | |
| raise ShortVideoError( | |
| f"sample={background_path.name}, " | |
| f"BG={frame_counts[0]}, MASK={frame_counts[1]}, " | |
| f"FG_BG={frame_counts[2]}, required={self.num_frames}" | |
| ) | |
| if len(set(frame_counts)) != 1 and index not in self._warned_trimmed_indices: | |
| print( | |
| "[Dataset:Trim] " | |
| f"index={index}, sample={background_path.name}, " | |
| f"using first {usable_frame_count} aligned frames from " | |
| f"BG={frame_counts[0]}, MASK={frame_counts[1]}, FG_BG={frame_counts[2]}", | |
| flush=True, | |
| ) | |
| self._warned_trimmed_indices.add(index) | |
| spatial_sizes = [(item[1], item[2]) for item in metadata] | |
| if len(set(spatial_sizes)) != 1: | |
| raise ValueError( | |
| f"Source size mismatch for {background_path.name}: " | |
| f"BG={spatial_sizes[0]}, MASK={spatial_sizes[1]}, FG_BG={spatial_sizes[2]}" | |
| ) | |
| clip_rng = None | |
| if epoch is not None: | |
| clip_seed = ( | |
| self.seed * 6364136223846793005 | |
| + epoch * 1442695040888963407 | |
| + index | |
| ) % (2**63) | |
| clip_rng = random.Random(clip_seed) | |
| indices = self._sample_indices(usable_frame_count, rng=clip_rng) | |
| source_width, source_height = metadata[2][1], metadata[2][2] | |
| size = orientation_aware_size( | |
| source_width, | |
| source_height, | |
| self.height, | |
| self.width, | |
| ) | |
| return { | |
| "background": self._decode(background_path, indices, size, is_mask=False), | |
| "mask": self._decode(mask_path, indices, size, is_mask=True), | |
| "foreground_background": self._decode( | |
| foreground_background_path, | |
| indices, | |
| size, | |
| is_mask=False, | |
| ), | |
| "sample_name": ( | |
| f"{source_name}:" | |
| f"{background_path.relative_to(source_root / 'BG').as_posix()}" | |
| ), | |
| "dataset_source": source_name, | |
| } | |