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) @staticmethod 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)] @staticmethod 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, }