support individual objects and bboxes for single image detection
This commit is contained in:
+147
-193
@@ -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
|
||||
]
|
||||
}
|
||||
},
|
||||
|
||||
@@ -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
|
||||
]
|
||||
}
|
||||
},
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user