segment-everything 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- segment_everything/__init__.py +5 -0
- segment_everything/augmentation/albumentations_helper.py +0 -0
- segment_everything/detect_and_segment.py +131 -0
- segment_everything/napari_helper.py +15 -0
- segment_everything/prompt_generator.py +188 -0
- segment_everything/py.typed +5 -0
- segment_everything/stacked_label_dataset.py +113 -0
- segment_everything/stacked_labels.py +428 -0
- segment_everything/vendored/PromptGuidedDecoder/Prompt_guided_Mask_Decoder.pt +0 -0
- segment_everything/vendored/__init__.py +5 -0
- segment_everything/vendored/dice.py +158 -0
- segment_everything/vendored/efficientvit/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/data_provider/__init__.py +7 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/bbox.py +30 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/color_aug.py +78 -0
- segment_everything/vendored/efficientvit/apps/data_provider/base.py +254 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_loader.py +1538 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_worker.py +357 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/controller.py +100 -0
- segment_everything/vendored/efficientvit/apps/setup.py +150 -0
- segment_everything/vendored/efficientvit/apps/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/trainer/base.py +318 -0
- segment_everything/vendored/efficientvit/apps/trainer/run_config.py +129 -0
- segment_everything/vendored/efficientvit/apps/utils/__init__.py +12 -0
- segment_everything/vendored/efficientvit/apps/utils/dist.py +32 -0
- segment_everything/vendored/efficientvit/apps/utils/ema.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/export.py +45 -0
- segment_everything/vendored/efficientvit/apps/utils/init.py +66 -0
- segment_everything/vendored/efficientvit/apps/utils/lr.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/metric.py +43 -0
- segment_everything/vendored/efficientvit/apps/utils/misc.py +101 -0
- segment_everything/vendored/efficientvit/apps/utils/opt.py +28 -0
- segment_everything/vendored/efficientvit/cls_model_zoo.py +79 -0
- segment_everything/vendored/efficientvit/clscore/__init__.py +0 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/__init__.py +5 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/imagenet.py +142 -0
- segment_everything/vendored/efficientvit/clscore/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_run_config.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_trainer.py +265 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/label_smooth.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/metric.py +23 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/mixup.py +67 -0
- segment_everything/vendored/efficientvit/models/__init__.py +0 -0
- segment_everything/vendored/efficientvit/models/efficientvit/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/efficientvit/backbone.py +380 -0
- segment_everything/vendored/efficientvit/models/efficientvit/cls.py +188 -0
- segment_everything/vendored/efficientvit/models/efficientvit/sam.py +181 -0
- segment_everything/vendored/efficientvit/models/efficientvit/seg.py +373 -0
- segment_everything/vendored/efficientvit/models/nn/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/nn/act.py +30 -0
- segment_everything/vendored/efficientvit/models/nn/drop.py +104 -0
- segment_everything/vendored/efficientvit/models/nn/norm.py +164 -0
- segment_everything/vendored/efficientvit/models/nn/ops.py +597 -0
- segment_everything/vendored/efficientvit/models/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/models/utils/list.py +53 -0
- segment_everything/vendored/efficientvit/models/utils/network.py +73 -0
- segment_everything/vendored/efficientvit/models/utils/random.py +65 -0
- segment_everything/vendored/efficientvit/sam_model_zoo.py +45 -0
- segment_everything/vendored/efficientvit/seg_model_zoo.py +70 -0
- segment_everything/vendored/get_object_aware.py +26 -0
- segment_everything/vendored/mobilesamv2/__init__.py +16 -0
- segment_everything/vendored/mobilesamv2/automatic_mask_generator.py +415 -0
- segment_everything/vendored/mobilesamv2/build_sam.py +246 -0
- segment_everything/vendored/mobilesamv2/modeling/__init__.py +11 -0
- segment_everything/vendored/mobilesamv2/modeling/common.py +43 -0
- segment_everything/vendored/mobilesamv2/modeling/image_encoder.py +394 -0
- segment_everything/vendored/mobilesamv2/modeling/mask_decoder.py +213 -0
- segment_everything/vendored/mobilesamv2/modeling/prompt_encoder.py +217 -0
- segment_everything/vendored/mobilesamv2/modeling/sam.py +203 -0
- segment_everything/vendored/mobilesamv2/modeling/transformer.py +240 -0
- segment_everything/vendored/mobilesamv2/predictor.py +384 -0
- segment_everything/vendored/mobilesamv2/utils/__init__.py +5 -0
- segment_everything/vendored/mobilesamv2/utils/amg.py +347 -0
- segment_everything/vendored/mobilesamv2/utils/onnx.py +144 -0
- segment_everything/vendored/mobilesamv2/utils/transforms.py +103 -0
- segment_everything/vendored/object_detection/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/nn/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autobackend.py +658 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autoshape.py +397 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/__init__.py +110 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/block.py +304 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/conv.py +297 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/head.py +468 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/transformer.py +378 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/utils.py +78 -0
- segment_everything/vendored/object_detection/ultralytics/nn/tasks.py +1049 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/__init__.py +6 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/model.py +104 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/predict.py +95 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/__init__.py +588 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/default.yaml +117 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/annotator.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/augment.py +899 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/base.py +286 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/build.py +213 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/converter.py +358 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/stream_loaders.py +459 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset.py +274 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset_wrappers.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/download_weights.sh +18 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco.sh +60 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco128.sh +17 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_imagenet.sh +51 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/utils.py +716 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/exporter.py +1214 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/model.py +641 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/predictor.py +461 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/results.py +741 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/__init__.py +893 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/autobatch.py +108 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/base.py +212 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/checks.py +547 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/dist.py +67 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/downloads.py +353 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/errors.py +12 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/files.py +100 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/instance.py +391 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/loss.py +579 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/metrics.py +1189 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/ops.py +870 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/patches.py +45 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/plotting.py +767 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tal.py +276 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/torch_utils.py +684 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tuner.py +54 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/predict.py +69 -0
- segment_everything/vendored/tinyvit/__init__.py +2 -0
- segment_everything/vendored/tinyvit/tiny_vit.py +867 -0
- segment_everything/weights_helper.py +124 -0
- segment_everything-0.1.0.dist-info/METADATA +53 -0
- segment_everything-0.1.0.dist-info/RECORD +145 -0
- segment_everything-0.1.0.dist-info/WHEEL +4 -0
- segment_everything-0.1.0.dist-info/licenses/LICENSE +28 -0
|
File without changes
|
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
# -*- coding: utf-8 -*-
|
|
3
|
+
|
|
4
|
+
from segment_everything.weights_helper import create_sam_model
|
|
5
|
+
from segment_everything.stacked_labels import StackedLabels
|
|
6
|
+
from segment_everything.vendored.mobilesamv2 import SamPredictor as SamPredictorV2
|
|
7
|
+
from segment_everything.weights_helper import get_device
|
|
8
|
+
from typing import Any, Generator, List
|
|
9
|
+
import torch
|
|
10
|
+
import os
|
|
11
|
+
from segment_anything.utils.amg import calculate_stability_score
|
|
12
|
+
import gc
|
|
13
|
+
|
|
14
|
+
current_dir = os.path.dirname(__file__)
|
|
15
|
+
|
|
16
|
+
def batch_iterator(batch_size: int, *args) -> Generator[List[Any], None, None]:
|
|
17
|
+
assert len(args) > 0 and all(
|
|
18
|
+
len(a) == len(args[0]) for a in args
|
|
19
|
+
), "Batched iteration must have inputs of all the same size."
|
|
20
|
+
n_batches = len(args[0]) // batch_size + int(
|
|
21
|
+
len(args[0]) % batch_size != 0
|
|
22
|
+
)
|
|
23
|
+
for b in range(n_batches):
|
|
24
|
+
yield [arg[b * batch_size : (b + 1) * batch_size] for arg in args]
|
|
25
|
+
|
|
26
|
+
def segment_from_stacked_labels(stacked_labels, model_type, device=None):
|
|
27
|
+
""" given stacked labels and a model re-segment all masks by calling sam on each bounding box
|
|
28
|
+
|
|
29
|
+
this function is useful when we have stacked labels with bounding boxes only (or other non-optimal masks)
|
|
30
|
+
and we want to resegment the masks using sam
|
|
31
|
+
|
|
32
|
+
Args:
|
|
33
|
+
stacked_labels (StackedLabels): input stacked labels
|
|
34
|
+
model (SAM model): sam model
|
|
35
|
+
device (string): device type
|
|
36
|
+
|
|
37
|
+
Returns:
|
|
38
|
+
StackedLabels: new stacked labels with resegmented labels
|
|
39
|
+
"""
|
|
40
|
+
if device is None:
|
|
41
|
+
device = get_device()
|
|
42
|
+
model = create_sam_model(model_type, device)
|
|
43
|
+
bbox_array = stacked_labels.get_bbox_np()
|
|
44
|
+
sam_masks = segment_from_bbox(stacked_labels.image, bbox_array, model, device)
|
|
45
|
+
|
|
46
|
+
return StackedLabels(sam_masks, stacked_labels.image)
|
|
47
|
+
|
|
48
|
+
def segment_from_bbox(img, bounding_boxes, model, device):
|
|
49
|
+
"""
|
|
50
|
+
Segments everything given the bounding boxes of the objects and the mobileSAMv2 prediction model.
|
|
51
|
+
Code from mobileSAMv2
|
|
52
|
+
"""
|
|
53
|
+
predictor = SamPredictorV2(model)
|
|
54
|
+
predictor.set_image(img)
|
|
55
|
+
|
|
56
|
+
input_boxes = predictor.transform.apply_boxes(
|
|
57
|
+
bounding_boxes, predictor.original_size
|
|
58
|
+
) # Does this need to be transformed?
|
|
59
|
+
if device == "cuda":
|
|
60
|
+
input_boxes = torch.from_numpy(input_boxes).cuda()
|
|
61
|
+
elif device == "cpu":
|
|
62
|
+
input_boxes = torch.from_numpy(input_boxes)
|
|
63
|
+
sam_mask = []
|
|
64
|
+
|
|
65
|
+
predicted_ious = []
|
|
66
|
+
stability_scores = []
|
|
67
|
+
|
|
68
|
+
image_embedding = predictor.features
|
|
69
|
+
image_embedding = torch.repeat_interleave(image_embedding, 400, dim=0)
|
|
70
|
+
|
|
71
|
+
prompt_embedding = model.prompt_encoder.get_dense_pe()
|
|
72
|
+
prompt_embedding = torch.repeat_interleave(prompt_embedding, 400, dim=0)
|
|
73
|
+
|
|
74
|
+
for (boxes,) in batch_iterator(200, input_boxes):
|
|
75
|
+
with torch.no_grad():
|
|
76
|
+
image_embedding = image_embedding[0 : boxes.shape[0], :, :, :]
|
|
77
|
+
prompt_embedding = prompt_embedding[0 : boxes.shape[0], :, :, :]
|
|
78
|
+
sparse_embeddings, dense_embeddings = model.prompt_encoder(
|
|
79
|
+
points=None,
|
|
80
|
+
boxes=boxes,
|
|
81
|
+
masks=None,
|
|
82
|
+
)
|
|
83
|
+
low_res_masks, pred_ious = model.mask_decoder(
|
|
84
|
+
image_embeddings=image_embedding,
|
|
85
|
+
image_pe=prompt_embedding,
|
|
86
|
+
sparse_prompt_embeddings=sparse_embeddings,
|
|
87
|
+
dense_prompt_embeddings=dense_embeddings,
|
|
88
|
+
multimask_output=False,
|
|
89
|
+
simple_type=True,
|
|
90
|
+
)
|
|
91
|
+
low_res_masks = predictor.model.postprocess_masks(
|
|
92
|
+
low_res_masks, predictor.input_size, predictor.original_size
|
|
93
|
+
)
|
|
94
|
+
model.threshold_offset = 1
|
|
95
|
+
stability_score = (
|
|
96
|
+
calculate_stability_score(
|
|
97
|
+
low_res_masks,
|
|
98
|
+
model.mask_threshold,
|
|
99
|
+
model.threshold_offset,
|
|
100
|
+
)
|
|
101
|
+
.cpu()
|
|
102
|
+
.numpy()
|
|
103
|
+
)
|
|
104
|
+
sam_mask_pre = (low_res_masks > model.mask_threshold) * 1.0
|
|
105
|
+
sam_mask.append(sam_mask_pre.squeeze(1))
|
|
106
|
+
predicted_ious.extend(pred_ious.cpu().numpy().flatten().tolist())
|
|
107
|
+
stability_scores.extend(stability_score.flatten().tolist())
|
|
108
|
+
|
|
109
|
+
sam_mask = torch.cat(sam_mask)
|
|
110
|
+
# predicted_ious = pred_ious.cpu().numpy()
|
|
111
|
+
cpu_segmentations = sam_mask.cpu().numpy()
|
|
112
|
+
del sam_mask
|
|
113
|
+
|
|
114
|
+
gc.collect()
|
|
115
|
+
torch.cuda.empty_cache()
|
|
116
|
+
|
|
117
|
+
curr_anns = []
|
|
118
|
+
for idx in range(len(cpu_segmentations)):
|
|
119
|
+
ann = {
|
|
120
|
+
"segmentation": cpu_segmentations[idx],
|
|
121
|
+
"area": sum(sum(cpu_segmentations[idx])),
|
|
122
|
+
"predicted_iou": predicted_ious[idx],
|
|
123
|
+
"stability_score": stability_scores[idx],
|
|
124
|
+
"prompt_bbox": bounding_boxes[idx],
|
|
125
|
+
}
|
|
126
|
+
if (
|
|
127
|
+
cpu_segmentations[idx].max() < 1
|
|
128
|
+
): # this means that bboxes won't always == segmentations
|
|
129
|
+
continue
|
|
130
|
+
curr_anns.append(ann)
|
|
131
|
+
return curr_anns
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
import napari
|
|
2
|
+
from segment_everything.stacked_labels import StackedLabels
|
|
3
|
+
from napari_segment_everything import segment_everything
|
|
4
|
+
|
|
5
|
+
def stacked_labels_to_napari(stacked_labels):
|
|
6
|
+
viewer = napari.Viewer()
|
|
7
|
+
segment_everything_widget=segment_everything.NapariSegmentEverything(viewer)
|
|
8
|
+
viewer.window.add_dock_widget(segment_everything_widget)
|
|
9
|
+
segment_everything_widget.load_project(stacked_labels.image, stacked_labels.mask_list)
|
|
10
|
+
|
|
11
|
+
def to_napari(image, mask_list):
|
|
12
|
+
viewer = napari.Viewer()
|
|
13
|
+
segment_everything_widget=segment_everything.NapariSegmentEverything(viewer)
|
|
14
|
+
viewer.window.add_dock_widget(segment_everything_widget)
|
|
15
|
+
segment_everything_widget.load_project(image, mask_list)
|
|
@@ -0,0 +1,188 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
# -*- coding: utf-8 -*-
|
|
3
|
+
"""
|
|
4
|
+
Created on Mon Apr 22 23:58:14 2024
|
|
5
|
+
|
|
6
|
+
@author: ian
|
|
7
|
+
"""
|
|
8
|
+
import cv2
|
|
9
|
+
import os, sys
|
|
10
|
+
from torchvision.models.detection import fasterrcnn_mobilenet_v3_large_fpn
|
|
11
|
+
from torchvision.transforms import ToTensor
|
|
12
|
+
from torchvision.ops import nms
|
|
13
|
+
import torch
|
|
14
|
+
|
|
15
|
+
from segment_everything.vendored.get_object_aware import get_object_aware_model
|
|
16
|
+
|
|
17
|
+
class BaseDetector:
|
|
18
|
+
def __init__(self, model_path, trainable=False):
|
|
19
|
+
self.model_path = model_path
|
|
20
|
+
self.model_name = self.model_path.split("/")[-1]
|
|
21
|
+
self.trainable = trainable
|
|
22
|
+
|
|
23
|
+
def train(self, training_data):
|
|
24
|
+
raise NotImplementedError()
|
|
25
|
+
|
|
26
|
+
def predict(self, image_data):
|
|
27
|
+
raise NotImplementedError()
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class YoloDetector(BaseDetector):
|
|
31
|
+
def __init__(self, model_path, model_type, device, trainable=False):
|
|
32
|
+
super().__init__(model_path)
|
|
33
|
+
self.model_type = model_type
|
|
34
|
+
|
|
35
|
+
if (model_type == "ObjectAwareModelFromMobileSamV2"):
|
|
36
|
+
self.model = get_object_aware_model(model_path) #ObjectAwareModel(model_path)
|
|
37
|
+
else:
|
|
38
|
+
from ultralytics import YOLO
|
|
39
|
+
self.model = YOLO(model_path)
|
|
40
|
+
|
|
41
|
+
self.device = device
|
|
42
|
+
|
|
43
|
+
def train(self):
|
|
44
|
+
print(
|
|
45
|
+
"YOLO detector is not yet trainable, use RcnnDetector for training"
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
def get_results(
|
|
49
|
+
self,
|
|
50
|
+
image_data,
|
|
51
|
+
retina_masks=True,
|
|
52
|
+
imgsz=1024,
|
|
53
|
+
conf=0.4,
|
|
54
|
+
iou=0.9,
|
|
55
|
+
max_det=400,
|
|
56
|
+
):
|
|
57
|
+
"""
|
|
58
|
+
Runs YOLO and returns the YOLO results
|
|
59
|
+
|
|
60
|
+
Parameters
|
|
61
|
+
----------
|
|
62
|
+
image : numpy.ndarray
|
|
63
|
+
A 2D-image in grayscale or RGB.
|
|
64
|
+
imgsz : INT, optional
|
|
65
|
+
Size of the input image. The default is 1024.
|
|
66
|
+
conf : FLOAT, optional
|
|
67
|
+
Confidence threshold for the bounding boxes. Lower means more boxes will be detected. The default is 0.4.
|
|
68
|
+
iou : FLOAT, optional
|
|
69
|
+
Threshold for how many intersecting bounding boxes should be allowed. Lower means fewer intersecting boxes will be returned. The default is 0.9.
|
|
70
|
+
max_det : INT, optional
|
|
71
|
+
Maximum number of detections that will be returned. The default is 400.
|
|
72
|
+
|
|
73
|
+
Returns
|
|
74
|
+
-------
|
|
75
|
+
obj_results: YOLO results objects
|
|
76
|
+
"""
|
|
77
|
+
|
|
78
|
+
image_cv2 = cv2.cvtColor(image_data, cv2.COLOR_BGR2RGB)
|
|
79
|
+
|
|
80
|
+
obj_results = self.model.predict(
|
|
81
|
+
image_cv2,
|
|
82
|
+
device=self.device,
|
|
83
|
+
retina_masks=True,
|
|
84
|
+
imgsz=imgsz,
|
|
85
|
+
conf=conf,
|
|
86
|
+
iou=iou,
|
|
87
|
+
max_det=max_det,
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
return obj_results
|
|
91
|
+
|
|
92
|
+
def get_bounding_boxes(
|
|
93
|
+
self,
|
|
94
|
+
image_data,
|
|
95
|
+
retina_masks=True,
|
|
96
|
+
imgsz=1024,
|
|
97
|
+
conf=0.4,
|
|
98
|
+
iou=0.9,
|
|
99
|
+
max_det=400,
|
|
100
|
+
):
|
|
101
|
+
"""
|
|
102
|
+
Generates a series of bounding boxes in xyxy-format from an image, using the YOLOv8 ObjectAwareModel.
|
|
103
|
+
|
|
104
|
+
Parameters
|
|
105
|
+
----------
|
|
106
|
+
image : numpy.ndarray
|
|
107
|
+
A 2D-image in grayscale or RGB.
|
|
108
|
+
imgsz : INT, optional
|
|
109
|
+
Size of the input image. The default is 1024.
|
|
110
|
+
conf : FLOAT, optional
|
|
111
|
+
Confidence threshold for the bounding boxes. Lower means more boxes will be detected. The default is 0.4.
|
|
112
|
+
iou : FLOAT, optional
|
|
113
|
+
Threshold for how many intersecting bounding boxes should be allowed. Lower means fewer intersecting boxes will be returned. The default is 0.9.
|
|
114
|
+
max_det : INT, optional
|
|
115
|
+
Maximum number of detections that will be returned. The default is 400.
|
|
116
|
+
|
|
117
|
+
Returns
|
|
118
|
+
-------
|
|
119
|
+
bounding_boxes : numpy.ndarray
|
|
120
|
+
An array of boxes in xyxy-coordinates.
|
|
121
|
+
"""
|
|
122
|
+
|
|
123
|
+
print("Predicting bounding boxes for image data")
|
|
124
|
+
obj_results = self.get_results(image_data, retina_masks, imgsz, conf, iou, max_det)
|
|
125
|
+
|
|
126
|
+
return obj_results[0].boxes.xyxy.cpu().numpy()
|
|
127
|
+
|
|
128
|
+
def __str__(self):
|
|
129
|
+
s = f"\n{'Model':<10}: {self.model_name}\n"
|
|
130
|
+
s += f"{'Type':<10}: {str(self.model_type)}\n"
|
|
131
|
+
s += f"{'Trainable':<10}: {str(self.trainable)}"
|
|
132
|
+
return s
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
class RcnnDetector(BaseDetector):
|
|
136
|
+
def __init__(self, model_path, device, trainable=True):
|
|
137
|
+
super().__init__(model_path, trainable)
|
|
138
|
+
self.model_type = "FasterRCNN"
|
|
139
|
+
if device == "mps":
|
|
140
|
+
device = "cpu"
|
|
141
|
+
self.device = device
|
|
142
|
+
self.model = fasterrcnn_mobilenet_v3_large_fpn(
|
|
143
|
+
box_detections_per_img=500,
|
|
144
|
+
).to(self.device)
|
|
145
|
+
self.model.load_state_dict(torch.load(model_path, map_location=self.device))
|
|
146
|
+
|
|
147
|
+
def train(self, training_data):
|
|
148
|
+
if self.trainable:
|
|
149
|
+
print("Training model")
|
|
150
|
+
print(self.model_path)
|
|
151
|
+
print(training_data)
|
|
152
|
+
|
|
153
|
+
def _get_transform(self, train):
|
|
154
|
+
from torchvision.transforms import v2 as T
|
|
155
|
+
|
|
156
|
+
transforms = []
|
|
157
|
+
if train:
|
|
158
|
+
transforms.append(T.RandomHorizontalFlip(0.5))
|
|
159
|
+
transforms.append(T.ToDtype(torch.float, scale=True))
|
|
160
|
+
transforms.append(T.ToPureTensor())
|
|
161
|
+
return T.Compose(transforms)
|
|
162
|
+
|
|
163
|
+
@torch.inference_mode()
|
|
164
|
+
def get_bounding_boxes(self, image_data, conf=0.5, iou=0.2):
|
|
165
|
+
image_cv2 = cv2.cvtColor(image_data, cv2.COLOR_BGR2RGB)
|
|
166
|
+
|
|
167
|
+
print("Predicting bounding boxes for image data")
|
|
168
|
+
convert_tensor = ToTensor()
|
|
169
|
+
eval_transform = self._get_transform(train=False)
|
|
170
|
+
tensor_image = convert_tensor(image_cv2)
|
|
171
|
+
x = eval_transform(tensor_image)
|
|
172
|
+
# convert RGBA -> RGB and move to device
|
|
173
|
+
x = x[:3, ...].to(self.device)
|
|
174
|
+
self.model.eval()
|
|
175
|
+
predictions = self.model([x])
|
|
176
|
+
pred = predictions[0]
|
|
177
|
+
# print(pred)
|
|
178
|
+
idx_after = nms(pred["boxes"], pred["scores"], iou_threshold=iou)
|
|
179
|
+
pred_boxes = pred["boxes"][idx_after]
|
|
180
|
+
pred_scores = pred["scores"][idx_after]
|
|
181
|
+
pred_boxes_conf = pred_boxes[pred_scores > conf]
|
|
182
|
+
return pred_boxes_conf.cpu().numpy()
|
|
183
|
+
|
|
184
|
+
def __str__(self):
|
|
185
|
+
s = f"\n{'Model':<10}: {self.model_name}\n"
|
|
186
|
+
s += f"{'Type':<10}: {str(self.model_type)}\n"
|
|
187
|
+
s += f"{'Trainable':<10}: {str(self.trainable)}"
|
|
188
|
+
return s
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
from torch.utils.data import Dataset
|
|
2
|
+
import random
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
class StackedLabelDataset(Dataset):
|
|
6
|
+
"""
|
|
7
|
+
This class is used to create a dataset with stacked labels. For an m by n image each label image is also
|
|
8
|
+
m by n array and contains only one label. This way overlapping labels can be handled.
|
|
9
|
+
"""
|
|
10
|
+
def __init__(self, data, processor):
|
|
11
|
+
""" initializes the StackedLabelDataset
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
data (list): A list of dictionary objects. Each dictionary object contains an image and a label collection
|
|
15
|
+
processor (_type_): _description_
|
|
16
|
+
"""
|
|
17
|
+
self.data = data
|
|
18
|
+
self.processor = processor
|
|
19
|
+
|
|
20
|
+
def __len__(self):
|
|
21
|
+
return len(self.data)
|
|
22
|
+
|
|
23
|
+
def __getitem__(self, idx):
|
|
24
|
+
""" returns an item from the dataset
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
idx (int): the index of the item to return
|
|
28
|
+
|
|
29
|
+
Returns:
|
|
30
|
+
dict: a dictionary containing the image and label collection
|
|
31
|
+
"""
|
|
32
|
+
result = self.data[idx]
|
|
33
|
+
image = result['image']
|
|
34
|
+
|
|
35
|
+
# if number dims 2
|
|
36
|
+
if len(image.shape) == 2:
|
|
37
|
+
# add fake channels
|
|
38
|
+
image = np.stack([image, image, image], axis=-1)
|
|
39
|
+
|
|
40
|
+
segmentation = result['segmentation']
|
|
41
|
+
|
|
42
|
+
if 'prompt_bbox' in result:
|
|
43
|
+
x_min, y_min, x_max, y_max = result['prompt_bbox']
|
|
44
|
+
H, W = segmentation.shape
|
|
45
|
+
x_min = max(0, x_min - np.random.randint(0, 10))
|
|
46
|
+
x_max = min(W, x_max + np.random.randint(0, 10))
|
|
47
|
+
y_min = max(0, y_min - np.random.randint(0, 10))
|
|
48
|
+
y_max = min(H, y_max + np.random.randint(0, 10))
|
|
49
|
+
box_prompt = [x_min, y_min, x_max, y_max]
|
|
50
|
+
box_prompt = [float(x) for x in box_prompt]
|
|
51
|
+
box_prompt = [[box_prompt]]
|
|
52
|
+
else:
|
|
53
|
+
box_prompt = None
|
|
54
|
+
|
|
55
|
+
if self.random_index == False:
|
|
56
|
+
point_prompt = result['point_coords']
|
|
57
|
+
else:
|
|
58
|
+
y, x = result['indexes']
|
|
59
|
+
idx = np.random.randint(len(x))
|
|
60
|
+
point_prompt = [[x[idx], y[idx]]]
|
|
61
|
+
|
|
62
|
+
'''
|
|
63
|
+
if np.sum(segmentation) == 0:
|
|
64
|
+
print('empty')
|
|
65
|
+
else:
|
|
66
|
+
print('mask')
|
|
67
|
+
|
|
68
|
+
print('centroid is', result['point_coords'])
|
|
69
|
+
print('random prompt', point_prompt)
|
|
70
|
+
print()
|
|
71
|
+
'''
|
|
72
|
+
'''
|
|
73
|
+
if len(x) > 0:
|
|
74
|
+
idx = np.random.randint(len(x))
|
|
75
|
+
prompt = [[x[idx], y[idx]]]
|
|
76
|
+
else:
|
|
77
|
+
prompt = result['point_coords']
|
|
78
|
+
'''
|
|
79
|
+
'''
|
|
80
|
+
number_results = len(results)
|
|
81
|
+
|
|
82
|
+
# random id for label
|
|
83
|
+
results_id = random.randint(0, 4*number_results)
|
|
84
|
+
|
|
85
|
+
# if the label is from the original labels
|
|
86
|
+
if results_id < number_results:
|
|
87
|
+
result = results[results_id]
|
|
88
|
+
segmentation = result['segmentation']
|
|
89
|
+
y, x = np.where(segmentation)
|
|
90
|
+
idx = np.random.randint(len(x))
|
|
91
|
+
prompt = [[x[idx], y[idx]]]
|
|
92
|
+
else:
|
|
93
|
+
all_labels = item['binary']
|
|
94
|
+
not_labeled = np.logical_not(all_labels)
|
|
95
|
+
y, x = np.where(not_labeled)
|
|
96
|
+
idx = np.random.randint(len(x))
|
|
97
|
+
prompt = [[x[idx], y[idx]]]
|
|
98
|
+
segmentation = np.zeros_like(all_labels, dtype=bool)
|
|
99
|
+
#print('random prompt', prompt)
|
|
100
|
+
#print()
|
|
101
|
+
'''
|
|
102
|
+
#print('image shape is',image.shape)
|
|
103
|
+
#print('prompt is',prompt)
|
|
104
|
+
|
|
105
|
+
inputs = self.processor(image, input_points=point_prompt, input_boxes = box_prompt, return_tensors="pt")
|
|
106
|
+
|
|
107
|
+
# remove batch dimension which the processor adds by default
|
|
108
|
+
inputs = {k:v.squeeze(0) for k,v in inputs.items()}
|
|
109
|
+
|
|
110
|
+
inputs['ground_truth'] = segmentation.astype('float32')
|
|
111
|
+
|
|
112
|
+
# prompt =
|
|
113
|
+
return inputs
|