fishxinyu's picture
download
raw
15.3 kB
#!/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.