support individual objects and bboxes for single image detection

This commit is contained in:
kijai
2024-08-02 01:37:57 +03:00
parent 05afefc447
commit 468184d7d7
3 changed files with 347 additions and 341 deletions
+147 -193
View File
@@ -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
]
}
},
+92 -88
View File
@@ -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
]
}
},
+108 -60
View File
@@ -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: