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,276 @@
1
+ # Ultralytics YOLO 🚀, AGPL-3.0 license
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+
6
+ from .checks import check_version
7
+ from .metrics import bbox_iou
8
+
9
+ TORCH_1_10 = check_version(torch.__version__, '1.10.0')
10
+
11
+
12
+ def select_candidates_in_gts(xy_centers, gt_bboxes, eps=1e-9):
13
+ """select the positive anchor center in gt
14
+
15
+ Args:
16
+ xy_centers (Tensor): shape(h*w, 4)
17
+ gt_bboxes (Tensor): shape(b, n_boxes, 4)
18
+ Return:
19
+ (Tensor): shape(b, n_boxes, h*w)
20
+ """
21
+ n_anchors = xy_centers.shape[0]
22
+ bs, n_boxes, _ = gt_bboxes.shape
23
+ lt, rb = gt_bboxes.view(-1, 1, 4).chunk(2, 2) # left-top, right-bottom
24
+ bbox_deltas = torch.cat((xy_centers[None] - lt, rb - xy_centers[None]), dim=2).view(bs, n_boxes, n_anchors, -1)
25
+ # return (bbox_deltas.min(3)[0] > eps).to(gt_bboxes.dtype)
26
+ return bbox_deltas.amin(3).gt_(eps)
27
+
28
+
29
+ def select_highest_overlaps(mask_pos, overlaps, n_max_boxes):
30
+ """if an anchor box is assigned to multiple gts,
31
+ the one with the highest iou will be selected.
32
+
33
+ Args:
34
+ mask_pos (Tensor): shape(b, n_max_boxes, h*w)
35
+ overlaps (Tensor): shape(b, n_max_boxes, h*w)
36
+ Return:
37
+ target_gt_idx (Tensor): shape(b, h*w)
38
+ fg_mask (Tensor): shape(b, h*w)
39
+ mask_pos (Tensor): shape(b, n_max_boxes, h*w)
40
+ """
41
+ # (b, n_max_boxes, h*w) -> (b, h*w)
42
+ fg_mask = mask_pos.sum(-2)
43
+ if fg_mask.max() > 1: # one anchor is assigned to multiple gt_bboxes
44
+ mask_multi_gts = (fg_mask.unsqueeze(1) > 1).expand(-1, n_max_boxes, -1) # (b, n_max_boxes, h*w)
45
+ max_overlaps_idx = overlaps.argmax(1) # (b, h*w)
46
+
47
+ is_max_overlaps = torch.zeros(mask_pos.shape, dtype=mask_pos.dtype, device=mask_pos.device)
48
+ is_max_overlaps.scatter_(1, max_overlaps_idx.unsqueeze(1), 1)
49
+
50
+ mask_pos = torch.where(mask_multi_gts, is_max_overlaps, mask_pos).float() # (b, n_max_boxes, h*w)
51
+ fg_mask = mask_pos.sum(-2)
52
+ # Find each grid serve which gt(index)
53
+ target_gt_idx = mask_pos.argmax(-2) # (b, h*w)
54
+ return target_gt_idx, fg_mask, mask_pos
55
+
56
+
57
+ class TaskAlignedAssigner(nn.Module):
58
+ """
59
+ A task-aligned assigner for object detection.
60
+
61
+ This class assigns ground-truth (gt) objects to anchors based on the task-aligned metric,
62
+ which combines both classification and localization information.
63
+
64
+ Attributes:
65
+ topk (int): The number of top candidates to consider.
66
+ num_classes (int): The number of object classes.
67
+ alpha (float): The alpha parameter for the classification component of the task-aligned metric.
68
+ beta (float): The beta parameter for the localization component of the task-aligned metric.
69
+ eps (float): A small value to prevent division by zero.
70
+ """
71
+
72
+ def __init__(self, topk=13, num_classes=80, alpha=1.0, beta=6.0, eps=1e-9):
73
+ """Initialize a TaskAlignedAssigner object with customizable hyperparameters."""
74
+ super().__init__()
75
+ self.topk = topk
76
+ self.num_classes = num_classes
77
+ self.bg_idx = num_classes
78
+ self.alpha = alpha
79
+ self.beta = beta
80
+ self.eps = eps
81
+
82
+ @torch.no_grad()
83
+ def forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt):
84
+ """
85
+ Compute the task-aligned assignment.
86
+ Reference https://github.com/Nioolek/PPYOLOE_pytorch/blob/master/ppyoloe/assigner/tal_assigner.py
87
+
88
+ Args:
89
+ pd_scores (Tensor): shape(bs, num_total_anchors, num_classes)
90
+ pd_bboxes (Tensor): shape(bs, num_total_anchors, 4)
91
+ anc_points (Tensor): shape(num_total_anchors, 2)
92
+ gt_labels (Tensor): shape(bs, n_max_boxes, 1)
93
+ gt_bboxes (Tensor): shape(bs, n_max_boxes, 4)
94
+ mask_gt (Tensor): shape(bs, n_max_boxes, 1)
95
+
96
+ Returns:
97
+ target_labels (Tensor): shape(bs, num_total_anchors)
98
+ target_bboxes (Tensor): shape(bs, num_total_anchors, 4)
99
+ target_scores (Tensor): shape(bs, num_total_anchors, num_classes)
100
+ fg_mask (Tensor): shape(bs, num_total_anchors)
101
+ target_gt_idx (Tensor): shape(bs, num_total_anchors)
102
+ """
103
+ self.bs = pd_scores.size(0)
104
+ self.n_max_boxes = gt_bboxes.size(1)
105
+
106
+ if self.n_max_boxes == 0:
107
+ device = gt_bboxes.device
108
+ return (torch.full_like(pd_scores[..., 0], self.bg_idx).to(device), torch.zeros_like(pd_bboxes).to(device),
109
+ torch.zeros_like(pd_scores).to(device), torch.zeros_like(pd_scores[..., 0]).to(device),
110
+ torch.zeros_like(pd_scores[..., 0]).to(device))
111
+
112
+ mask_pos, align_metric, overlaps = self.get_pos_mask(pd_scores, pd_bboxes, gt_labels, gt_bboxes, anc_points,
113
+ mask_gt)
114
+
115
+ target_gt_idx, fg_mask, mask_pos = select_highest_overlaps(mask_pos, overlaps, self.n_max_boxes)
116
+
117
+ # Assigned target
118
+ target_labels, target_bboxes, target_scores = self.get_targets(gt_labels, gt_bboxes, target_gt_idx, fg_mask)
119
+
120
+ # Normalize
121
+ align_metric *= mask_pos
122
+ pos_align_metrics = align_metric.amax(axis=-1, keepdim=True) # b, max_num_obj
123
+ pos_overlaps = (overlaps * mask_pos).amax(axis=-1, keepdim=True) # b, max_num_obj
124
+ norm_align_metric = (align_metric * pos_overlaps / (pos_align_metrics + self.eps)).amax(-2).unsqueeze(-1)
125
+ target_scores = target_scores * norm_align_metric
126
+
127
+ return target_labels, target_bboxes, target_scores, fg_mask.bool(), target_gt_idx
128
+
129
+ def get_pos_mask(self, pd_scores, pd_bboxes, gt_labels, gt_bboxes, anc_points, mask_gt):
130
+ """Get in_gts mask, (b, max_num_obj, h*w)."""
131
+ mask_in_gts = select_candidates_in_gts(anc_points, gt_bboxes)
132
+ # Get anchor_align metric, (b, max_num_obj, h*w)
133
+ align_metric, overlaps = self.get_box_metrics(pd_scores, pd_bboxes, gt_labels, gt_bboxes, mask_in_gts * mask_gt)
134
+ # Get topk_metric mask, (b, max_num_obj, h*w)
135
+ mask_topk = self.select_topk_candidates(align_metric, topk_mask=mask_gt.expand(-1, -1, self.topk).bool())
136
+ # Merge all mask to a final mask, (b, max_num_obj, h*w)
137
+ mask_pos = mask_topk * mask_in_gts * mask_gt
138
+
139
+ return mask_pos, align_metric, overlaps
140
+
141
+ def get_box_metrics(self, pd_scores, pd_bboxes, gt_labels, gt_bboxes, mask_gt):
142
+ """Compute alignment metric given predicted and ground truth bounding boxes."""
143
+ na = pd_bboxes.shape[-2]
144
+ mask_gt = mask_gt.bool() # b, max_num_obj, h*w
145
+ overlaps = torch.zeros([self.bs, self.n_max_boxes, na], dtype=pd_bboxes.dtype, device=pd_bboxes.device)
146
+ bbox_scores = torch.zeros([self.bs, self.n_max_boxes, na], dtype=pd_scores.dtype, device=pd_scores.device)
147
+
148
+ ind = torch.zeros([2, self.bs, self.n_max_boxes], dtype=torch.long) # 2, b, max_num_obj
149
+ ind[0] = torch.arange(end=self.bs).view(-1, 1).expand(-1, self.n_max_boxes) # b, max_num_obj
150
+ ind[1] = gt_labels.squeeze(-1) # b, max_num_obj
151
+ # Get the scores of each grid for each gt cls
152
+ bbox_scores[mask_gt] = pd_scores[ind[0], :, ind[1]][mask_gt] # b, max_num_obj, h*w
153
+
154
+ # (b, max_num_obj, 1, 4), (b, 1, h*w, 4)
155
+ pd_boxes = pd_bboxes.unsqueeze(1).expand(-1, self.n_max_boxes, -1, -1)[mask_gt]
156
+ gt_boxes = gt_bboxes.unsqueeze(2).expand(-1, -1, na, -1)[mask_gt]
157
+ overlaps[mask_gt] = bbox_iou(gt_boxes, pd_boxes, xywh=False, CIoU=True).squeeze(-1).clamp_(0)
158
+
159
+ align_metric = bbox_scores.pow(self.alpha) * overlaps.pow(self.beta)
160
+ return align_metric, overlaps
161
+
162
+ def select_topk_candidates(self, metrics, largest=True, topk_mask=None):
163
+ """
164
+ Select the top-k candidates based on the given metrics.
165
+
166
+ Args:
167
+ metrics (Tensor): A tensor of shape (b, max_num_obj, h*w), where b is the batch size,
168
+ max_num_obj is the maximum number of objects, and h*w represents the
169
+ total number of anchor points.
170
+ largest (bool): If True, select the largest values; otherwise, select the smallest values.
171
+ topk_mask (Tensor): An optional boolean tensor of shape (b, max_num_obj, topk), where
172
+ topk is the number of top candidates to consider. If not provided,
173
+ the top-k values are automatically computed based on the given metrics.
174
+
175
+ Returns:
176
+ (Tensor): A tensor of shape (b, max_num_obj, h*w) containing the selected top-k candidates.
177
+ """
178
+
179
+ # (b, max_num_obj, topk)
180
+ topk_metrics, topk_idxs = torch.topk(metrics, self.topk, dim=-1, largest=largest)
181
+ if topk_mask is None:
182
+ topk_mask = (topk_metrics.max(-1, keepdim=True)[0] > self.eps).expand_as(topk_idxs)
183
+ # (b, max_num_obj, topk)
184
+ topk_idxs.masked_fill_(~topk_mask, 0)
185
+
186
+ # (b, max_num_obj, topk, h*w) -> (b, max_num_obj, h*w)
187
+ count_tensor = torch.zeros(metrics.shape, dtype=torch.int8, device=topk_idxs.device)
188
+ ones = torch.ones_like(topk_idxs[:, :, :1], dtype=torch.int8, device=topk_idxs.device)
189
+ for k in range(self.topk):
190
+ # Expand topk_idxs for each value of k and add 1 at the specified positions
191
+ count_tensor.scatter_add_(-1, topk_idxs[:, :, k:k + 1], ones)
192
+ # count_tensor.scatter_add_(-1, topk_idxs, torch.ones_like(topk_idxs, dtype=torch.int8, device=topk_idxs.device))
193
+ # filter invalid bboxes
194
+ count_tensor.masked_fill_(count_tensor > 1, 0)
195
+
196
+ return count_tensor.to(metrics.dtype)
197
+
198
+ def get_targets(self, gt_labels, gt_bboxes, target_gt_idx, fg_mask):
199
+ """
200
+ Compute target labels, target bounding boxes, and target scores for the positive anchor points.
201
+
202
+ Args:
203
+ gt_labels (Tensor): Ground truth labels of shape (b, max_num_obj, 1), where b is the
204
+ batch size and max_num_obj is the maximum number of objects.
205
+ gt_bboxes (Tensor): Ground truth bounding boxes of shape (b, max_num_obj, 4).
206
+ target_gt_idx (Tensor): Indices of the assigned ground truth objects for positive
207
+ anchor points, with shape (b, h*w), where h*w is the total
208
+ number of anchor points.
209
+ fg_mask (Tensor): A boolean tensor of shape (b, h*w) indicating the positive
210
+ (foreground) anchor points.
211
+
212
+ Returns:
213
+ (Tuple[Tensor, Tensor, Tensor]): A tuple containing the following tensors:
214
+ - target_labels (Tensor): Shape (b, h*w), containing the target labels for
215
+ positive anchor points.
216
+ - target_bboxes (Tensor): Shape (b, h*w, 4), containing the target bounding boxes
217
+ for positive anchor points.
218
+ - target_scores (Tensor): Shape (b, h*w, num_classes), containing the target scores
219
+ for positive anchor points, where num_classes is the number
220
+ of object classes.
221
+ """
222
+
223
+ # Assigned target labels, (b, 1)
224
+ batch_ind = torch.arange(end=self.bs, dtype=torch.int64, device=gt_labels.device)[..., None]
225
+ target_gt_idx = target_gt_idx + batch_ind * self.n_max_boxes # (b, h*w)
226
+ target_labels = gt_labels.long().flatten()[target_gt_idx] # (b, h*w)
227
+
228
+ # Assigned target boxes, (b, max_num_obj, 4) -> (b, h*w)
229
+ target_bboxes = gt_bboxes.view(-1, 4)[target_gt_idx]
230
+
231
+ # Assigned target scores
232
+ target_labels.clamp_(0)
233
+
234
+ # 10x faster than F.one_hot()
235
+ target_scores = torch.zeros((target_labels.shape[0], target_labels.shape[1], self.num_classes),
236
+ dtype=torch.int64,
237
+ device=target_labels.device) # (b, h*w, 80)
238
+ target_scores.scatter_(2, target_labels.unsqueeze(-1), 1)
239
+
240
+ fg_scores_mask = fg_mask[:, :, None].repeat(1, 1, self.num_classes) # (b, h*w, 80)
241
+ target_scores = torch.where(fg_scores_mask > 0, target_scores, 0)
242
+
243
+ return target_labels, target_bboxes, target_scores
244
+
245
+
246
+ def make_anchors(feats, strides, grid_cell_offset=0.5):
247
+ """Generate anchors from features."""
248
+ anchor_points, stride_tensor = [], []
249
+ assert feats is not None
250
+ dtype, device = feats[0].dtype, feats[0].device
251
+ for i, stride in enumerate(strides):
252
+ _, _, h, w = feats[i].shape
253
+ sx = torch.arange(end=w, device=device, dtype=dtype) + grid_cell_offset # shift x
254
+ sy = torch.arange(end=h, device=device, dtype=dtype) + grid_cell_offset # shift y
255
+ sy, sx = torch.meshgrid(sy, sx, indexing='ij') if TORCH_1_10 else torch.meshgrid(sy, sx)
256
+ anchor_points.append(torch.stack((sx, sy), -1).view(-1, 2))
257
+ stride_tensor.append(torch.full((h * w, 1), stride, dtype=dtype, device=device))
258
+ return torch.cat(anchor_points), torch.cat(stride_tensor)
259
+
260
+
261
+ def dist2bbox(distance, anchor_points, xywh=True, dim=-1):
262
+ """Transform distance(ltrb) to box(xywh or xyxy)."""
263
+ lt, rb = distance.chunk(2, dim)
264
+ x1y1 = anchor_points - lt
265
+ x2y2 = anchor_points + rb
266
+ if xywh:
267
+ c_xy = (x1y1 + x2y2) / 2
268
+ wh = x2y2 - x1y1
269
+ return torch.cat((c_xy, wh), dim) # xywh bbox
270
+ return torch.cat((x1y1, x2y2), dim) # xyxy bbox
271
+
272
+
273
+ def bbox2dist(anchor_points, bbox, reg_max):
274
+ """Transform bbox(xyxy) to dist(ltrb)."""
275
+ x1y1, x2y2 = bbox.chunk(2, -1)
276
+ return torch.cat((anchor_points - x1y1, x2y2 - anchor_points), -1).clamp_(0, reg_max - 0.01) # dist (lt, rb)