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.
Files changed (145) hide show
  1. segment_everything/__init__.py +5 -0
  2. segment_everything/augmentation/albumentations_helper.py +0 -0
  3. segment_everything/detect_and_segment.py +131 -0
  4. segment_everything/napari_helper.py +15 -0
  5. segment_everything/prompt_generator.py +188 -0
  6. segment_everything/py.typed +5 -0
  7. segment_everything/stacked_label_dataset.py +113 -0
  8. segment_everything/stacked_labels.py +428 -0
  9. segment_everything/vendored/PromptGuidedDecoder/Prompt_guided_Mask_Decoder.pt +0 -0
  10. segment_everything/vendored/__init__.py +5 -0
  11. segment_everything/vendored/dice.py +158 -0
  12. segment_everything/vendored/efficientvit/__init__.py +0 -0
  13. segment_everything/vendored/efficientvit/apps/__init__.py +0 -0
  14. segment_everything/vendored/efficientvit/apps/data_provider/__init__.py +7 -0
  15. segment_everything/vendored/efficientvit/apps/data_provider/augment/__init__.py +6 -0
  16. segment_everything/vendored/efficientvit/apps/data_provider/augment/bbox.py +30 -0
  17. segment_everything/vendored/efficientvit/apps/data_provider/augment/color_aug.py +78 -0
  18. segment_everything/vendored/efficientvit/apps/data_provider/base.py +254 -0
  19. segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/__init__.py +6 -0
  20. segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_loader.py +1538 -0
  21. segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_worker.py +357 -0
  22. segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/controller.py +100 -0
  23. segment_everything/vendored/efficientvit/apps/setup.py +150 -0
  24. segment_everything/vendored/efficientvit/apps/trainer/__init__.py +6 -0
  25. segment_everything/vendored/efficientvit/apps/trainer/base.py +318 -0
  26. segment_everything/vendored/efficientvit/apps/trainer/run_config.py +129 -0
  27. segment_everything/vendored/efficientvit/apps/utils/__init__.py +12 -0
  28. segment_everything/vendored/efficientvit/apps/utils/dist.py +32 -0
  29. segment_everything/vendored/efficientvit/apps/utils/ema.py +52 -0
  30. segment_everything/vendored/efficientvit/apps/utils/export.py +45 -0
  31. segment_everything/vendored/efficientvit/apps/utils/init.py +66 -0
  32. segment_everything/vendored/efficientvit/apps/utils/lr.py +52 -0
  33. segment_everything/vendored/efficientvit/apps/utils/metric.py +43 -0
  34. segment_everything/vendored/efficientvit/apps/utils/misc.py +101 -0
  35. segment_everything/vendored/efficientvit/apps/utils/opt.py +28 -0
  36. segment_everything/vendored/efficientvit/cls_model_zoo.py +79 -0
  37. segment_everything/vendored/efficientvit/clscore/__init__.py +0 -0
  38. segment_everything/vendored/efficientvit/clscore/data_provider/__init__.py +5 -0
  39. segment_everything/vendored/efficientvit/clscore/data_provider/imagenet.py +142 -0
  40. segment_everything/vendored/efficientvit/clscore/trainer/__init__.py +6 -0
  41. segment_everything/vendored/efficientvit/clscore/trainer/cls_run_config.py +18 -0
  42. segment_everything/vendored/efficientvit/clscore/trainer/cls_trainer.py +265 -0
  43. segment_everything/vendored/efficientvit/clscore/trainer/utils/__init__.py +7 -0
  44. segment_everything/vendored/efficientvit/clscore/trainer/utils/label_smooth.py +18 -0
  45. segment_everything/vendored/efficientvit/clscore/trainer/utils/metric.py +23 -0
  46. segment_everything/vendored/efficientvit/clscore/trainer/utils/mixup.py +67 -0
  47. segment_everything/vendored/efficientvit/models/__init__.py +0 -0
  48. segment_everything/vendored/efficientvit/models/efficientvit/__init__.py +8 -0
  49. segment_everything/vendored/efficientvit/models/efficientvit/backbone.py +380 -0
  50. segment_everything/vendored/efficientvit/models/efficientvit/cls.py +188 -0
  51. segment_everything/vendored/efficientvit/models/efficientvit/sam.py +181 -0
  52. segment_everything/vendored/efficientvit/models/efficientvit/seg.py +373 -0
  53. segment_everything/vendored/efficientvit/models/nn/__init__.py +8 -0
  54. segment_everything/vendored/efficientvit/models/nn/act.py +30 -0
  55. segment_everything/vendored/efficientvit/models/nn/drop.py +104 -0
  56. segment_everything/vendored/efficientvit/models/nn/norm.py +164 -0
  57. segment_everything/vendored/efficientvit/models/nn/ops.py +597 -0
  58. segment_everything/vendored/efficientvit/models/utils/__init__.py +7 -0
  59. segment_everything/vendored/efficientvit/models/utils/list.py +53 -0
  60. segment_everything/vendored/efficientvit/models/utils/network.py +73 -0
  61. segment_everything/vendored/efficientvit/models/utils/random.py +65 -0
  62. segment_everything/vendored/efficientvit/sam_model_zoo.py +45 -0
  63. segment_everything/vendored/efficientvit/seg_model_zoo.py +70 -0
  64. segment_everything/vendored/get_object_aware.py +26 -0
  65. segment_everything/vendored/mobilesamv2/__init__.py +16 -0
  66. segment_everything/vendored/mobilesamv2/automatic_mask_generator.py +415 -0
  67. segment_everything/vendored/mobilesamv2/build_sam.py +246 -0
  68. segment_everything/vendored/mobilesamv2/modeling/__init__.py +11 -0
  69. segment_everything/vendored/mobilesamv2/modeling/common.py +43 -0
  70. segment_everything/vendored/mobilesamv2/modeling/image_encoder.py +394 -0
  71. segment_everything/vendored/mobilesamv2/modeling/mask_decoder.py +213 -0
  72. segment_everything/vendored/mobilesamv2/modeling/prompt_encoder.py +217 -0
  73. segment_everything/vendored/mobilesamv2/modeling/sam.py +203 -0
  74. segment_everything/vendored/mobilesamv2/modeling/transformer.py +240 -0
  75. segment_everything/vendored/mobilesamv2/predictor.py +384 -0
  76. segment_everything/vendored/mobilesamv2/utils/__init__.py +5 -0
  77. segment_everything/vendored/mobilesamv2/utils/amg.py +347 -0
  78. segment_everything/vendored/mobilesamv2/utils/onnx.py +144 -0
  79. segment_everything/vendored/mobilesamv2/utils/transforms.py +103 -0
  80. segment_everything/vendored/object_detection/__init__.py +0 -0
  81. segment_everything/vendored/object_detection/ultralytics/__init__.py +5 -0
  82. segment_everything/vendored/object_detection/ultralytics/nn/__init__.py +9 -0
  83. segment_everything/vendored/object_detection/ultralytics/nn/autobackend.py +658 -0
  84. segment_everything/vendored/object_detection/ultralytics/nn/autoshape.py +397 -0
  85. segment_everything/vendored/object_detection/ultralytics/nn/modules/__init__.py +110 -0
  86. segment_everything/vendored/object_detection/ultralytics/nn/modules/block.py +304 -0
  87. segment_everything/vendored/object_detection/ultralytics/nn/modules/conv.py +297 -0
  88. segment_everything/vendored/object_detection/ultralytics/nn/modules/head.py +468 -0
  89. segment_everything/vendored/object_detection/ultralytics/nn/modules/transformer.py +378 -0
  90. segment_everything/vendored/object_detection/ultralytics/nn/modules/utils.py +78 -0
  91. segment_everything/vendored/object_detection/ultralytics/nn/tasks.py +1049 -0
  92. segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/__init__.py +6 -0
  93. segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/model.py +104 -0
  94. segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/predict.py +95 -0
  95. segment_everything/vendored/object_detection/ultralytics/yolo/__init__.py +5 -0
  96. segment_everything/vendored/object_detection/ultralytics/yolo/cfg/__init__.py +588 -0
  97. segment_everything/vendored/object_detection/ultralytics/yolo/cfg/default.yaml +117 -0
  98. segment_everything/vendored/object_detection/ultralytics/yolo/data/__init__.py +9 -0
  99. segment_everything/vendored/object_detection/ultralytics/yolo/data/annotator.py +53 -0
  100. segment_everything/vendored/object_detection/ultralytics/yolo/data/augment.py +899 -0
  101. segment_everything/vendored/object_detection/ultralytics/yolo/data/base.py +286 -0
  102. segment_everything/vendored/object_detection/ultralytics/yolo/data/build.py +213 -0
  103. segment_everything/vendored/object_detection/ultralytics/yolo/data/converter.py +358 -0
  104. segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/__init__.py +0 -0
  105. segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/stream_loaders.py +459 -0
  106. segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset.py +274 -0
  107. segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset_wrappers.py +53 -0
  108. segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/download_weights.sh +18 -0
  109. segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco.sh +60 -0
  110. segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco128.sh +17 -0
  111. segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_imagenet.sh +51 -0
  112. segment_everything/vendored/object_detection/ultralytics/yolo/data/utils.py +716 -0
  113. segment_everything/vendored/object_detection/ultralytics/yolo/engine/__init__.py +0 -0
  114. segment_everything/vendored/object_detection/ultralytics/yolo/engine/exporter.py +1214 -0
  115. segment_everything/vendored/object_detection/ultralytics/yolo/engine/model.py +641 -0
  116. segment_everything/vendored/object_detection/ultralytics/yolo/engine/predictor.py +461 -0
  117. segment_everything/vendored/object_detection/ultralytics/yolo/engine/results.py +741 -0
  118. segment_everything/vendored/object_detection/ultralytics/yolo/utils/__init__.py +893 -0
  119. segment_everything/vendored/object_detection/ultralytics/yolo/utils/autobatch.py +108 -0
  120. segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/__init__.py +5 -0
  121. segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/base.py +212 -0
  122. segment_everything/vendored/object_detection/ultralytics/yolo/utils/checks.py +547 -0
  123. segment_everything/vendored/object_detection/ultralytics/yolo/utils/dist.py +67 -0
  124. segment_everything/vendored/object_detection/ultralytics/yolo/utils/downloads.py +353 -0
  125. segment_everything/vendored/object_detection/ultralytics/yolo/utils/errors.py +12 -0
  126. segment_everything/vendored/object_detection/ultralytics/yolo/utils/files.py +100 -0
  127. segment_everything/vendored/object_detection/ultralytics/yolo/utils/instance.py +391 -0
  128. segment_everything/vendored/object_detection/ultralytics/yolo/utils/loss.py +579 -0
  129. segment_everything/vendored/object_detection/ultralytics/yolo/utils/metrics.py +1189 -0
  130. segment_everything/vendored/object_detection/ultralytics/yolo/utils/ops.py +870 -0
  131. segment_everything/vendored/object_detection/ultralytics/yolo/utils/patches.py +45 -0
  132. segment_everything/vendored/object_detection/ultralytics/yolo/utils/plotting.py +767 -0
  133. segment_everything/vendored/object_detection/ultralytics/yolo/utils/tal.py +276 -0
  134. segment_everything/vendored/object_detection/ultralytics/yolo/utils/torch_utils.py +684 -0
  135. segment_everything/vendored/object_detection/ultralytics/yolo/utils/tuner.py +54 -0
  136. segment_everything/vendored/object_detection/ultralytics/yolo/v8/__init__.py +5 -0
  137. segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/__init__.py +5 -0
  138. segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/predict.py +69 -0
  139. segment_everything/vendored/tinyvit/__init__.py +2 -0
  140. segment_everything/vendored/tinyvit/tiny_vit.py +867 -0
  141. segment_everything/weights_helper.py +124 -0
  142. segment_everything-0.1.0.dist-info/METADATA +53 -0
  143. segment_everything-0.1.0.dist-info/RECORD +145 -0
  144. segment_everything-0.1.0.dist-info/WHEEL +4 -0
  145. segment_everything-0.1.0.dist-info/licenses/LICENSE +28 -0
@@ -0,0 +1,5 @@
1
+ """Utilities for rendering and training prompt-based and overlapping segmentation models (like SAM). """
2
+
3
+ __version__ = '0.1.0'
4
+ __author__ = "Brian Northan"
5
+ __email__ = "bnorthan@gmail.com"
@@ -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,5 @@
1
+ You may remove this file if you don't intend to add types to your package
2
+
3
+ Details at:
4
+
5
+ https://mypy.readthedocs.io/en/stable/installed_packages.html#creating-pep-561-compatible-packages
@@ -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