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