From 06061eb074f14dd2aeffb06f9a9868dc87e6c55b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 1 Aug 2024 01:28:30 +0300 Subject: [PATCH] segmenged image output from automask node --- nodes.py | 73 +++++++++++++++++++++++++++++++++++++++++--------------- 1 file changed, 54 insertions(+), 19 deletions(-) diff --git a/nodes.py b/nodes.py index 9c15287..1cf763e 100644 --- a/nodes.py +++ b/nodes.py @@ -2,6 +2,7 @@ import torch import os import numpy as np import json +import random from contextlib import nullcontext @@ -203,12 +204,13 @@ class Sam2Segmentation: combined_coords = np.concatenate((positive_point_coords, negative_point_coords), axis=0) combined_labels = np.concatenate((positive_point_labels, negative_point_labels), axis=0) - autocast_condition = not mm.is_device_mps(device) mask_list = [] try: model.to(device) except: model.model.to(device) + + autocast_condition = not mm.is_device_mps(device) with torch.autocast(mm.get_autocast_device(model.device), dtype=dtype) if autocast_condition else nullcontext(): if image.shape[0] == 1: model.set_image(image_np) @@ -275,7 +277,7 @@ class Sam2AutoSegmentation: "image": ("IMAGE", ), "points_per_side": ("INT", {"default": 32}), "points_per_batch": ("INT", {"default": 64}), - "pred_iou_thresh": ("FLOAT", {"default": 0.8}), + "pred_iou_thresh": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.01}), "stability_score_thresh": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.01}), "stability_score_offset": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), "mask_threshold": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}), @@ -292,8 +294,8 @@ class Sam2AutoSegmentation: } - RETURN_TYPES = ("MASK", "BBOX",) - RETURN_NAMES =("mask", "bbox" ,) + RETURN_TYPES = ("MASK", "IMAGE", "BBOX",) + RETURN_NAMES =("mask", "segmented_image", "bbox" ,) FUNCTION = "segment" CATEGORY = "SAM2" @@ -305,9 +307,7 @@ class Sam2AutoSegmentation: device = sam2_model["device"] dtype = sam2_model["dtype"] segmentor = sam2_model["segmentor"] - B, H, W, C = image.shape - image_np = (image[0].contiguous() * 255).byte().numpy() - + if segmentor != 'automaskgenerator': raise ValueError("Loaded model is not SAM2AutomaticMaskGenerator") @@ -325,26 +325,61 @@ class Sam2AutoSegmentation: model.use_m2m=use_m2m model.mask_threshold=mask_threshold - autocast_condition = not mm.is_device_mps(device) - - model.predictor.model.to(device) + + B, H, W, C = image.shape + image_np = (image.contiguous() * 255).byte().numpy() + + out_list = [] + segment_out_list = [] mask_list=[] + + pbar = ProgressBar(B) + autocast_condition = not mm.is_device_mps(device) with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): - result_dict = model.generate(image_np) - print(result_dict[0].keys()) - mask_list = [item['segmentation'] for item in result_dict] - bbox_list = [item['bbox'] for item in result_dict] + for img_np in image_np: + result_dict = model.generate(img_np) + mask_list = [item['segmentation'] for item in result_dict] + bbox_list = [item['bbox'] for item in result_dict] + + # Generate random colors for each mask + num_masks = len(mask_list) + colors = [tuple(random.choices(range(256), k=3)) for _ in range(num_masks)] + + # Create a blank image to overlay masks + overlay_image = np.zeros((H, W, 3), dtype=np.uint8) + + # Create a combined mask initialized to zeros + combined_mask = np.zeros((H, W), dtype=np.uint8) + + # Iterate through masks and color them + for mask, color in zip(mask_list, colors): + + # Combine masks using logical OR + combined_mask = np.logical_or(combined_mask, mask).astype(np.uint8) + + # Convert mask to numpy array + mask_np = mask.astype(np.uint8) + + # Color the mask + colored_mask = np.zeros_like(overlay_image) + for i in range(3): # Apply color channel-wise + colored_mask[:, :, i] = mask_np * color[i] + + # Blend the colored mask with the overlay image + overlay_image = np.where(colored_mask > 0, colored_mask, overlay_image) + out_list.append(torch.from_numpy(combined_mask)) + segment_out_list.append(overlay_image) + pbar.update(1) + + stacked_array = np.stack(segment_out_list, axis=0) + segment_image_tensor = torch.from_numpy(stacked_array).float() / 255 if not keep_model_loaded: model.predictor.model.to(offload_device) - out_list = [] - for mask in mask_list: - mask_tensor = torch.from_numpy(mask) - out_list.append(mask_tensor) mask_tensor = torch.stack(out_list, dim=0) - return (mask_tensor.cpu().float(), bbox_list) + return (mask_tensor.cpu().float(), segment_image_tensor.cpu().float(), bbox_list) NODE_CLASS_MAPPINGS = { "DownloadAndLoadSAM2Model": DownloadAndLoadSAM2Model,