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,468 @@
1
+ # Ultralytics YOLO 🚀, AGPL-3.0 license
2
+ """
3
+ Model head modules
4
+ """
5
+
6
+ import math
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+ from torch.nn.init import constant_, xavier_uniform_
11
+
12
+ from ...yolo.utils.tal import dist2bbox, make_anchors
13
+
14
+ from .block import DFL, Proto
15
+ from .conv import Conv
16
+ from .transformer import (
17
+ MLP,
18
+ DeformableTransformerDecoder,
19
+ DeformableTransformerDecoderLayer,
20
+ )
21
+ from .utils import bias_init_with_prob, linear_init_
22
+
23
+ __all__ = "Detect", "Segment", "Pose", "Classify", "RTDETRDecoder"
24
+
25
+
26
+ class Detect(nn.Module):
27
+ """YOLOv8 Detect head for detection models."""
28
+
29
+ dynamic = False # force grid reconstruction
30
+ export = False # export mode
31
+ shape = None
32
+ anchors = torch.empty(0) # init
33
+ strides = torch.empty(0) # init
34
+
35
+ def __init__(self, nc=80, ch=()): # detection layer
36
+ super().__init__()
37
+ self.nc = nc # number of classes
38
+ self.nl = len(ch) # number of detection layers
39
+ self.reg_max = 26 # DFL channels (ch[0] // 16 to scale 4/8/12/16/20 for n/s/m/l/x)
40
+ self.no = nc + self.reg_max * 4 # number of outputs per anchor
41
+ self.stride = torch.zeros(self.nl) # strides computed during build
42
+ c2, c3 = max((16, ch[0] // 4, self.reg_max * 4)), max(
43
+ ch[0], self.nc
44
+ ) # channels
45
+ self.cv2 = nn.ModuleList(
46
+ nn.Sequential(
47
+ Conv(x, c2, 3),
48
+ Conv(c2, c2, 3),
49
+ nn.Conv2d(c2, 4 * self.reg_max, 1),
50
+ )
51
+ for x in ch
52
+ )
53
+ self.cv3 = nn.ModuleList(
54
+ nn.Sequential(
55
+ Conv(x, c3, 3), Conv(c3, c3, 3), nn.Conv2d(c3, self.nc, 1)
56
+ )
57
+ for x in ch
58
+ )
59
+ self.dfl = DFL(self.reg_max) if self.reg_max > 1 else nn.Identity()
60
+
61
+ def forward(self, x):
62
+ """Concatenates and returns predicted bounding boxes and class probabilities."""
63
+ shape = x[0].shape # BCHW
64
+ for i in range(self.nl):
65
+ x[i] = torch.cat((self.cv2[i](x[i]), self.cv3[i](x[i])), 1)
66
+ if self.training:
67
+ return x
68
+ elif self.dynamic or self.shape != shape:
69
+ self.anchors, self.strides = (
70
+ x.transpose(0, 1) for x in make_anchors(x, self.stride, 0.5)
71
+ )
72
+ self.shape = shape
73
+
74
+ x_cat = torch.cat([xi.view(shape[0], self.no, -1) for xi in x], 2)
75
+ if self.export and self.format in (
76
+ "saved_model",
77
+ "pb",
78
+ "tflite",
79
+ "edgetpu",
80
+ "tfjs",
81
+ ): # avoid TF FlexSplitV ops
82
+ box = x_cat[:, : self.reg_max * 4]
83
+ cls = x_cat[:, self.reg_max * 4 :]
84
+ else:
85
+ box, cls = x_cat.split((self.reg_max * 4, self.nc), 1)
86
+ dbox = (
87
+ dist2bbox(
88
+ self.dfl(box), self.anchors.unsqueeze(0), xywh=True, dim=1
89
+ )
90
+ * self.strides
91
+ )
92
+ y = torch.cat((dbox, cls.sigmoid()), 1)
93
+ return y if self.export else (y, x)
94
+
95
+ def bias_init(self):
96
+ """Initialize Detect() biases, WARNING: requires stride availability."""
97
+ m = self # self.model[-1] # Detect() module
98
+ # cf = torch.bincount(torch.tensor(np.concatenate(dataset.labels, 0)[:, 0]).long(), minlength=nc) + 1
99
+ # ncf = math.log(0.6 / (m.nc - 0.999999)) if cf is None else torch.log(cf / cf.sum()) # nominal class frequency
100
+ for a, b, s in zip(m.cv2, m.cv3, m.stride): # from
101
+ a[-1].bias.data[:] = 1.0 # box
102
+ b[-1].bias.data[: m.nc] = math.log(
103
+ 5 / m.nc / (640 / s) ** 2
104
+ ) # cls (.01 objects, 80 classes, 640 img)
105
+
106
+
107
+ class Segment(Detect):
108
+ """YOLOv8 Segment head for segmentation models."""
109
+
110
+ def __init__(self, nc=80, nm=32, npr=256, ch=()):
111
+ """Initialize the YOLO model attributes such as the number of masks, prototypes, and the convolution layers."""
112
+ super().__init__(nc, ch)
113
+ self.nm = nm # number of masks
114
+ self.npr = npr # number of protos
115
+ # self.proto = Proto(ch[0], self.npr, self.nm) # protos
116
+ self.detect = Detect.forward
117
+
118
+ c4 = max(ch[0] // 4, self.nm)
119
+ self.cv4 = nn.ModuleList(
120
+ nn.Sequential(
121
+ Conv(x, c4, 3), Conv(c4, c4, 3), nn.Conv2d(c4, self.nm, 1)
122
+ )
123
+ for x in ch
124
+ )
125
+
126
+ def forward(self, x):
127
+ """Return model outputs and mask coefficients if training, otherwise return outputs and mask coefficients."""
128
+ # p = self.proto(x[0]) # mask protos #mobilesamv2 change
129
+ p = 0
130
+ # import pdb;pdb.set_trace()
131
+ bs = x[0].shape[0] # batch size
132
+
133
+ mc = torch.cat(
134
+ [self.cv4[i](x[i]).view(bs, self.nm, -1) for i in range(self.nl)],
135
+ 2,
136
+ ) # mask coefficients
137
+ x = self.detect(self, x)
138
+ if self.training:
139
+ return x, mc, p
140
+ return (
141
+ (torch.cat([x, mc], 1), p)
142
+ if self.export
143
+ else (torch.cat([x[0], mc], 1), (x[1], mc, p))
144
+ )
145
+
146
+
147
+ class Pose(Detect):
148
+ """YOLOv8 Pose head for keypoints models."""
149
+
150
+ def __init__(self, nc=80, kpt_shape=(17, 3), ch=()):
151
+ """Initialize YOLO network with default parameters and Convolutional Layers."""
152
+ super().__init__(nc, ch)
153
+ self.kpt_shape = kpt_shape # number of keypoints, number of dims (2 for x,y or 3 for x,y,visible)
154
+ self.nk = kpt_shape[0] * kpt_shape[1] # number of keypoints total
155
+ self.detect = Detect.forward
156
+
157
+ c4 = max(ch[0] // 4, self.nk)
158
+ self.cv4 = nn.ModuleList(
159
+ nn.Sequential(
160
+ Conv(x, c4, 3), Conv(c4, c4, 3), nn.Conv2d(c4, self.nk, 1)
161
+ )
162
+ for x in ch
163
+ )
164
+
165
+ def forward(self, x):
166
+ """Perform forward pass through YOLO model and return predictions."""
167
+ bs = x[0].shape[0] # batch size
168
+ kpt = torch.cat(
169
+ [self.cv4[i](x[i]).view(bs, self.nk, -1) for i in range(self.nl)],
170
+ -1,
171
+ ) # (bs, 17*3, h*w)
172
+ x = self.detect(self, x)
173
+ if self.training:
174
+ return x, kpt
175
+ pred_kpt = self.kpts_decode(bs, kpt)
176
+ return (
177
+ torch.cat([x, pred_kpt], 1)
178
+ if self.export
179
+ else (torch.cat([x[0], pred_kpt], 1), (x[1], kpt))
180
+ )
181
+
182
+ def kpts_decode(self, bs, kpts):
183
+ """Decodes keypoints."""
184
+ ndim = self.kpt_shape[1]
185
+ if (
186
+ self.export
187
+ ): # required for TFLite export to avoid 'PLACEHOLDER_FOR_GREATER_OP_CODES' bug
188
+ y = kpts.view(bs, *self.kpt_shape, -1)
189
+ a = (y[:, :, :2] * 2.0 + (self.anchors - 0.5)) * self.strides
190
+ if ndim == 3:
191
+ a = torch.cat((a, y[:, :, 2:3].sigmoid()), 2)
192
+ return a.view(bs, self.nk, -1)
193
+ else:
194
+ y = kpts.clone()
195
+ if ndim == 3:
196
+ y[:, 2::3].sigmoid_() # inplace sigmoid
197
+ y[:, 0::ndim] = (
198
+ y[:, 0::ndim] * 2.0 + (self.anchors[0] - 0.5)
199
+ ) * self.strides
200
+ y[:, 1::ndim] = (
201
+ y[:, 1::ndim] * 2.0 + (self.anchors[1] - 0.5)
202
+ ) * self.strides
203
+ return y
204
+
205
+
206
+ class Classify(nn.Module):
207
+ """YOLOv8 classification head, i.e. x(b,c1,20,20) to x(b,c2)."""
208
+
209
+ def __init__(
210
+ self, c1, c2, k=1, s=1, p=None, g=1
211
+ ): # ch_in, ch_out, kernel, stride, padding, groups
212
+ super().__init__()
213
+ c_ = 1280 # efficientnet_b0 size
214
+ self.conv = Conv(c1, c_, k, s, p, g)
215
+ self.pool = nn.AdaptiveAvgPool2d(1) # to x(b,c_,1,1)
216
+ self.drop = nn.Dropout(p=0.0, inplace=True)
217
+ self.linear = nn.Linear(c_, c2) # to x(b,c2)
218
+
219
+ def forward(self, x):
220
+ """Performs a forward pass of the YOLO model on input image data."""
221
+ if isinstance(x, list):
222
+ x = torch.cat(x, 1)
223
+ x = self.linear(self.drop(self.pool(self.conv(x)).flatten(1)))
224
+ return x if self.training else x.softmax(1)
225
+
226
+
227
+ class RTDETRDecoder(nn.Module):
228
+
229
+ def __init__(
230
+ self,
231
+ nc=80,
232
+ ch=(512, 1024, 2048),
233
+ hd=256, # hidden dim
234
+ nq=300, # num queries
235
+ ndp=4, # num decoder points
236
+ nh=8, # num head
237
+ ndl=6, # num decoder layers
238
+ d_ffn=1024, # dim of feedforward
239
+ dropout=0.0,
240
+ act=nn.ReLU(),
241
+ eval_idx=-1,
242
+ # training args
243
+ nd=100, # num denoising
244
+ label_noise_ratio=0.5,
245
+ box_noise_scale=1.0,
246
+ learnt_init_query=False,
247
+ ):
248
+ super().__init__()
249
+ self.hidden_dim = hd
250
+ self.nhead = nh
251
+ self.nl = len(ch) # num level
252
+ self.nc = nc
253
+ self.num_queries = nq
254
+ self.num_decoder_layers = ndl
255
+
256
+ # backbone feature projection
257
+ self.input_proj = nn.ModuleList(
258
+ nn.Sequential(nn.Conv2d(x, hd, 1, bias=False), nn.BatchNorm2d(hd))
259
+ for x in ch
260
+ )
261
+ # NOTE: simplified version but it's not consistent with .pt weights.
262
+ # self.input_proj = nn.ModuleList(Conv(x, hd, act=False) for x in ch)
263
+
264
+ # Transformer module
265
+ decoder_layer = DeformableTransformerDecoderLayer(
266
+ hd, nh, d_ffn, dropout, act, self.nl, ndp
267
+ )
268
+ self.decoder = DeformableTransformerDecoder(
269
+ hd, decoder_layer, ndl, eval_idx
270
+ )
271
+
272
+ # denoising part
273
+ self.denoising_class_embed = nn.Embedding(nc, hd)
274
+ self.num_denoising = nd
275
+ self.label_noise_ratio = label_noise_ratio
276
+ self.box_noise_scale = box_noise_scale
277
+
278
+ # decoder embedding
279
+ self.learnt_init_query = learnt_init_query
280
+ if learnt_init_query:
281
+ self.tgt_embed = nn.Embedding(nq, hd)
282
+ self.query_pos_head = MLP(4, 2 * hd, hd, num_layers=2)
283
+
284
+ # encoder head
285
+ self.enc_output = nn.Sequential(nn.Linear(hd, hd), nn.LayerNorm(hd))
286
+ self.enc_score_head = nn.Linear(hd, nc)
287
+ self.enc_bbox_head = MLP(hd, hd, 4, num_layers=3)
288
+
289
+ # decoder head
290
+ self.dec_score_head = nn.ModuleList(
291
+ [nn.Linear(hd, nc) for _ in range(ndl)]
292
+ )
293
+ self.dec_bbox_head = nn.ModuleList(
294
+ [MLP(hd, hd, 4, num_layers=3) for _ in range(ndl)]
295
+ )
296
+
297
+ self._reset_parameters()
298
+
299
+ def forward(self, x, batch=None):
300
+ from ultralytics.vit.utils.ops import get_cdn_group
301
+
302
+ # input projection and embedding
303
+ feats, shapes = self._get_encoder_input(x)
304
+
305
+ # prepare denoising training
306
+ dn_embed, dn_bbox, attn_mask, dn_meta = get_cdn_group(
307
+ batch,
308
+ self.nc,
309
+ self.num_queries,
310
+ self.denoising_class_embed.weight,
311
+ self.num_denoising,
312
+ self.label_noise_ratio,
313
+ self.box_noise_scale,
314
+ self.training,
315
+ )
316
+
317
+ embed, refer_bbox, enc_bboxes, enc_scores = self._get_decoder_input(
318
+ feats, shapes, dn_embed, dn_bbox
319
+ )
320
+
321
+ # decoder
322
+ dec_bboxes, dec_scores = self.decoder(
323
+ embed,
324
+ refer_bbox,
325
+ feats,
326
+ shapes,
327
+ self.dec_bbox_head,
328
+ self.dec_score_head,
329
+ self.query_pos_head,
330
+ attn_mask=attn_mask,
331
+ )
332
+ if not self.training:
333
+ dec_scores = dec_scores.sigmoid_()
334
+ return dec_bboxes, dec_scores, enc_bboxes, enc_scores, dn_meta
335
+
336
+ def _generate_anchors(
337
+ self,
338
+ shapes,
339
+ grid_size=0.05,
340
+ dtype=torch.float32,
341
+ device="cpu",
342
+ eps=1e-2,
343
+ ):
344
+ anchors = []
345
+ for i, (h, w) in enumerate(shapes):
346
+ grid_y, grid_x = torch.meshgrid(
347
+ torch.arange(end=h, dtype=dtype, device=device),
348
+ torch.arange(end=w, dtype=dtype, device=device),
349
+ indexing="ij",
350
+ )
351
+ grid_xy = torch.stack([grid_x, grid_y], -1) # (h, w, 2)
352
+
353
+ valid_WH = torch.tensor([h, w], dtype=dtype, device=device)
354
+ grid_xy = (grid_xy.unsqueeze(0) + 0.5) / valid_WH # (1, h, w, 2)
355
+ wh = (
356
+ torch.ones_like(grid_xy, dtype=dtype, device=device)
357
+ * grid_size
358
+ * (2.0**i)
359
+ )
360
+ anchors.append(
361
+ torch.cat([grid_xy, wh], -1).view(-1, h * w, 4)
362
+ ) # (1, h*w, 4)
363
+
364
+ anchors = torch.cat(anchors, 1) # (1, h*w*nl, 4)
365
+ valid_mask = ((anchors > eps) * (anchors < 1 - eps)).all(
366
+ -1, keepdim=True
367
+ ) # 1, h*w*nl, 1
368
+ anchors = torch.log(anchors / (1 - anchors))
369
+ anchors = torch.where(valid_mask, anchors, torch.inf)
370
+ return anchors, valid_mask
371
+
372
+ def _get_encoder_input(self, x):
373
+ # get projection features
374
+ x = [self.input_proj[i](feat) for i, feat in enumerate(x)]
375
+ # get encoder inputs
376
+ feats = []
377
+ shapes = []
378
+ for feat in x:
379
+ h, w = feat.shape[2:]
380
+ # [b, c, h, w] -> [b, h*w, c]
381
+ feats.append(feat.flatten(2).permute(0, 2, 1))
382
+ # [nl, 2]
383
+ shapes.append([h, w])
384
+
385
+ # [b, h*w, c]
386
+ feats = torch.cat(feats, 1)
387
+ return feats, shapes
388
+
389
+ def _get_decoder_input(self, feats, shapes, dn_embed=None, dn_bbox=None):
390
+ bs = len(feats)
391
+ # prepare input for decoder
392
+ anchors, valid_mask = self._generate_anchors(
393
+ shapes, dtype=feats.dtype, device=feats.device
394
+ )
395
+ features = self.enc_output(
396
+ torch.where(valid_mask, feats, 0)
397
+ ) # bs, h*w, 256
398
+
399
+ enc_outputs_scores = self.enc_score_head(features) # (bs, h*w, nc)
400
+ # dynamic anchors + static content
401
+ enc_outputs_bboxes = (
402
+ self.enc_bbox_head(features) + anchors
403
+ ) # (bs, h*w, 4)
404
+
405
+ # query selection
406
+ # (bs, num_queries)
407
+ topk_ind = torch.topk(
408
+ enc_outputs_scores.max(-1).values, self.num_queries, dim=1
409
+ ).indices.view(-1)
410
+ # (bs, num_queries)
411
+ batch_ind = (
412
+ torch.arange(end=bs, dtype=topk_ind.dtype)
413
+ .unsqueeze(-1)
414
+ .repeat(1, self.num_queries)
415
+ .view(-1)
416
+ )
417
+
418
+ # Unsigmoided
419
+ refer_bbox = enc_outputs_bboxes[batch_ind, topk_ind].view(
420
+ bs, self.num_queries, -1
421
+ )
422
+ # refer_bbox = torch.gather(enc_outputs_bboxes, 1, topk_ind.reshape(bs, self.num_queries).unsqueeze(-1).repeat(1, 1, 4))
423
+
424
+ enc_bboxes = refer_bbox.sigmoid()
425
+ if dn_bbox is not None:
426
+ refer_bbox = torch.cat([dn_bbox, refer_bbox], 1)
427
+ if self.training:
428
+ refer_bbox = refer_bbox.detach()
429
+ enc_scores = enc_outputs_scores[batch_ind, topk_ind].view(
430
+ bs, self.num_queries, -1
431
+ )
432
+
433
+ if self.learnt_init_query:
434
+ embeddings = self.tgt_embed.weight.unsqueeze(0).repeat(bs, 1, 1)
435
+ else:
436
+ embeddings = features[batch_ind, topk_ind].view(
437
+ bs, self.num_queries, -1
438
+ )
439
+ if self.training:
440
+ embeddings = embeddings.detach()
441
+ if dn_embed is not None:
442
+ embeddings = torch.cat([dn_embed, embeddings], 1)
443
+
444
+ return embeddings, refer_bbox, enc_bboxes, enc_scores
445
+
446
+ # TODO
447
+ def _reset_parameters(self):
448
+ # class and bbox head init
449
+ bias_cls = bias_init_with_prob(0.01) / 80 * self.nc
450
+ # NOTE: the weight initialization in `linear_init_` would cause NaN when training with custom datasets.
451
+ # linear_init_(self.enc_score_head)
452
+ constant_(self.enc_score_head.bias, bias_cls)
453
+ constant_(self.enc_bbox_head.layers[-1].weight, 0.0)
454
+ constant_(self.enc_bbox_head.layers[-1].bias, 0.0)
455
+ for cls_, reg_ in zip(self.dec_score_head, self.dec_bbox_head):
456
+ # linear_init_(cls_)
457
+ constant_(cls_.bias, bias_cls)
458
+ constant_(reg_.layers[-1].weight, 0.0)
459
+ constant_(reg_.layers[-1].bias, 0.0)
460
+
461
+ linear_init_(self.enc_output[0])
462
+ xavier_uniform_(self.enc_output[0].weight)
463
+ if self.learnt_init_query:
464
+ xavier_uniform_(self.tgt_embed.weight)
465
+ xavier_uniform_(self.query_pos_head.layers[0].weight)
466
+ xavier_uniform_(self.query_pos_head.layers[1].weight)
467
+ for layer in self.input_proj:
468
+ xavier_uniform_(layer[0].weight)