diff --git a/examples/florence_segment_2.json b/examples/florence_segment_2.json index e7c28b6..93b5aad 100644 --- a/examples/florence_segment_2.json +++ b/examples/florence_segment_2.json @@ -1,6 +1,6 @@ { - "last_node_id": 101, - "last_link_id": 235, + "last_node_id": 102, + "last_link_id": 239, "nodes": [ { "id": 83, @@ -41,31 +41,6 @@ "image" ] }, - { - "id": 90, - "type": "PreviewImage", - "pos": [ - 409, - -775 - ], - "size": { - "0": 568.406494140625, - "1": 384.9489440917969 - }, - "flags": {}, - "order": 6, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 200 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } - }, { "id": 66, "type": "DownloadAndLoadSAM2Model", @@ -85,7 +60,7 @@ "name": "sam2_model", "type": "SAM2MODEL", "links": [ - 231 + 236 ], "shape": 3, "slot_index": 0 @@ -113,7 +88,7 @@ "1": 541.2733154296875 }, "flags": {}, - "order": 10, + "order": 9, "mode": 0, "inputs": [ { @@ -124,7 +99,7 @@ { "name": "mask", "type": "MASK", - "link": 233, + "link": 238, "slot_index": 1 } ], @@ -145,40 +120,6 @@ false ] }, - { - "id": 88, - "type": "DownloadAndLoadFlorence2Model", - "pos": [ - -512, - -707 - ], - "size": { - "0": 315, - "1": 106 - }, - "flags": {}, - "order": 2, - "mode": 0, - "outputs": [ - { - "name": "florence2_model", - "type": "FL2MODEL", - "links": [ - 197 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "DownloadAndLoadFlorence2Model" - }, - "widgets_values": [ - "microsoft/Florence-2-base", - "fp16", - "sdpa" - ] - }, { "id": 72, "type": "ImageResizeKJ", @@ -229,7 +170,7 @@ 192, 210, 226, - 232 + 237 ], "shape": 3, "slot_index": 0 @@ -285,6 +226,77 @@ "Node name for S&R": "PreviewImage" } }, + { + "id": 90, + "type": "PreviewImage", + "pos": [ + 422, + -800 + ], + "size": { + "0": 568.406494140625, + "1": 384.9489440917969 + }, + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 200 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 93, + "type": "Florence2toCoordinates", + "pos": [ + 399, + -314 + ], + "size": { + "0": 210, + "1": 78 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "data", + "type": "JSON", + "link": 204 + } + ], + "outputs": [ + { + "name": "coordinates", + "type": "STRING", + "links": [], + "shape": 3, + "slot_index": 0 + }, + { + "name": "bboxes", + "type": "BBOX", + "links": [ + 239 + ], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "Florence2toCoordinates" + }, + "widgets_values": [ + "" + ] + }, { "id": 87, "type": "Florence2Run", @@ -340,8 +352,7 @@ "name": "data", "type": "JSON", "links": [ - 204, - 225 + 204 ], "shape": 3, "slot_index": 3 @@ -362,125 +373,50 @@ ] }, { - "id": 93, - "type": "Florence2toCoordinates", + "id": 102, + "type": "Sam2Segmentation", "pos": [ - 388, - -336 + 440, + -120 ], - "size": { - "0": 210, - "1": 58 - }, - "flags": {}, - "order": 7, - "mode": 0, - "inputs": [ - { - "name": "data", - "type": "JSON", - "link": 204 - } + "size": [ + 314.5386123916544, + 162 ], - "outputs": [ - { - "name": "coordinates", - "type": "STRING", - "links": [ - 234 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "Florence2toCoordinates" - }, - "widgets_values": [ - "0, 1" - ] - }, - { - "id": 98, - "type": "Florence2toCoordinates", - "pos": [ - 392, - -231 - ], - "size": { - "0": 210, - "1": 58 - }, "flags": {}, "order": 8, "mode": 0, - "inputs": [ - { - "name": "data", - "type": "JSON", - "link": 225 - } - ], - "outputs": [ - { - "name": "coordinates", - "type": "STRING", - "links": [ - 235 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "Florence2toCoordinates" - }, - "widgets_values": [ - "2" - ] - }, - { - "id": 101, - "type": "Sam2Segmentation", - "pos": [ - 437, - -89 - ], - "size": [ - 315, - 126 - ], - "flags": {}, - "order": 9, - "mode": 0, "inputs": [ { "name": "sam2_model", "type": "SAM2MODEL", - "link": 231 + "link": 236 }, { "name": "image", "type": "IMAGE", - "link": 232 + "link": 237 + }, + { + "name": "bboxes", + "type": "BBOX", + "link": 239 }, { "name": "coordinates_positive", "type": "STRING", - "link": 234, + "link": null, "widget": { "name": "coordinates_positive" - }, - "slot_index": 2 + } }, { "name": "coordinates_negative", "type": "STRING", - "link": 235, + "link": null, "widget": { "name": "coordinates_negative" - }, - "slot_index": 3 + } } ], "outputs": [ @@ -488,19 +424,53 @@ "name": "mask", "type": "MASK", "links": [ - 233 + 238 ], - "shape": 3, - "slot_index": 0 + "shape": 3 } ], "properties": { "Node name for S&R": "Sam2Segmentation" }, "widgets_values": [ - "", true, - "" + "", + "", + true + ] + }, + { + "id": 88, + "type": "DownloadAndLoadFlorence2Model", + "pos": [ + -470, + -777 + ], + "size": { + "0": 315, + "1": 106 + }, + "flags": {}, + "order": 2, + "mode": 0, + "outputs": [ + { + "name": "florence2_model", + "type": "FL2MODEL", + "links": [ + 197 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadFlorence2Model" + }, + "widgets_values": [ + "microsoft/Florence-2-base", + "fp16", + "sdpa" ] } ], @@ -553,14 +523,6 @@ 0, "IMAGE" ], - [ - 225, - 87, - 3, - 98, - 0, - "JSON" - ], [ 226, 72, @@ -570,44 +532,36 @@ "IMAGE" ], [ - 231, + 236, 66, 0, - 101, + 102, 0, "SAM2MODEL" ], [ - 232, + 237, 72, 0, - 101, + 102, 1, "IMAGE" ], [ - 233, - 101, + 238, + 102, 0, 84, 1, "MASK" ], [ - 234, + 239, 93, - 0, - 101, + 1, + 102, 2, - "STRING" - ], - [ - 235, - 98, - 0, - 101, - 3, - "STRING" + "BBOX" ] ], "groups": [], @@ -616,8 +570,8 @@ "ds": { "scale": 0.7627768444385467, "offset": [ - 663.0888928776299, - 959.3311266780172 + 564.3268832902941, + 896.4031145502903 ] } }, diff --git a/examples/points_segment_example.json b/examples/points_segment_example.json index 62fe6e0..7c97f42 100644 --- a/examples/points_segment_example.json +++ b/examples/points_segment_example.json @@ -1,6 +1,6 @@ { - "last_node_id": 97, - "last_link_id": 218, + "last_node_id": 98, + "last_link_id": 222, "nodes": [ { "id": 83, @@ -60,7 +60,7 @@ "name": "sam2_model", "type": "SAM2MODEL", "links": [ - 214 + 219 ], "shape": 3, "slot_index": 0 @@ -124,7 +124,7 @@ "type": "IMAGE", "links": [ 192, - 215 + 220 ], "shape": 3, "slot_index": 0 @@ -178,7 +178,7 @@ { "name": "mask", "type": "MASK", - "link": 216, + "link": 222, "slot_index": 1 } ], @@ -199,69 +199,6 @@ false ] }, - { - "id": 97, - "type": "Sam2Segmentation", - "pos": [ - 374, - -65 - ], - "size": [ - 315, - 126 - ], - "flags": {}, - "order": 4, - "mode": 0, - "inputs": [ - { - "name": "sam2_model", - "type": "SAM2MODEL", - "link": 214 - }, - { - "name": "image", - "type": "IMAGE", - "link": 215 - }, - { - "name": "coordinates_positive", - "type": "STRING", - "link": 218, - "widget": { - "name": "coordinates_positive" - }, - "slot_index": 2 - }, - { - "name": "coordinates_negative", - "type": "STRING", - "link": null, - "widget": { - "name": "coordinates_negative" - } - } - ], - "outputs": [ - { - "name": "mask", - "type": "MASK", - "links": [ - 216 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "Sam2Segmentation" - }, - "widgets_values": [ - "", - true, - "" - ] - }, { "id": 96, "type": "SplineEditor", @@ -287,7 +224,7 @@ "name": "coord_str", "type": "STRING", "links": [ - 218 + 221 ], "shape": 3, "slot_index": 1 @@ -317,12 +254,12 @@ }, "widgets_values": [ "[{\"x\":295.0535292738977,\"y\":285.53567349086876},{\"x\":377.54161272681534,\"y\":287.12198278804027}]", - "[{\"x\":295.05352783203125,\"y\":285.5356750488281},{\"x\":377.5416259765625,\"y\":287.1219787597656}]", + "[{\"x\":295.0535292738977,\"y\":285.53567349086876},{\"x\":377.54161272681534,\"y\":287.12198278804027}]", 768, 512, 2, - "path", - "cardinal", + "controlpoints", + "linear", 0.5, 1, "list", @@ -331,6 +268,73 @@ null, null ] + }, + { + "id": 98, + "type": "Sam2Segmentation", + "pos": [ + 380, + -100 + ], + "size": [ + 316.01859966191455, + 162 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "sam2_model", + "type": "SAM2MODEL", + "link": 219 + }, + { + "name": "image", + "type": "IMAGE", + "link": 220 + }, + { + "name": "bboxes", + "type": "BBOX", + "link": null + }, + { + "name": "coordinates_positive", + "type": "STRING", + "link": 221, + "widget": { + "name": "coordinates_positive" + } + }, + { + "name": "coordinates_negative", + "type": "STRING", + "link": null, + "widget": { + "name": "coordinates_negative" + } + } + ], + "outputs": [ + { + "name": "mask", + "type": "MASK", + "links": [ + 222 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "Sam2Segmentation" + }, + "widgets_values": [ + true, + "", + "", + false + ] } ], "links": [ @@ -351,36 +355,36 @@ "IMAGE" ], [ - 214, + 219, 66, 0, - 97, + 98, 0, "SAM2MODEL" ], [ - 215, + 220, 72, 0, - 97, + 98, 1, "IMAGE" ], [ - 216, - 97, + 221, + 96, + 1, + 98, + 3, + "STRING" + ], + [ + 222, + 98, 0, 84, 1, "MASK" - ], - [ - 218, - 96, - 1, - 97, - 2, - "STRING" ] ], "groups": [], @@ -389,8 +393,8 @@ "ds": { "scale": 0.6303940863128483, "offset": [ - 725.3145306037434, - 1267.6359603669778 + 672.4375863049054, + 1167.1697371529935 ] } }, diff --git a/nodes.py b/nodes.py index 790cd77..38fd3f5 100644 --- a/nodes.py +++ b/nodes.py @@ -97,12 +97,13 @@ class Florence2toCoordinates: }, } - RETURN_TYPES = ("STRING", ) - RETURN_NAMES =("coordinates", ) + RETURN_TYPES = ("STRING", "BBOX") + RETURN_NAMES =("center_coordinates", "bboxes") FUNCTION = "segment" CATEGORY = "SAM2" def segment(self, data, index): + print(data) try: coordinates = coordinates.replace("'", '"') coordinates = json.loads(coordinates) @@ -118,8 +119,9 @@ class Florence2toCoordinates: indexes = [int(i) for i in index.split(",")] else: # If index is empty, use all indices from data[0] indexes = list(range(len(data[0]))) - + print("Indexes:", indexes) + bboxes = [] for idx in indexes: if 0 <= idx < len(data[0]): @@ -129,12 +131,13 @@ class Florence2toCoordinates: center_x = int((min_x + max_x) / 2) center_y = int((min_y + max_y) / 2) center_points.append({"x": center_x, "y": center_y}) + bboxes.append(bbox) else: raise ValueError(f"There's nothing in index: {idx}") coordinates = json.dumps(center_points) print("Coordinates:", coordinates) - return (coordinates,) + return (coordinates, bboxes) class Sam2Segmentation: @classmethod @@ -143,13 +146,14 @@ class Sam2Segmentation: "required": { "sam2_model": ("SAM2MODEL", ), "image": ("IMAGE", ), - "coordinates_positive": ("STRING", {"forceInput": True}), - "keep_model_loaded": ("BOOLEAN", {"default": True}), }, "optional": { + "coordinates_positive": ("STRING", {"forceInput": True}), "coordinates_negative": ("STRING", {"forceInput": True}), - "individual_points": ("BOOLEAN", {"default": False}), + "bboxes": ("BBOX", ), + "individual_objects": ("BOOLEAN", {"default": False}), + }, } @@ -158,7 +162,7 @@ class Sam2Segmentation: FUNCTION = "segment" CATEGORY = "SAM2" - def segment(self, image, sam2_model, coordinates_positive, keep_model_loaded, coordinates_negative=None, individual_points=False): + def segment(self, image, sam2_model, keep_model_loaded, coordinates_positive=None, coordinates_negative=None, individual_objects=False, bboxes=None): offload_device = mm.unet_offload_device() model = sam2_model["model"] device = sam2_model["device"] @@ -172,47 +176,82 @@ class Sam2Segmentation: if segmentor == 'single_image' and B > 1: raise ValueError("Use video segmentor for multiple frames") + if segmentor == 'video' and bboxes is not None: + raise ValueError("Video segmentor doesn't support bboxes") + if segmentor == 'video': # video model needs images resized first thing model_input_image_size = model.image_size print("Resizing to model input image size: ", model_input_image_size) image = common_upscale(image.movedim(-1,1), model_input_image_size, model_input_image_size, "bilinear", "disabled").movedim(1,-1) - try: - coordinates_positive = json.loads(coordinates_positive.replace("'", '"')) - coordinates_positive = [(coord['x'], coord['y']) for coord in coordinates_positive] + #handle point coordinates + if coordinates_positive is not None: + try: + coordinates_positive = json.loads(coordinates_positive.replace("'", '"')) + coordinates_positive = [(coord['x'], coord['y']) for coord in coordinates_positive] + if coordinates_negative is not None: + coordinates_negative = json.loads(coordinates_negative.replace("'", '"')) + coordinates_negative = [(coord['x'], coord['y']) for coord in coordinates_negative] + except: + pass + + if not individual_objects: + positive_point_coords = np.atleast_2d(np.array(coordinates_positive)) + else: + positive_point_coords = np.array([np.atleast_2d(coord) for coord in coordinates_positive]) + if coordinates_negative is not None: - coordinates_negative = json.loads(coordinates_negative.replace("'", '"')) - coordinates_negative = [(coord['x'], coord['y']) for coord in coordinates_negative] - except: - coordinates_positive = coordinates_positive + negative_point_coords = np.array(coordinates_negative) + # Ensure both positive and negative coords are lists of 2D arrays if individual_objects is True + if individual_objects: + if negative_point_coords.ndim == 2: + negative_point_coords = negative_point_coords[:, np.newaxis, :] + final_coords = np.concatenate((positive_point_coords, negative_point_coords), axis=1) + else: + final_coords = np.concatenate((positive_point_coords, negative_point_coords), axis=0) + else: + final_coords = positive_point_coords + + # Handle possible bboxes + if bboxes is not None: + boxes_np_batch = [] + for bbox_list in bboxes: + boxes_np = [] + for bbox in bbox_list: + boxes_np.append(bbox) + boxes_np = np.array(boxes_np) + boxes_np_batch.append(boxes_np) + if individual_objects: + final_box = np.array(boxes_np_batch) + else: + final_box = np.array(boxes_np) + final_labels = None + + #handle labels + if coordinates_positive is not None: + if not individual_objects: + positive_point_labels = np.ones(len(positive_point_coords)) + else: + positive_labels = [] + for point in positive_point_coords: + positive_labels.append(np.array([1])) # 1) + positive_point_labels = np.stack(positive_labels, axis=0) + if coordinates_negative is not None: - coordinates_negative = coordinates_negative - - positive_point_coords = np.array(coordinates_positive) - positive_point_labels = [1] * len(positive_point_coords) # 1 = positive - positive_point_labels = np.array(positive_point_labels) - print("positive coordinates: ", positive_point_coords) - - if coordinates_negative is not None: - negative_point_coords = np.array(coordinates_negative) - negative_point_labels = [0] * len(negative_point_coords) # 0 = negative - negative_point_labels = np.array(negative_point_labels) - print("negative coordinates: ", negative_point_coords) - - # Combine coordinates and labels - else: - negative_point_coords = np.empty((0, 2)) - negative_point_labels = np.array([]) - # Ensure both positive and negative coordinates are 2D arrays - positive_point_coords = np.atleast_2d(positive_point_coords) - negative_point_coords = np.atleast_2d(negative_point_coords) - - # Ensure both positive and negative labels are 1D arrays - positive_point_labels = np.atleast_1d(positive_point_labels) - negative_point_labels = np.atleast_1d(negative_point_labels) - - combined_coords = np.concatenate((positive_point_coords, negative_point_coords), axis=0) - combined_labels = np.concatenate((positive_point_labels, negative_point_labels), axis=0) + if not individual_objects: + negative_point_labels = np.zeros(len(negative_point_coords)) # 0 = negative + final_labels = np.concatenate((positive_point_labels, negative_point_labels), axis=0) + else: + negative_labels = [] + for point in positive_point_coords: + negative_labels.append(np.array([0])) # 1) + negative_point_labels = np.stack(negative_labels, axis=0) + #combine labels + final_labels = np.concatenate((positive_point_labels, negative_point_labels), axis=1) + else: + final_labels = positive_point_labels + print("combined labels: ", final_labels) + print("combined labels shape: ", final_labels.shape) mask_list = [] try: @@ -225,16 +264,26 @@ class Sam2Segmentation: if image.shape[0] == 1: model.set_image(image_np) masks, scores, logits = model.predict( - point_coords=combined_coords, - point_labels=combined_labels, - multimask_output=True, + point_coords=final_coords if coordinates_positive is not None else None, + point_labels=final_labels, + box=final_box if bboxes is not None else None, + multimask_output=True if not individual_objects else False, ) - - sorted_ind = np.argsort(scores)[::-1] - masks = masks[sorted_ind][0] #choose only the best result for now - scores = scores[sorted_ind] - logits = logits[sorted_ind] - mask_list.append(np.expand_dims(masks, axis=0)) + + if masks.ndim == 3: + sorted_ind = np.argsort(scores)[::-1] + masks = masks[sorted_ind][0] #choose only the best result for now + scores = scores[sorted_ind] + logits = logits[sorted_ind] + mask_list.append(np.expand_dims(masks, axis=0)) + else: + _, _, H, W = masks.shape + # Combine masks for all object IDs in the frame + combined_mask = np.zeros((H, W), dtype=bool) + for mask in masks: + combined_mask = np.logical_or(combined_mask, mask) + combined_mask = combined_mask.astype(np.uint8) + mask_list.append(combined_mask) else: mask_list = [] @@ -242,23 +291,22 @@ class Sam2Segmentation: model.reset_state(self.inference_state) self.inference_state = model.init_state(image.permute(0, 3, 1, 2).contiguous(), H, W, device=device) - if individual_points: - for i, (coord, label) in enumerate(zip(combined_coords, combined_labels)): + if individual_objects: + for i, (coord, label) in enumerate(zip(final_coords, final_labels)): _, out_obj_ids, out_mask_logits = model.add_new_points( inference_state=self.inference_state, frame_idx=0, obj_id=i, - points=[combined_coords[i]], - labels=[combined_labels[i]], + points=final_coords[i], + labels=final_labels[i], ) - else: _, out_obj_ids, out_mask_logits = model.add_new_points( inference_state=self.inference_state, frame_idx=0, obj_id=1, - points=combined_coords, - labels=combined_labels, + points=final_coords, + labels=final_labels, ) pbar = ProgressBar(B) @@ -269,7 +317,7 @@ class Sam2Segmentation: for i, out_obj_id in enumerate(out_obj_ids) } pbar.update(1) - if individual_points: + if individual_objects: _, _, H, W = out_mask_logits.shape # Combine masks for all object IDs in the frame combined_mask = np.zeros((H, W), dtype=np.uint8) @@ -278,7 +326,7 @@ class Sam2Segmentation: combined_mask = np.logical_or(combined_mask, out_mask) video_segments[out_frame_idx] = combined_mask - if individual_points: + if individual_objects: for frame_idx, combined_mask in video_segments.items(): mask_list.append(combined_mask) else: