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
|
@@ -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
|