Spaces:
Paused
Paused
File size: 14,062 Bytes
09462dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 | import os
from random import shuffle
import cv2
import numpy as np
from decord import VideoReader
# Minimum mask ratio threshold (percentage of frame). Override via env var
# SCAIL_MIN_MASK_RATIO for small-subject scenes (e.g. paper figures where the
# subject occupies <1% of the frame) without editing this file.
MIN_MASK_RATIO = float(os.environ.get('SCAIL_MIN_MASK_RATIO', '1.0'))
# Default cap on number of targets when caller does not override
DEFAULT_MAX_TARGETS = 4
# Deterministic BGR palette used when callers want stable colors across runs.
DEFAULT_PALETTE_BGR = [
(255, 0, 0), # Blue
(0, 0, 255), # Red
(0, 255, 0), # Green
(255, 0, 255), # Magenta
(255, 255, 0), # Cyan
(0, 255, 255), # Yellow
]
def remove_small_tracks_from_predictor(predictor, invalid_track_ids):
"""Remove invalid track IDs from predictor's internal tracker state."""
if not invalid_track_ids:
return
metadata = predictor.inference_state.get("tracker_metadata", {})
if not metadata:
return
obj_ids = metadata.get("obj_ids_all_gpu", np.array([]))
if len(obj_ids) == 0:
return
keep_mask = np.array([int(oid) not in invalid_track_ids for oid in obj_ids])
metadata["obj_ids_all_gpu"] = obj_ids[keep_mask]
arrays_to_filter = [
"obj_id_to_score", "obj_id_to_cls", "obj_id_to_tracker_score"
]
for key in arrays_to_filter:
if key in metadata and isinstance(metadata[key], dict):
metadata[key] = {k: v for k, v in metadata[key].items() if int(k) not in invalid_track_ids}
tracker_states = predictor.inference_state.get("tracker_inference_states", [])
if tracker_states:
for state in tracker_states:
if hasattr(state, 'obj_ids') and state.obj_ids is not None:
state_keep = np.array([int(oid) not in invalid_track_ids for oid in state.obj_ids])
state.obj_ids = state.obj_ids[state_keep]
print(f"Removed track IDs {invalid_track_ids} from tracker state")
def visualize_and_save_mask(results, width, height, predictor, new_indices, full_length,
max_targets=DEFAULT_MAX_TARGETS, shuffle_colors=True,
direct_return=False):
"""Run through SAM3 streaming results and gather per-track binary masks.
Returns (valid_track_ids_ordered, mask_arrays, track_colors) ordered by descending
mask area in the first frame; or None if no valid track is detected.
"""
colors = list(DEFAULT_PALETTE_BGR)
if shuffle_colors:
shuffle(colors)
frame_idx = 0
valid_track_ids = None
total_pixels = height * width
valid_track_ids_ordered = []
mask_arrays = {}
track_colors = {}
_color_counter = 0
for result_idx, result in enumerate(results):
index_result = new_indices[result_idx]
if result.masks is not None:
masks = result.masks.data.cpu().numpy() # (N, H, W)
track_ids = result.boxes.id.cpu().numpy() if result.boxes.id is not None else np.arange(len(masks))
if frame_idx == 0:
valid_track_ids = set()
invalid_track_ids = set()
candidates = []
for i, (mask, track_id) in enumerate(zip(masks, track_ids)):
if mask.shape[:2] != (height, width):
mask_resized = cv2.resize(mask.astype(np.float32), (width, height))
else:
mask_resized = mask
mask_bool = mask_resized > 0.5
mask_ratio = np.sum(mask_bool) / total_pixels * 100
if mask_ratio >= MIN_MASK_RATIO:
candidates.append((int(track_id), mask_ratio))
else:
invalid_track_ids.add(int(track_id))
candidates.sort(key=lambda x: x[1], reverse=True)
if len(candidates) == 0 and direct_return:
print(f" No valid candidates (all < MIN_MASK_RATIO={MIN_MASK_RATIO}%) in first frame, return")
return
if len(candidates) > max_targets:
if direct_return:
print(f" Found {len(candidates)} candidates, return")
return
print(f" Found {len(candidates)} candidates, limiting to top {max_targets}")
kept_candidates = candidates[:max_targets]
dropped_candidates = candidates[max_targets:]
for track_id, _ in kept_candidates:
valid_track_ids.add(track_id)
for track_id, _ in dropped_candidates:
invalid_track_ids.add(track_id)
else:
kept_candidates = candidates
for track_id, _ in candidates:
valid_track_ids.add(track_id)
if kept_candidates and direct_return:
max_ratio = kept_candidates[0][1]
if max_ratio < 1.5 or max_ratio > 50:
print(f" Max mask ratio {max_ratio:.2f}% out of valid range [1.5, 50], return")
return
if len(kept_candidates) >= 2 and direct_return:
max_ratio = kept_candidates[0][1]
min_ratio = kept_candidates[-1][1]
if min_ratio < max_ratio / 3:
print(f" Smallest person ({min_ratio:.2f}%) < 1/3 of largest ({max_ratio:.2f}%), return")
return
valid_track_ids_ordered = [tid for tid, _ in kept_candidates]
mask_arrays = {tid: np.zeros((full_length, height, width), dtype=bool)
for tid in valid_track_ids_ordered}
if invalid_track_ids:
remove_small_tracks_from_predictor(predictor, invalid_track_ids)
for i, (mask, track_id) in enumerate(zip(masks, track_ids)):
if valid_track_ids is not None and int(track_id) not in valid_track_ids:
continue
tid = int(track_id)
if tid not in track_colors:
track_colors[tid] = colors[_color_counter % len(colors)]
_color_counter += 1
if mask.shape[:2] != (height, width):
mask = cv2.resize(mask.astype(np.float32), (width, height))
mask_bool = mask > 0.5
if tid in mask_arrays:
mask_arrays[tid][index_result] = mask_bool
frame_idx += 1
if not valid_track_ids_ordered:
return None
return valid_track_ids_ordered, mask_arrays, track_colors
def _centroid_x(mask_2d):
"""X-coordinate of the centroid of a 2D bool mask. Returns +inf if mask is empty."""
cols = np.where(mask_2d.any(axis=0))[0]
if len(cols) == 0:
return float('inf')
rows = np.where(mask_2d.any(axis=1))[0]
# use bounding-box center (cheap and stable)
return 0.5 * (cols[0] + cols[-1])
def _reorder_and_color(valid_track_ids_ordered, mask_arrays, sort_by, fixed_colors):
"""Apply left-to-right sort and deterministic color assignment.
Returns (masks, colors) where masks is a list of (T, H, W) bool ndarray and
colors is a list of BGR tuples, both in the chosen ordering.
"""
if sort_by == 'x':
ordered = sorted(valid_track_ids_ordered,
key=lambda tid: _centroid_x(mask_arrays[tid][0]))
elif sort_by == 'area':
ordered = list(valid_track_ids_ordered)
else:
raise ValueError(f"unknown sort_by: {sort_by}")
n = len(ordered)
if fixed_colors is not None:
if len(fixed_colors) < n:
raise ValueError(f"fixed_colors has {len(fixed_colors)} entries but {n} tracks")
colors = [tuple(c) for c in fixed_colors[:n]]
else:
colors = [DEFAULT_PALETTE_BGR[i % len(DEFAULT_PALETTE_BGR)] for i in range(n)]
masks = [mask_arrays[tid] for tid in ordered]
return masks, colors
def get_mask_from_video(video_path, predictor, max_targets=DEFAULT_MAX_TARGETS,
sort_by='area', fixed_colors=None,
text=("human", "character")):
"""Run SAM3 tracking on a video file and return per-person binary masks and colors.
Args:
video_path: path to input video (str or Path).
predictor: SAM3VideoSemanticPredictor instance (state will be reset).
max_targets: cap on the number of tracked persons (kept by descending area).
sort_by: 'area' (default, descending area) or 'x' (left-to-right by first-frame
centroid x).
fixed_colors: optional list of BGR tuples assigned to ordered tracks instead of
the default palette. Must have at least len(tracks) entries.
Returns:
masks: list of (T, H, W) bool ndarray, one per tracked person.
colors: list of BGR color tuples corresponding to each person.
Both lists are empty if no valid persons are detected.
"""
video_path = str(video_path)
predictor.inference_state = {}
if hasattr(predictor, 'dataset'):
predictor.dataset = None
vr = VideoReader(video_path)
full_length = len(vr)
height, width = vr[0].asnumpy().shape[:2]
del vr
results = predictor(source=video_path, text=list(text), stream=True)
ret = visualize_and_save_mask(
results, width, height, predictor,
new_indices=np.arange(full_length), full_length=full_length,
max_targets=max_targets, shuffle_colors=fixed_colors is None,
direct_return=False,
)
if ret is None:
return [], []
valid_track_ids_ordered, mask_arrays, _ = ret
return _reorder_and_color(valid_track_ids_ordered, mask_arrays, sort_by, fixed_colors)
def get_mask_from_image_via_video(image_path, video_predictor, max_targets=DEFAULT_MAX_TARGETS,
sort_by='x', fixed_colors=None,
text=("human", "character"), n_repeat=4, fps=8):
"""Detect persons in a still image by wrapping it as a tiny mp4 and routing through
SAM3VideoSemanticPredictor. Workaround for image-mode SAM3 missing small / distant
subjects that the video pipeline picks up reliably.
Returns (masks, colors) with each mask shaped (1, H, W) bool — only the first frame
of the synthetic clip is kept.
"""
import tempfile
from NLFPoseExtract.v2_helper import imread_bgr
image_path = str(image_path)
img = imread_bgr(image_path)
H, W = img.shape[:2]
tmp_fd, tmp_path = tempfile.mkstemp(suffix='.mp4')
os.close(tmp_fd)
try:
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
vw = cv2.VideoWriter(tmp_path, fourcc, float(fps), (W, H))
if not vw.isOpened():
raise RuntimeError(f"cv2.VideoWriter failed to open {tmp_path}")
for _ in range(n_repeat):
vw.write(img)
vw.release()
masks, colors = get_mask_from_video(
tmp_path, video_predictor,
max_targets=max_targets, sort_by=sort_by,
fixed_colors=fixed_colors, text=text,
)
finally:
try:
os.unlink(tmp_path)
except OSError:
pass
masks = [m[:1] for m in masks]
return masks, colors
def get_mask_from_image(image_path, predictor, max_targets=DEFAULT_MAX_TARGETS,
sort_by='x', fixed_colors=None,
text=("human", "character")):
"""Run SAM3SemanticPredictor (image variant) on a single image.
Args:
image_path: path to input image (str or Path).
predictor: SAM3SemanticPredictor instance.
max_targets: cap on number of persons (kept by descending area).
sort_by: 'x' (default, left-to-right) or 'area'.
fixed_colors: optional list of BGR tuples assigned in order; otherwise the
deterministic palette is used.
Returns:
masks: list of (1, H, W) bool ndarray, one per detected person.
colors: list of BGR color tuples corresponding to each person.
"""
image_path = str(image_path)
results = predictor(source=image_path, text=list(text))
if not results:
return [], []
result = results[0]
if result.masks is None or len(result.masks) == 0:
return [], []
masks_NHW = result.masks.data.cpu().numpy() # (N, H, W)
masks_NHW = masks_NHW > 0.5
return _filter_image_masks(masks_NHW, max_targets, sort_by, fixed_colors)
def _filter_image_masks(masks_NHW, max_targets, sort_by, fixed_colors):
"""Apply MIN_MASK_RATIO + max_targets filter to per-image SAM masks, then order
and color them. Returns (masks_list, colors) where each mask is (1, H, W) bool.
"""
N, H, W = masks_NHW.shape
total_pixels = H * W
candidates = [] # (idx, mask_ratio)
for i in range(N):
ratio = float(np.sum(masks_NHW[i])) / total_pixels * 100
if ratio >= MIN_MASK_RATIO:
candidates.append((i, ratio))
candidates.sort(key=lambda x: x[1], reverse=True)
candidates = candidates[:max_targets]
if not candidates:
return [], []
kept = [masks_NHW[idx] for idx, _ in candidates] # list of (H, W) bool
if sort_by == 'x':
order = sorted(range(len(kept)), key=lambda i: _centroid_x(kept[i]))
kept = [kept[i] for i in order]
elif sort_by != 'area':
raise ValueError(f"unknown sort_by: {sort_by}")
n = len(kept)
if fixed_colors is not None:
if len(fixed_colors) < n:
raise ValueError(f"fixed_colors has {len(fixed_colors)} entries but {n} masks")
colors = [tuple(c) for c in fixed_colors[:n]]
else:
colors = [DEFAULT_PALETTE_BGR[i % len(DEFAULT_PALETTE_BGR)] for i in range(n)]
masks_out = [m[None] for m in kept] # add T=1 axis
return masks_out, colors
|