Buckets:
| #!/usr/bin/env python3 | |
| """ | |
| Extract subjects from videos using SAM3 text-prompted segmentation. | |
| Usage: | |
| python segment_subjects.py \ | |
| --video_dir /home/xinyuy/dataset_processing/flexcombine-bench/datasets/multi/raw \ | |
| --csv /home/xinyuy/dataset_processing/flexcombine-bench/datasets/multi/raw/captions.csv \ | |
| --output_dir /home/xinyuy/dataset_processing/flexcombine-bench/datasets/multi/inpainting \ | |
| --mode inpaint \ | |
| --box_color 0,0,0 \ | |
| --save_video | |
| python segment_subjects.py \ | |
| --video_dir /home/xinyuy/dataset_processing/flexcombine-bench/datasets/multi/raw \ | |
| --csv /home/xinyuy/dataset_processing/flexcombine-bench/datasets/multi/raw/captions.csv \ | |
| --output_dir /home/xinyuy/dataset_processing/flexcombine-bench/datasets/multi/subject | |
| CSV format (comma-separated, with header): | |
| clip_id,video_path,caption,subject,subject_detailed | |
| 8ecaf6...,/path/to/8ecaf6....mp4,...,man in a striped shirt,... | |
| Required columns: clip_id, subject | |
| Videos are looked up as <video_dir>/<clip_id>.mp4 | |
| Output structure: | |
| output_dir/ | |
| <clip_id>/ | |
| frames/ | |
| frame_0000.png # RGBA PNG, background is transparent | |
| frame_0001.png | |
| ... | |
| """ | |
| import argparse | |
| import csv | |
| import glob | |
| import os | |
| import re | |
| import subprocess | |
| import sys | |
| from pathlib import Path | |
| import cv2 | |
| import numpy as np | |
| from PIL import Image | |
| # --------------------------------------------------------------------------- | |
| # Utilities | |
| # --------------------------------------------------------------------------- | |
| def sanitize_name(s: str) -> str: | |
| return re.sub(r"[^\w\-]", "_", s).strip("_") | |
| def get_video_fps(video_path: str) -> float: | |
| if video_path.endswith(".mp4"): | |
| cap = cv2.VideoCapture(video_path) | |
| fps = cap.get(cv2.CAP_PROP_FPS) | |
| cap.release() | |
| return fps if fps > 0 else 30.0 | |
| return 30.0 | |
| def load_video_frames(video_path: str) -> list: | |
| """Return list of RGB numpy arrays (H, W, 3).""" | |
| if video_path.endswith(".mp4"): | |
| cap = cv2.VideoCapture(video_path) | |
| frames = [] | |
| while True: | |
| ret, frame = cap.read() | |
| if not ret: | |
| break | |
| frames.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) | |
| cap.release() | |
| return frames | |
| else: | |
| paths = glob.glob(os.path.join(video_path, "*.jpg")) | |
| try: | |
| paths.sort(key=lambda p: int(os.path.splitext(os.path.basename(p))[0])) | |
| except ValueError: | |
| paths.sort() | |
| return [np.array(Image.open(p).convert("RGB")) for p in paths] | |
| def mask_to_numpy(mask) -> np.ndarray: | |
| """Convert a SAM3 binary mask (tensor or ndarray) to a bool (H, W) array.""" | |
| if hasattr(mask, "numpy"): | |
| return mask.numpy().astype(bool) | |
| return np.asarray(mask, dtype=bool) | |
| def apply_mask_rgba(frame_rgb: np.ndarray, mask: np.ndarray) -> np.ndarray: | |
| """Return RGBA image: subject pixels have alpha=255, background alpha=0.""" | |
| h, w = frame_rgb.shape[:2] | |
| rgba = np.zeros((h, w, 4), dtype=np.uint8) | |
| rgba[..., :3] = frame_rgb | |
| rgba[..., 3] = (mask * 255).astype(np.uint8) | |
| return rgba | |
| def mask_to_bbox(mask: np.ndarray): | |
| """Return (x0, y0, x1, y1) bounding box of the True region, or None if empty.""" | |
| rows = np.any(mask, axis=1) | |
| cols = np.any(mask, axis=0) | |
| if not rows.any(): | |
| return None | |
| ys = np.where(rows)[0] | |
| xs = np.where(cols)[0] | |
| y0, y1 = int(ys[0]), int(ys[-1]) | |
| x0, x1 = int(xs[0]), int(xs[-1]) | |
| return x0, y0, x1, y1 | |
| def apply_bbox_inpaint(frame_rgb: np.ndarray, mask: np.ndarray, box_color: tuple) -> np.ndarray: | |
| """Return RGB frame with the mask bounding box filled with box_color.""" | |
| result = frame_rgb.copy() | |
| bbox = mask_to_bbox(mask) | |
| if bbox is not None: | |
| x0, y0, x1, y1 = bbox | |
| result[y0:y1 + 1, x0:x1 + 1] = box_color | |
| return result | |
| def save_frames(frames_rgba: list, out_dir: str) -> None: | |
| from tqdm import tqdm | |
| os.makedirs(out_dir, exist_ok=True) | |
| for idx, rgba in enumerate(tqdm(frames_rgba, desc="saving PNG frames", unit="fr")): | |
| path = os.path.join(out_dir, f"frame_{idx:04d}.png") | |
| Image.fromarray(rgba, mode="RGBA").save(path) | |
| def save_video( | |
| frames_rgba: list, | |
| out_path: str, | |
| fps: float, | |
| bg_color: tuple = (0, 0, 0), | |
| ) -> None: | |
| """Save frames as an MP4 composited over a solid background colour.""" | |
| if not frames_rgba: | |
| return | |
| h, w = frames_rgba[0].shape[:2] | |
| tmp = out_path + ".tmp.mp4" | |
| fourcc = cv2.VideoWriter_fourcc(*"mp4v") | |
| writer = cv2.VideoWriter(tmp, fourcc, fps, (w, h)) | |
| bg = np.array(bg_color, dtype=np.float32) | |
| for rgba in frames_rgba: | |
| alpha = rgba[..., 3:4].astype(np.float32) / 255.0 | |
| fg = rgba[..., :3].astype(np.float32) | |
| composite = (fg * alpha + bg * (1.0 - alpha)).clip(0, 255).astype(np.uint8) | |
| writer.write(cv2.cvtColor(composite, cv2.COLOR_RGB2BGR)) | |
| writer.release() | |
| # Re-encode for broad compatibility | |
| ret = subprocess.run( | |
| ["ffmpeg", "-y", "-i", tmp, out_path], | |
| capture_output=True, | |
| ) | |
| os.remove(tmp) | |
| if ret.returncode != 0: | |
| print(f" ffmpeg re-encode failed; raw file left at {tmp}") | |
| def save_frames_rgb(frames_rgb: list, out_dir: str) -> None: | |
| from tqdm import tqdm | |
| os.makedirs(out_dir, exist_ok=True) | |
| for idx, rgb in enumerate(tqdm(frames_rgb, desc="saving PNG frames", unit="fr")): | |
| path = os.path.join(out_dir, f"frame_{idx:04d}.png") | |
| Image.fromarray(rgb, mode="RGB").save(path) | |
| def save_video_rgb(frames_rgb: list, out_path: str, fps: float) -> None: | |
| from tqdm import tqdm | |
| if not frames_rgb: | |
| return | |
| h, w = frames_rgb[0].shape[:2] | |
| tmp = out_path + ".tmp.mp4" | |
| fourcc = cv2.VideoWriter_fourcc(*"mp4v") | |
| writer = cv2.VideoWriter(tmp, fourcc, fps, (w, h)) | |
| for rgb in tqdm(frames_rgb, desc="encoding video", unit="fr"): | |
| writer.write(cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR)) | |
| writer.release() | |
| print(f" re-encoding with ffmpeg …") | |
| ret = subprocess.run(["ffmpeg", "-y", "-i", tmp, out_path], capture_output=True) | |
| os.remove(tmp) | |
| if ret.returncode != 0: | |
| print(f" ffmpeg re-encode failed; raw file left at {tmp}") | |
| # --------------------------------------------------------------------------- | |
| # Core processing | |
| # --------------------------------------------------------------------------- | |
| def process_video( | |
| predictor, | |
| video_path: str, | |
| prompt: str, | |
| output_dir: str, | |
| fps: float, | |
| save_png_frames: bool, | |
| save_mp4: bool, | |
| bg_color: tuple, | |
| mode: str = "extract", | |
| box_color: tuple = (0, 0, 0), | |
| ) -> None: | |
| frames = load_video_frames(video_path) | |
| if not frames: | |
| print(f" Warning: no frames loaded from {video_path}") | |
| return | |
| # --- start session --- | |
| response = predictor.handle_request( | |
| request=dict(type="start_session", resource_path=video_path) | |
| ) | |
| session_id = response["session_id"] | |
| try: | |
| predictor.handle_request( | |
| request=dict(type="reset_session", session_id=session_id) | |
| ) | |
| # Add text prompt and inspect first-frame detections | |
| response = predictor.handle_request( | |
| request=dict( | |
| type="add_prompt", | |
| session_id=session_id, | |
| frame_index=0, | |
| text=prompt, | |
| ) | |
| ) | |
| out0 = response["outputs"] | |
| obj_ids_0 = out0["out_obj_ids"].tolist() | |
| masks_0 = out0["out_binary_masks"] | |
| if not obj_ids_0: | |
| print(f" No detections for prompt '{prompt}'") | |
| return | |
| # Keep only the object with the largest mask area in frame 0 | |
| areas = [mask_to_numpy(masks_0[i]).sum() for i in range(len(obj_ids_0))] | |
| best_i = int(np.argmax(areas)) | |
| keep_id = int(obj_ids_0[best_i]) | |
| print(f" Detected {len(obj_ids_0)} object(s); keeping obj_id={keep_id} (area={areas[best_i]})") | |
| for oid in obj_ids_0: | |
| if int(oid) != keep_id: | |
| predictor.handle_request( | |
| request=dict(type="remove_object", session_id=session_id, obj_id=int(oid)) | |
| ) | |
| # --- propagate --- | |
| # Drain the stream with minimal work per iteration: numpy ops inside the loop | |
| # would suspend rank 0 long enough for worker ranks to advance to the next | |
| # frame's distributed op and deadlock. Collect raw outputs first, then | |
| # process them after the generator (and its barrier()) has fully completed. | |
| raw_outputs: dict[int, object] = {} | |
| for resp in predictor.handle_stream_request( | |
| request=dict(type="propagate_in_video", session_id=session_id) | |
| ): | |
| frame_idx = resp["frame_index"] | |
| if frame_idx < len(frames): | |
| raw_outputs[frame_idx] = resp["outputs"] | |
| obj_frames: list[np.ndarray | None] = [None] * len(frames) | |
| for frame_idx, out in raw_outputs.items(): | |
| frame_rgb = frames[frame_idx] | |
| h, w = frame_rgb.shape[:2] | |
| mask = np.zeros((h, w), dtype=bool) | |
| for i, oid in enumerate(out["out_obj_ids"].tolist()): | |
| if int(oid) == keep_id: | |
| mask = mask_to_numpy(out["out_binary_masks"][i]) | |
| break | |
| if mode == "extract": | |
| obj_frames[frame_idx] = apply_mask_rgba(frame_rgb, mask) | |
| else: | |
| obj_frames[frame_idx] = apply_bbox_inpaint(frame_rgb, mask, box_color) | |
| finally: | |
| predictor.handle_request( | |
| request=dict(type="close_session", session_id=session_id) | |
| ) | |
| # Fill frames where the object was not tracked | |
| for i, frame in enumerate(obj_frames): | |
| if frame is None: | |
| h, w = frames[i].shape[:2] | |
| if mode == "extract": | |
| obj_frames[i] = apply_mask_rgba(frames[i], np.zeros((h, w), dtype=bool)) | |
| else: | |
| obj_frames[i] = frames[i].copy() | |
| if mode == "extract": | |
| if save_png_frames: | |
| save_frames(obj_frames, os.path.join(output_dir, "frames")) | |
| print(f" Saved {len(obj_frames)} PNG frames → {output_dir}/frames/") | |
| if save_mp4: | |
| vp = os.path.join(output_dir, "subject.mp4") | |
| save_video(obj_frames, vp, fps, bg_color) | |
| print(f" Saved video → {vp}") | |
| else: | |
| if save_png_frames: | |
| save_frames_rgb(obj_frames, os.path.join(output_dir, "frames")) | |
| print(f" Saved {len(obj_frames)} PNG frames → {output_dir}/frames/") | |
| if save_mp4: | |
| vp = os.path.join(output_dir, "inpaint.mp4") | |
| save_video_rgb(obj_frames, vp, fps) | |
| print(f" Saved video → {vp}") | |
| # --------------------------------------------------------------------------- | |
| # Entry point | |
| # --------------------------------------------------------------------------- | |
| def main() -> None: | |
| parser = argparse.ArgumentParser( | |
| description="Extract subjects from videos using SAM3 text-prompted segmentation" | |
| ) | |
| parser.add_argument( | |
| "--video_dir", required=True, | |
| help="Folder containing .mp4 files or JPEG frame sub-folders", | |
| ) | |
| parser.add_argument( | |
| "--csv", required=True, | |
| help="CSV file with columns 'video' and 'prompt'", | |
| ) | |
| parser.add_argument( | |
| "--output_dir", required=True, | |
| help="Root directory for extracted outputs", | |
| ) | |
| parser.add_argument( | |
| "--checkpoint", default=None, | |
| help="Path to SAM 3 checkpoint .pt file (default: sam3/sam3/sam3.pt, or auto-download from HuggingFace)", | |
| ) | |
| parser.add_argument( | |
| "--no_frames", action="store_true", | |
| help="Do not save per-frame RGBA PNG files", | |
| ) | |
| parser.add_argument( | |
| "--save_video", action="store_true", | |
| help="Also save an output MP4 video (off by default)", | |
| ) | |
| parser.add_argument( | |
| "--fps", type=float, default=None, | |
| help="FPS for output video (default: auto-detected from source, fallback 30)", | |
| ) | |
| parser.add_argument( | |
| "--bg_color", type=str, default="0,0,0", | |
| help="Background colour for the output video as R,G,B (default: 0,0,0 = black)", | |
| ) | |
| parser.add_argument( | |
| "--mode", choices=["extract", "inpaint"], default="extract", | |
| help="'extract': output subject with transparent background (default); " | |
| "'inpaint': output original video with a filled bounding box over the subject", | |
| ) | |
| parser.add_argument( | |
| "--box_color", type=str, default="0,0,0", | |
| help="Fill colour for the inpaint bounding box as R,G,B (default: 0,0,0 = black)", | |
| ) | |
| args = parser.parse_args() | |
| bg_color = tuple(int(x) for x in args.bg_color.split(",")) | |
| if len(bg_color) != 3: | |
| parser.error("--bg_color must be three comma-separated integers, e.g. 0,0,0") | |
| box_color = tuple(int(x) for x in args.box_color.split(",")) | |
| if len(box_color) != 3: | |
| parser.error("--box_color must be three comma-separated integers, e.g. 128,128,128") | |
| # Read CSV — expects columns: clip_id, subject | |
| # video is resolved as <video_dir>/<clip_id>.mp4 | |
| entries: list[tuple[str, str, str]] = [] # (clip_id, video_path, prompt) | |
| with open(args.csv, newline="", encoding="utf-8") as f: | |
| reader = csv.DictReader(f) | |
| for row in reader: | |
| clip_id = row["clip_id"].strip() | |
| prompt = row['subject'].strip() | |
| video_path = os.path.join(args.video_dir, clip_id + ".mp4") | |
| entries.append((clip_id, video_path, prompt)) | |
| if not entries: | |
| print("CSV is empty — nothing to do.") | |
| sys.exit(0) | |
| print(f"Loaded {len(entries)} entries from {args.csv}") | |
| # Build model | |
| print("Loading SAM 3 model …") | |
| import torch | |
| from sam3.model_builder import build_sam3_video_predictor | |
| _checkpoint = args.checkpoint | |
| if _checkpoint is None: | |
| _script_dir = os.path.dirname(os.path.abspath(__file__)) | |
| _local = os.path.join(_script_dir, "sam3", "sam3", "sam3.pt") | |
| _checkpoint = _local if os.path.exists(_local) else None | |
| gpus_to_use = list(range(torch.cuda.device_count())) or [torch.cuda.current_device()] | |
| predictor = build_sam3_video_predictor(checkpoint_path=_checkpoint, gpus_to_use=gpus_to_use) | |
| print("Model ready.\n") | |
| # Process | |
| for clip_id, video_path, prompt in entries: | |
| if not os.path.exists(video_path): | |
| print(f"[SKIP] Not found: {video_path}") | |
| continue | |
| fps = args.fps if args.fps is not None else get_video_fps(video_path) | |
| out_dir = os.path.join(args.output_dir, clip_id) | |
| os.makedirs(out_dir, exist_ok=True) | |
| print(f"[{clip_id}] prompt='{prompt}' fps={fps:.1f} → {out_dir}") | |
| process_video( | |
| predictor=predictor, | |
| video_path=video_path, | |
| prompt=prompt, | |
| output_dir=out_dir, | |
| fps=fps, | |
| save_png_frames=not args.no_frames, | |
| save_mp4=args.save_video, | |
| bg_color=bg_color, | |
| mode=args.mode, | |
| box_color=box_color, | |
| ) | |
| print("\nAll done.") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 15.3 kB
- Xet hash:
- 53973e1a0cb7aae9033db936f5ffcce0972d6de0223bae8f88d78afebcb24093
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.