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,415 @@
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ import numpy as np
8
+ import torch
9
+ from torchvision.ops.boxes import batched_nms, box_area # type: ignore
10
+
11
+ from typing import Any, Dict, List, Optional, Tuple
12
+
13
+ from .modeling import Sam
14
+ from .predictor import SamPredictor
15
+ from .utils.amg import (
16
+ MaskData,
17
+ area_from_rle,
18
+ batch_iterator,
19
+ batched_mask_to_box,
20
+ box_xyxy_to_xywh,
21
+ build_all_layer_point_grids,
22
+ calculate_stability_score,
23
+ coco_encode_rle,
24
+ generate_crop_boxes,
25
+ is_box_near_crop_edge,
26
+ mask_to_rle_pytorch,
27
+ remove_small_regions,
28
+ rle_to_mask,
29
+ uncrop_boxes_xyxy,
30
+ uncrop_masks,
31
+ uncrop_points,
32
+ )
33
+
34
+
35
+ class SamAutomaticMaskGenerator:
36
+ def __init__(
37
+ self,
38
+ model: Sam,
39
+ points_per_side: Optional[int] = 32,#32
40
+ points_per_batch: int = 64,
41
+ pred_iou_thresh: float = 0.88,
42
+ stability_score_thresh: float = 0.95,
43
+ stability_score_offset: float = 1.0,
44
+ box_nms_thresh: float = 0.7,
45
+ crop_n_layers: int = 0,
46
+ crop_nms_thresh: float = 0.7,
47
+ crop_overlap_ratio: float = 512 / 1500,
48
+ crop_n_points_downscale_factor: int = 1,
49
+ point_grids: Optional[List[np.ndarray]] = None,
50
+ min_mask_region_area: int = 0,
51
+ output_mode: str = "binary_mask",
52
+ ) -> None:
53
+ """
54
+ Using a SAM model, generates masks for the entire image.
55
+ Generates a grid of point prompts over the image, then filters
56
+ low quality and duplicate masks. The default settings are chosen
57
+ for SAM with a ViT-H backbone.
58
+
59
+ Arguments:
60
+ model (Sam): The SAM model to use for mask prediction.
61
+ points_per_side (int or None): The number of points to be sampled
62
+ along one side of the image. The total number of points is
63
+ points_per_side**2. If None, 'point_grids' must provide explicit
64
+ point sampling.
65
+ points_per_batch (int): Sets the number of points run simultaneously
66
+ by the model. Higher numbers may be faster but use more GPU memory.
67
+ pred_iou_thresh (float): A filtering threshold in [0,1], using the
68
+ model's predicted mask quality.
69
+ stability_score_thresh (float): A filtering threshold in [0,1], using
70
+ the stability of the mask under changes to the cutoff used to binarize
71
+ the model's mask predictions.
72
+ stability_score_offset (float): The amount to shift the cutoff when
73
+ calculated the stability score.
74
+ box_nms_thresh (float): The box IoU cutoff used by non-maximal
75
+ suppression to filter duplicate masks.
76
+ crop_n_layers (int): If >0, mask prediction will be run again on
77
+ crops of the image. Sets the number of layers to run, where each
78
+ layer has 2**i_layer number of image crops.
79
+ crop_nms_thresh (float): The box IoU cutoff used by non-maximal
80
+ suppression to filter duplicate masks between different crops.
81
+ crop_overlap_ratio (float): Sets the degree to which crops overlap.
82
+ In the first crop layer, crops will overlap by this fraction of
83
+ the image length. Later layers with more crops scale down this overlap.
84
+ crop_n_points_downscale_factor (int): The number of points-per-side
85
+ sampled in layer n is scaled down by crop_n_points_downscale_factor**n.
86
+ point_grids (list(np.ndarray) or None): A list over explicit grids
87
+ of points used for sampling, normalized to [0,1]. The nth grid in the
88
+ list is used in the nth crop layer. Exclusive with points_per_side.
89
+ min_mask_region_area (int): If >0, postprocessing will be applied
90
+ to remove disconnected regions and holes in masks with area smaller
91
+ than min_mask_region_area. Requires opencv.
92
+ output_mode (str): The form masks are returned in. Can be 'binary_mask',
93
+ 'uncompressed_rle', or 'coco_rle'. 'coco_rle' requires pycocotools.
94
+ For large resolutions, 'binary_mask' may consume large amounts of
95
+ memory.
96
+ """
97
+
98
+ assert (points_per_side is None) != (
99
+ point_grids is None
100
+ ), "Exactly one of points_per_side or point_grid must be provided."
101
+ if points_per_side is not None:
102
+ #points position(0-1)
103
+ self.point_grids = build_all_layer_point_grids(
104
+ points_per_side,
105
+ crop_n_layers,
106
+ crop_n_points_downscale_factor,
107
+ )
108
+ #import pdb;pdb.set_trace()
109
+ elif point_grids is not None:
110
+ self.point_grids = point_grids
111
+ else:
112
+ raise ValueError("Can't have both points_per_side and point_grid be None.")
113
+ assert output_mode in [
114
+ "binary_mask",
115
+ "uncompressed_rle",
116
+ "coco_rle",
117
+ ], f"Unknown output_mode {output_mode}."
118
+ if output_mode == "coco_rle":
119
+ from pycocotools import mask as mask_utils # type: ignore # noqa: F401
120
+
121
+ if min_mask_region_area > 0:
122
+ import cv2 # type: ignore # noqa: F401
123
+
124
+ self.predictor = SamPredictor(model)
125
+ self.points_per_batch = points_per_batch
126
+ self.pred_iou_thresh = pred_iou_thresh
127
+ self.stability_score_thresh = stability_score_thresh
128
+ self.stability_score_offset = stability_score_offset
129
+ self.box_nms_thresh = box_nms_thresh
130
+ self.crop_n_layers = crop_n_layers
131
+ self.crop_nms_thresh = crop_nms_thresh
132
+ self.crop_overlap_ratio = crop_overlap_ratio
133
+ self.crop_n_points_downscale_factor = crop_n_points_downscale_factor
134
+ self.min_mask_region_area = min_mask_region_area
135
+ self.output_mode = output_mode
136
+
137
+ @torch.no_grad()
138
+ def generate(self, image: np.ndarray) -> List[Dict[str, Any]]:
139
+ """
140
+ Generates masks for the given image.
141
+
142
+ Arguments:
143
+ image (np.ndarray): The image to generate masks for, in HWC uint8 format.
144
+
145
+ Returns:
146
+ list(dict(str, any)): A list over records for masks. Each record is
147
+ a dict containing the following keys:
148
+ segmentation (dict(str, any) or np.ndarray): The mask. If
149
+ output_mode='binary_mask', is an array of shape HW. Otherwise,
150
+ is a dictionary containing the RLE.
151
+ bbox (list(float)): The box around the mask, in XYWH format.
152
+ area (int): The area in pixels of the mask.
153
+ predicted_iou (float): The model's own prediction of the mask's
154
+ quality. This is filtered by the pred_iou_thresh parameter.
155
+ point_coords (list(list(float))): The point coordinates input
156
+ to the model to generate this mask.
157
+ stability_score (float): A measure of the mask's quality. This
158
+ is filtered on using the stability_score_thresh parameter.
159
+ crop_box (list(float)): The crop of the image used to generate
160
+ the mask, given in XYWH format.
161
+ """
162
+
163
+ # Generate masks
164
+ mask_data = self._generate_masks(image)
165
+
166
+ # Filter small disconnected regions and holes in masks
167
+ if self.min_mask_region_area > 0:
168
+ mask_data = self.postprocess_small_regions(
169
+ mask_data,
170
+ self.min_mask_region_area,
171
+ max(self.box_nms_thresh, self.crop_nms_thresh),
172
+ )
173
+
174
+ # Encode masks
175
+ #import pdb;pdb.set_trace()
176
+ if self.output_mode == "coco_rle":
177
+ mask_data["segmentations"] = [coco_encode_rle(rle) for rle in mask_data["rles"]]
178
+ elif self.output_mode == "binary_mask":
179
+ mask_data["segmentations"] = [rle_to_mask(rle) for rle in mask_data["rles"]]
180
+ else:
181
+ mask_data["segmentations"] = mask_data["rles"]
182
+ # Write mask records
183
+ curr_anns = []
184
+ for idx in range(len(mask_data["segmentations"])):
185
+ ann = {
186
+ "segmentation": mask_data["segmentations"][idx],
187
+ "area": area_from_rle(mask_data["rles"][idx]),
188
+ "bbox": box_xyxy_to_xywh(mask_data["boxes"][idx]).tolist(),
189
+ "predicted_iou": mask_data["iou_preds"][idx].item(),
190
+ "point_coords": [mask_data["points"][idx].tolist()],
191
+ "stability_score": mask_data["stability_score"][idx].item(),
192
+ "crop_box": box_xyxy_to_xywh(mask_data["crop_boxes"][idx]).tolist(),
193
+ }
194
+ curr_anns.append(ann)
195
+ return curr_anns
196
+
197
+ def _generate_masks(self, image: np.ndarray) -> MaskData:
198
+ orig_size = image.shape[:2]
199
+ #import pdb;pdb.set_trace()
200
+ crop_boxes, layer_idxs = generate_crop_boxes(
201
+ orig_size, self.crop_n_layers, self.crop_overlap_ratio
202
+ )
203
+
204
+ # Iterate over image crops
205
+ data = MaskData()
206
+ #import pdb;pdb.set_trace()
207
+ for crop_box, layer_idx in zip(crop_boxes, layer_idxs):
208
+ crop_data = self._process_crop(image, crop_box, layer_idx, orig_size)
209
+ data.cat(crop_data)
210
+
211
+ # Remove duplicate masks between crops
212
+ if len(crop_boxes) > 1:
213
+ # Prefer masks from smaller crops
214
+ scores = 1 / box_area(data["crop_boxes"])
215
+ scores = scores.to(data["boxes"].device)
216
+ keep_by_nms = batched_nms(
217
+ data["boxes"].float(),
218
+ scores,
219
+ torch.zeros_like(data["boxes"][:, 0]), # categories
220
+ iou_threshold=self.crop_nms_thresh,
221
+ )
222
+ data.filter(keep_by_nms)
223
+
224
+ data.to_numpy()
225
+ return data
226
+
227
+ def _process_crop(
228
+ self,
229
+ image: np.ndarray,
230
+ crop_box: List[int],
231
+ crop_layer_idx: int,
232
+ orig_size: Tuple[int, ...],
233
+ ) -> MaskData:
234
+ # Crop the image and calculate embeddings
235
+ x0, y0, x1, y1 = crop_box
236
+ cropped_im = image[y0:y1, x0:x1, :]
237
+ cropped_im_size = cropped_im.shape[:2]
238
+ self.predictor.set_image(cropped_im)
239
+
240
+ # Get points for this crop
241
+ points_scale = np.array(cropped_im_size)[None, ::-1]
242
+ points_for_image = self.point_grids[crop_layer_idx] * points_scale
243
+
244
+ # Generate masks for this crop in batches
245
+ data = MaskData()
246
+ #import pdb;pdb.set_trace()
247
+
248
+ for (points,) in batch_iterator(self.points_per_batch, points_for_image):
249
+ #import pdb;pdb.set_trace()
250
+ batch_data = self._process_batch(points, cropped_im_size, crop_box, orig_size)
251
+ data.cat(batch_data)
252
+ del batch_data
253
+ self.predictor.reset_image()
254
+ # Remove duplicates within this crop.
255
+
256
+ keep_by_nms = batched_nms(
257
+ data["boxes"].float(),
258
+ data["iou_preds"],
259
+ torch.zeros_like(data["boxes"][:, 0]), # categories
260
+ iou_threshold=self.box_nms_thresh,
261
+ )
262
+
263
+ data.filter(keep_by_nms)
264
+
265
+ ###########################################################################
266
+ # cc = time.time()
267
+ # for (points,) in batch_iterator(self.points_per_batch, points_for_image):
268
+ # batch_data = self._process_batch(points, cropped_im_size, crop_box, orig_size)
269
+ # data.cat(batch_data)
270
+ # del batch_data
271
+ # self.predictor.reset_image()
272
+
273
+ # # Remove duplicates within this crop.
274
+ # keep_by_nms = batched_nms(
275
+ # data["boxes"].float(),
276
+ # data["iou_preds"],
277
+ # torch.zeros_like(data["boxes"][:, 0]), # categories
278
+ # iou_threshold=self.box_nms_thresh,
279
+ # )
280
+ # data.filter(keep_by_nms)
281
+ # dd = time.time(); print('cv2 read:', cc-dd)
282
+ # ###############################################################################
283
+ # import time
284
+ # ee = time.time()
285
+ # for (points,) in batch_iterator(self.points_per_batch, points_for_image):
286
+ # batch_data = self._process_batch(points, cropped_im_size, crop_box, orig_size)
287
+ # data.cat(batch_data)
288
+ # del batch_data
289
+ # self.predictor.reset_image()
290
+
291
+ # # Remove duplicates within this crop.
292
+ # keep_by_nms = batched_nms(
293
+ # data["boxes"].float(),
294
+ # data["iou_preds"],
295
+ # torch.zeros_like(data["boxes"][:, 0]), # categories
296
+ # iou_threshold=self.box_nms_thresh,
297
+ # )
298
+ # data.filter(keep_by_nms)
299
+ # ff = time.time(); print('cv2 read:', ff-ee)
300
+ #import pdb;pdb.set_trace()
301
+
302
+
303
+ # Return to the original image frame
304
+
305
+ data["boxes"] = uncrop_boxes_xyxy(data["boxes"], crop_box)
306
+ data["points"] = uncrop_points(data["points"], crop_box)
307
+ data["crop_boxes"] = torch.tensor([crop_box for _ in range(len(data["rles"]))])
308
+
309
+ return data
310
+
311
+ def _process_batch(
312
+ self,
313
+ points: np.ndarray,
314
+ im_size: Tuple[int, ...],
315
+ crop_box: List[int],
316
+ orig_size: Tuple[int, ...],
317
+ ) -> MaskData:
318
+ orig_h, orig_w = orig_size
319
+
320
+ # Run model on this batch
321
+ #import pdb;pdb.set_trace()
322
+ transformed_points = self.predictor.transform.apply_coords(points, im_size)
323
+ in_points = torch.as_tensor(transformed_points, device=self.predictor.device)
324
+ in_labels = torch.ones(in_points.shape[0], dtype=torch.int, device=in_points.device)
325
+ #import pdb;pdb.set_trace()
326
+ masks, iou_preds, _ = self.predictor.predict_torch(
327
+ in_points[:, None, :],
328
+ in_labels[:, None],
329
+ multimask_output=True,
330
+ return_logits=True,
331
+ )
332
+ # Serialize predictions and store in MaskData
333
+ data = MaskData(
334
+ masks=masks.flatten(0, 1),
335
+ iou_preds=iou_preds.flatten(0, 1),
336
+ points=torch.as_tensor(points.repeat(masks.shape[1], axis=0)),
337
+ )
338
+ del masks
339
+
340
+ # Filter by predicted IoU
341
+ if self.pred_iou_thresh > 0.0:
342
+ keep_mask = data["iou_preds"] > self.pred_iou_thresh
343
+ data.filter(keep_mask)
344
+ # Calculate stability score
345
+ #import pdb;pdb.set_trace()
346
+ data["stability_score"] = calculate_stability_score(
347
+ data["masks"], self.predictor.model.mask_threshold, self.stability_score_offset
348
+ )
349
+ if self.stability_score_thresh > 0.0:
350
+ keep_mask = data["stability_score"] >= self.stability_score_thresh
351
+ data.filter(keep_mask)
352
+ # Threshold masks and calculate boxes
353
+ data["masks"] = data["masks"] > self.predictor.model.mask_threshold
354
+ data["boxes"] = batched_mask_to_box(data["masks"])
355
+ # Filter boxes that touch crop boundaries
356
+ keep_mask = ~is_box_near_crop_edge(data["boxes"], crop_box, [0, 0, orig_w, orig_h])
357
+ if not torch.all(keep_mask):
358
+ data.filter(keep_mask)
359
+ # Compress to RLE
360
+ data["masks"] = uncrop_masks(data["masks"], crop_box, orig_h, orig_w)
361
+ data["rles"] = mask_to_rle_pytorch(data["masks"])
362
+ del data["masks"]
363
+
364
+ return data
365
+
366
+ @staticmethod
367
+ def postprocess_small_regions(
368
+ mask_data: MaskData, min_area: int, nms_thresh: float
369
+ ) -> MaskData:
370
+ """
371
+ Removes small disconnected regions and holes in masks, then reruns
372
+ box NMS to remove any new duplicates.
373
+
374
+ Edits mask_data in place.
375
+
376
+ Requires open-cv as a dependency.
377
+ """
378
+ if len(mask_data["rles"]) == 0:
379
+ return mask_data
380
+
381
+ # Filter small disconnected regions and holes
382
+ new_masks = []
383
+ scores = []
384
+ for rle in mask_data["rles"]:
385
+ mask = rle_to_mask(rle)
386
+
387
+ mask, changed = remove_small_regions(mask, min_area, mode="holes")
388
+ unchanged = not changed
389
+ mask, changed = remove_small_regions(mask, min_area, mode="islands")
390
+ unchanged = unchanged and not changed
391
+
392
+ new_masks.append(torch.as_tensor(mask).unsqueeze(0))
393
+ # Give score=0 to changed masks and score=1 to unchanged masks
394
+ # so NMS will prefer ones that didn't need postprocessing
395
+ scores.append(float(unchanged))
396
+
397
+ # Recalculate boxes and remove any new duplicates
398
+ masks = torch.cat(new_masks, dim=0)
399
+ boxes = batched_mask_to_box(masks)
400
+ keep_by_nms = batched_nms(
401
+ boxes.float(),
402
+ torch.as_tensor(scores),
403
+ torch.zeros_like(boxes[:, 0]), # categories
404
+ iou_threshold=nms_thresh,
405
+ )
406
+
407
+ # Only recalculate RLEs for masks that have changed
408
+ for i_mask in keep_by_nms:
409
+ if scores[i_mask] == 0.0:
410
+ mask_torch = masks[i_mask].unsqueeze(0)
411
+ mask_data["rles"][i_mask] = mask_to_rle_pytorch(mask_torch)[0]
412
+ mask_data["boxes"][i_mask] = boxes[i_mask] # update res directly
413
+ mask_data.filter(keep_by_nms)
414
+
415
+ return mask_data
@@ -0,0 +1,246 @@
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+
10
+ from functools import partial
11
+
12
+ from .modeling import (
13
+ ImageEncoderViT,
14
+ MaskDecoder,
15
+ PromptEncoder,
16
+ Sam,
17
+ TwoWayTransformer,
18
+ )
19
+ from ..tinyvit.tiny_vit import TinyViT # 11000
20
+ from ..efficientvit.models.efficientvit.backbone import (
21
+ EfficientViTLargeBackbone,
22
+ )
23
+ from ..efficientvit.models.efficientvit.sam import (
24
+ SamNeck,
25
+ EfficientViTSamImageEncoder,
26
+ )
27
+ from ..efficientvit.models.nn.norm import set_norm_eps
28
+
29
+
30
+ def build_sam_vit_h(checkpoint=None):
31
+ return _build_sam(
32
+ encoder_embed_dim=1280,
33
+ encoder_depth=32,
34
+ encoder_num_heads=16,
35
+ encoder_global_attn_indexes=[7, 15, 23, 31],
36
+ checkpoint=checkpoint,
37
+ )
38
+
39
+
40
+ def build_sam_vit_l(checkpoint=None):
41
+ return _build_sam(
42
+ encoder_embed_dim=1024,
43
+ encoder_depth=24,
44
+ encoder_num_heads=16,
45
+ encoder_global_attn_indexes=[5, 11, 17, 23],
46
+ checkpoint=checkpoint,
47
+ )
48
+
49
+
50
+ def build_sam_vit_b(checkpoint=None):
51
+ return _build_sam(
52
+ encoder_embed_dim=768,
53
+ encoder_depth=12,
54
+ encoder_num_heads=12,
55
+ encoder_global_attn_indexes=[2, 5, 8, 11],
56
+ checkpoint=checkpoint,
57
+ )
58
+
59
+
60
+ def _build_sam(
61
+ encoder_embed_dim,
62
+ encoder_depth,
63
+ encoder_num_heads,
64
+ encoder_global_attn_indexes,
65
+ checkpoint=None,
66
+ ):
67
+ prompt_embed_dim = 256
68
+ image_size = 1024
69
+ vit_patch_size = 16
70
+ image_embedding_size = image_size // vit_patch_size
71
+ sam = Sam(
72
+ image_encoder=ImageEncoderViT(
73
+ depth=encoder_depth,
74
+ embed_dim=encoder_embed_dim,
75
+ img_size=image_size,
76
+ mlp_ratio=4,
77
+ norm_layer=partial(torch.nn.LayerNorm, eps=1e-6),
78
+ num_heads=encoder_num_heads,
79
+ patch_size=vit_patch_size,
80
+ qkv_bias=True,
81
+ use_rel_pos=True,
82
+ global_attn_indexes=encoder_global_attn_indexes,
83
+ window_size=14,
84
+ out_chans=prompt_embed_dim,
85
+ ),
86
+ prompt_encoder=PromptEncoder(
87
+ embed_dim=prompt_embed_dim,
88
+ image_embedding_size=(image_embedding_size, image_embedding_size),
89
+ input_image_size=(image_size, image_size),
90
+ mask_in_chans=16,
91
+ ),
92
+ mask_decoder=MaskDecoder(
93
+ num_multimask_outputs=3,
94
+ transformer=TwoWayTransformer(
95
+ depth=2,
96
+ embedding_dim=prompt_embed_dim,
97
+ mlp_dim=2048,
98
+ num_heads=8,
99
+ ),
100
+ transformer_dim=prompt_embed_dim,
101
+ iou_head_depth=3,
102
+ iou_head_hidden_dim=256,
103
+ ),
104
+ pixel_mean=[123.675, 116.28, 103.53],
105
+ pixel_std=[58.395, 57.12, 57.375],
106
+ )
107
+ sam.eval()
108
+ if checkpoint is not None:
109
+ with open(checkpoint, "rb") as f:
110
+ state_dict = torch.load(f)
111
+ sam.load_state_dict(state_dict, strict=False)
112
+ return sam
113
+
114
+
115
+ def build_sam_vit_t_encoder(checkpoint=None):
116
+ mobile_sam = TinyViT(
117
+ img_size=1024,
118
+ in_chans=3,
119
+ num_classes=1000,
120
+ embed_dims=[64, 128, 160, 320],
121
+ depths=[2, 2, 6, 2],
122
+ num_heads=[2, 4, 5, 10],
123
+ window_sizes=[7, 7, 14, 7],
124
+ mlp_ratio=4.0,
125
+ drop_rate=0.0,
126
+ drop_path_rate=0.0,
127
+ use_checkpoint=False,
128
+ mbconv_expand_ratio=4.0,
129
+ local_conv_size=3,
130
+ layer_lr_decay=0.8,
131
+ )
132
+ if checkpoint is not None:
133
+ with open(checkpoint, "rb") as f:
134
+ state_dict = torch.load(f)
135
+ mobile_sam.load_state_dict(state_dict["model"], strict=False)
136
+ return mobile_sam
137
+
138
+
139
+ def build_efficientvit_l2_encoder(checkpoint=None):
140
+ backbone = EfficientViTLargeBackbone(
141
+ width_list=[32, 64, 128, 256, 512],
142
+ depth_list=[1, 2, 2, 8, 8],
143
+ in_channels=3,
144
+ qkv_dim=32,
145
+ norm="bn2d",
146
+ act_func="gelu",
147
+ )
148
+ neck = SamNeck(
149
+ fid_list=["stage4", "stage3", "stage2"],
150
+ in_channel_list=[512, 256, 128],
151
+ head_width=256,
152
+ head_depth=12,
153
+ expand_ratio=1,
154
+ middle_op="fmbconv",
155
+ out_dim=256,
156
+ )
157
+ image_encoder = EfficientViTSamImageEncoder(backbone, neck)
158
+ set_norm_eps(image_encoder, 1e-6)
159
+ checkpoints = torch.load(checkpoint)
160
+ checkpoint = checkpoints["state_dict"]
161
+ new_state_dict = {}
162
+ if checkpoint != None:
163
+ for key, value in checkpoint.items():
164
+ index = key.find("image_encoder.")
165
+ if index != -1:
166
+ new_key = key[index + len("image_encoder.") :]
167
+ new_state_dict[new_key] = value
168
+ else:
169
+ continue
170
+ image_encoder.load_state_dict(new_state_dict, strict=True) # origin
171
+ print("VIT checkpoint loaded successfully")
172
+ return image_encoder
173
+
174
+
175
+ def build_sam_vit_h_encoder(checkpoint=None):
176
+ prompt_embed_dim = 256
177
+ image_size = 1024
178
+ vit_patch_size = 16
179
+ encoder_embed_dim = 1280
180
+ encoder_depth = 32
181
+ encoder_num_heads = 16
182
+ encoder_global_attn_indexes = [7, 15, 23, 31]
183
+ image_encoder = ImageEncoderViT(
184
+ depth=encoder_depth,
185
+ embed_dim=encoder_embed_dim,
186
+ img_size=image_size,
187
+ mlp_ratio=4,
188
+ norm_layer=partial(torch.nn.LayerNorm, eps=1e-6),
189
+ num_heads=encoder_num_heads,
190
+ patch_size=vit_patch_size,
191
+ qkv_bias=True,
192
+ use_rel_pos=True,
193
+ global_attn_indexes=encoder_global_attn_indexes,
194
+ window_size=14,
195
+ out_chans=prompt_embed_dim,
196
+ )
197
+ if checkpoint is not None:
198
+ with open(checkpoint, "rb") as f:
199
+ state_dict = torch.load(f)
200
+ image_encoder.load_state_dict(state_dict, strict=True)
201
+ return image_encoder
202
+
203
+
204
+ def build_PromptGuidedDecoder(checkpoint=None):
205
+ prompt_embed_dim = 256
206
+ image_size = 1024
207
+ vit_patch_size = 16
208
+ image_embedding_size = image_size // vit_patch_size
209
+ prompt_encoder = PromptEncoder(
210
+ embed_dim=prompt_embed_dim,
211
+ image_embedding_size=(image_embedding_size, image_embedding_size),
212
+ input_image_size=(image_size, image_size),
213
+ mask_in_chans=16,
214
+ )
215
+ mask_decoder = MaskDecoder(
216
+ num_multimask_outputs=3,
217
+ transformer=TwoWayTransformer(
218
+ depth=2,
219
+ embedding_dim=prompt_embed_dim,
220
+ mlp_dim=2048,
221
+ num_heads=8,
222
+ ),
223
+ transformer_dim=prompt_embed_dim,
224
+ iou_head_depth=3,
225
+ iou_head_hidden_dim=256,
226
+ )
227
+ if checkpoint is not None:
228
+ with open(checkpoint, "rb") as f:
229
+ state_dict = torch.load(f)
230
+ promt_dict = state_dict["PromtEncoder"]
231
+ mask_dict = state_dict["MaskDecoder"]
232
+ prompt_encoder.load_state_dict(promt_dict)
233
+ mask_decoder.load_state_dict(mask_dict)
234
+ return {"PromtEncoder": prompt_encoder, "MaskDecoder": mask_decoder}
235
+
236
+
237
+ sam_model_registry = {
238
+ "default": build_sam_vit_h,
239
+ "vit_h": build_sam_vit_h,
240
+ "vit_l": build_sam_vit_l,
241
+ "vit_b": build_sam_vit_b,
242
+ "tiny_vit": build_sam_vit_t_encoder,
243
+ "efficientvit_l2": build_efficientvit_l2_encoder,
244
+ "PromptGuidedDecoder": build_PromptGuidedDecoder,
245
+ "sam_vit_h": build_sam_vit_h_encoder,
246
+ }
@@ -0,0 +1,11 @@
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ from .sam import Sam
8
+ from .image_encoder import ImageEncoderViT
9
+ from .mask_decoder import MaskDecoder
10
+ from .prompt_encoder import PromptEncoder
11
+ from .transformer import TwoWayTransformer