From cc4caff2861c1ad2da873e0d153bde14e65ed989 Mon Sep 17 00:00:00 2001 From: Peter Schroedl Date: Sat, 30 Nov 2024 18:18:52 +0100 Subject: [PATCH] wip fixup automasking, show coord point --- nodes.py | 74 +++++++++++++++++++++++++++++++++++++++++--------------- 1 file changed, 55 insertions(+), 19 deletions(-) diff --git a/nodes.py b/nodes.py index 5cc19f9..85a1f7c 100644 --- a/nodes.py +++ b/nodes.py @@ -8,6 +8,8 @@ import ast import sys +import cv2 + # Add the directory containing 'sam2_realtime' to sys.path current_directory = os.path.dirname(os.path.abspath(__file__)) sam2_realtime_path = os.path.join(current_directory) # Adjust the relative path @@ -83,7 +85,7 @@ class DownloadAndLoadSAM2RealtimeModel: config_dir = os.path.join(script_directory, "sam2_configs") # Code ripped out of sam2.build_sam.build_sam2_camera_predictor to appease Hydra - model_cfg = "sam2_hiera_t.yaml" #TODO: remove hardcoded config and path + model_cfg = "sam2_hiera_t.yaml" #TODO(pschroedl): remove hardcoded config and path with initialize_config_dir(config_dir=config_dir, version_base=None): cfg = compose(config_name=model_cfg) @@ -148,6 +150,8 @@ class Sam2RealtimeSegmentation: # "bboxes": ("BBOX", ), # "individual_objects": ("BOOLEAN", {"default": False}), # "mask": ("MASK", ), + "threshold": ("FLOAT", {"forceInput": True}), + "show_point": ("BOOLEAN", {"default": False}), }, } @@ -160,6 +164,40 @@ class Sam2RealtimeSegmentation: self.predictor = None self.if_init = False + + def _process_mask(self, mask: np.ndarray, frame_shape: tuple) -> np.ndarray: + if mask.shape[0] == 0: + logging.warning("Empty mask received") + return np.zeros((frame_shape[0], frame_shape[1]), dtype="uint8") + + colors = [ + [255, 0, 255], # Purple + [0, 255, 255], # Yellow + [255, 255, 0], # Cyan + [0, 255, 0], # Green + [255, 0, 0], # Blue + ] + + combined_colored_mask = np.zeros((frame_shape[0], frame_shape[1], 4), dtype="uint8") + + for i in range(mask.shape[0]): + current_mask = (mask[i, 0] > 0).cpu().numpy().astype("uint8") * 255 + if current_mask.shape[:2] != frame_shape[:2]: + current_mask = cv2.resize(current_mask, (frame_shape[1], frame_shape[0])) + + # Create BGRA mask with transparency + colored_mask = np.zeros((frame_shape[0], frame_shape[1], 4), dtype="uint8") + color = colors[i % len(colors)] + colored_mask[current_mask > 0] = color + [128] # Add alpha value of 128 + + # Alpha blend with existing masks + alpha = colored_mask[:, :, 3:4] / 255.0 + combined_colored_mask = (1 - alpha) * combined_colored_mask + alpha * colored_mask + + # Convert back to BGR for display + combined_colored_mask = combined_colored_mask[:, :, :3].astype("uint8") + return combined_colored_mask + def segment_images( self, images, @@ -171,6 +209,8 @@ class Sam2RealtimeSegmentation: # bboxes=None, # individual_objects=False, # mask=None, + threshold=0.5, + show_point=False, ): model = sam2_model["model"] device = sam2_model["device"] @@ -180,21 +220,17 @@ class Sam2RealtimeSegmentation: processed_frames = [] mask_list = [] - # The `model` variable is now ready and equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor + # The `model` is equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor if self.predictor is None: self.predictor = model def process_frame(frame, frame_idx): with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16): - frame = frame.to(device).float() # Keep everything in torch + frame = frame.to(device).float() if not self.if_init: self.predictor.load_first_frame(frame) self.if_init = True - # obj_id = 1 - # point = [256, 256] - # points = [point] - # labels = [1] coordinates_positive_list = ast.literal_eval(coordinates_positive) point_labels_list = ast.literal_eval(point_labels) @@ -208,34 +244,34 @@ class Sam2RealtimeSegmentation: points=[point_tuple], labels=[point_labels_list[idx]] ) - - # _, _, _ = self.predictor.add_new_prompt(frame_idx, obj_id, points=points, labels=labels) else: out_obj_ids, out_mask_logits = self.predictor.track(frame) if out_mask_logits.shape[0] > 0: - # Ensure out_mask_logits is on the same device - out_mask_logits = out_mask_logits.to(device) - mask = (out_mask_logits[0, 0] > 0.5).byte() + mask = (out_mask_logits[0, 0] > threshold).byte() mask = torch.nn.functional.interpolate( mask.unsqueeze(0).unsqueeze(0).float(), size=(frame.shape[0], frame.shape[1]), mode='nearest' - ).squeeze(0).squeeze(0).byte().to(device) # Move the interpolated mask to the correct device + ).squeeze(0).squeeze(0).byte() # Move the interpolated mask to the correct device else: mask = torch.ones((frame.shape[0], frame.shape[1]), device=device, dtype=torch.uint8) - # Ensure frame is on the same device - frame = frame.to(device) + automask_colored = self._process_mask(mask,frame.shape) - mask_colored = torch.stack([mask] * 3, dim=2).to(device) # Create 3-channel mask and move to device + # Draw points on the mask + if show_point: + for point in coordinates_positive: + cv2.circle(automask_colored, tuple(point), radius=5, color=(0, 0, 255), thickness=-1) - overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3).to(device) - processed_frames.append(overlayed_frame) + automasked_frame = torch.add(frame * 0.7, automask_colored * 0.3) + processed_frames.append(automasked_frame) - constructed_mask = torch.add(frame * 0.1, mask_colored * 0.9).to(device) + # TODO: This "mask" should be 1 channel to be returned as MASK type + constructed_mask = torch.add(frame * 0.1, mask * 0.9) mask_list.append(constructed_mask) + for frame_idx, img in enumerate(images): process_frame(img, frame_idx)