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,579 @@
1
+ # Ultralytics YOLO 🚀, AGPL-3.0 license
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+
7
+ from ..utils.metrics import OKS_SIGMA
8
+ from ..utils.ops import crop_mask, xywh2xyxy, xyxy2xywh
9
+ from ..utils.tal import TaskAlignedAssigner, dist2bbox, make_anchors
10
+
11
+ from .metrics import bbox_iou
12
+ from .tal import bbox2dist
13
+
14
+
15
+ class VarifocalLoss(nn.Module):
16
+ """Varifocal loss by Zhang et al. https://arxiv.org/abs/2008.13367."""
17
+
18
+ def __init__(self):
19
+ """Initialize the VarifocalLoss class."""
20
+ super().__init__()
21
+
22
+ def forward(self, pred_score, gt_score, label, alpha=0.75, gamma=2.0):
23
+ """Computes varfocal loss."""
24
+ weight = (
25
+ alpha * pred_score.sigmoid().pow(gamma) * (1 - label)
26
+ + gt_score * label
27
+ )
28
+ with torch.cuda.amp.autocast(enabled=False):
29
+ loss = (
30
+ (
31
+ F.binary_cross_entropy_with_logits(
32
+ pred_score.float(), gt_score.float(), reduction="none"
33
+ )
34
+ * weight
35
+ )
36
+ .mean(1)
37
+ .sum()
38
+ )
39
+ return loss
40
+
41
+
42
+ # Losses
43
+ class FocalLoss(nn.Module):
44
+ """Wraps focal loss around existing loss_fcn(), i.e. criteria = FocalLoss(nn.BCEWithLogitsLoss(), gamma=1.5)."""
45
+
46
+ def __init__(
47
+ self,
48
+ ):
49
+ super().__init__()
50
+
51
+ def forward(self, pred, label, gamma=1.5, alpha=0.25):
52
+ """Calculates and updates confusion matrix for object detection/classification tasks."""
53
+ loss = F.binary_cross_entropy_with_logits(
54
+ pred, label, reduction="none"
55
+ )
56
+ # p_t = torch.exp(-loss)
57
+ # loss *= self.alpha * (1.000001 - p_t) ** self.gamma # non-zero power for gradient stability
58
+
59
+ # TF implementation https://github.com/tensorflow/addons/blob/v0.7.1/tensorflow_addons/losses/focal_loss.py
60
+ pred_prob = pred.sigmoid() # prob from logits
61
+ p_t = label * pred_prob + (1 - label) * (1 - pred_prob)
62
+ modulating_factor = (1.0 - p_t) ** gamma
63
+ loss *= modulating_factor
64
+ if alpha > 0:
65
+ alpha_factor = label * alpha + (1 - label) * (1 - alpha)
66
+ loss *= alpha_factor
67
+ return loss.mean(1).sum()
68
+
69
+
70
+ class BboxLoss(nn.Module):
71
+
72
+ def __init__(self, reg_max, use_dfl=False):
73
+ """Initialize the BboxLoss module with regularization maximum and DFL settings."""
74
+ super().__init__()
75
+ self.reg_max = reg_max
76
+ self.use_dfl = use_dfl
77
+
78
+ def forward(
79
+ self,
80
+ pred_dist,
81
+ pred_bboxes,
82
+ anchor_points,
83
+ target_bboxes,
84
+ target_scores,
85
+ target_scores_sum,
86
+ fg_mask,
87
+ ):
88
+ """IoU loss."""
89
+ weight = target_scores.sum(-1)[fg_mask].unsqueeze(-1)
90
+ iou = bbox_iou(
91
+ pred_bboxes[fg_mask], target_bboxes[fg_mask], xywh=False, CIoU=True
92
+ )
93
+ loss_iou = ((1.0 - iou) * weight).sum() / target_scores_sum
94
+
95
+ # DFL loss
96
+ if self.use_dfl:
97
+ target_ltrb = bbox2dist(anchor_points, target_bboxes, self.reg_max)
98
+ loss_dfl = (
99
+ self._df_loss(
100
+ pred_dist[fg_mask].view(-1, self.reg_max + 1),
101
+ target_ltrb[fg_mask],
102
+ )
103
+ * weight
104
+ )
105
+ loss_dfl = loss_dfl.sum() / target_scores_sum
106
+ else:
107
+ loss_dfl = torch.tensor(0.0).to(pred_dist.device)
108
+
109
+ return loss_iou, loss_dfl
110
+
111
+ @staticmethod
112
+ def _df_loss(pred_dist, target):
113
+ """Return sum of left and right DFL losses."""
114
+ # Distribution Focal Loss (DFL) proposed in Generalized Focal Loss https://ieeexplore.ieee.org/document/9792391
115
+ tl = target.long() # target left
116
+ tr = tl + 1 # target right
117
+ wl = tr - target # weight left
118
+ wr = 1 - wl # weight right
119
+ return (
120
+ F.cross_entropy(pred_dist, tl.view(-1), reduction="none").view(
121
+ tl.shape
122
+ )
123
+ * wl
124
+ + F.cross_entropy(pred_dist, tr.view(-1), reduction="none").view(
125
+ tl.shape
126
+ )
127
+ * wr
128
+ ).mean(-1, keepdim=True)
129
+
130
+
131
+ class KeypointLoss(nn.Module):
132
+
133
+ def __init__(self, sigmas) -> None:
134
+ super().__init__()
135
+ self.sigmas = sigmas
136
+
137
+ def forward(self, pred_kpts, gt_kpts, kpt_mask, area):
138
+ """Calculates keypoint loss factor and Euclidean distance loss for predicted and actual keypoints."""
139
+ d = (pred_kpts[..., 0] - gt_kpts[..., 0]) ** 2 + (
140
+ pred_kpts[..., 1] - gt_kpts[..., 1]
141
+ ) ** 2
142
+ kpt_loss_factor = (
143
+ torch.sum(kpt_mask != 0) + torch.sum(kpt_mask == 0)
144
+ ) / (torch.sum(kpt_mask != 0) + 1e-9)
145
+ # e = d / (2 * (area * self.sigmas) ** 2 + 1e-9) # from formula
146
+ e = d / (2 * self.sigmas) ** 2 / (area + 1e-9) / 2 # from cocoeval
147
+ return kpt_loss_factor * ((1 - torch.exp(-e)) * kpt_mask).mean()
148
+
149
+
150
+ # Criterion class for computing Detection training losses
151
+ class v8DetectionLoss:
152
+
153
+ def __init__(self, model): # model must be de-paralleled
154
+
155
+ device = next(model.parameters()).device # get model device
156
+ h = model.args # hyperparameters
157
+
158
+ m = model.model[-1] # Detect() module
159
+ self.bce = nn.BCEWithLogitsLoss(reduction="none")
160
+ self.hyp = h
161
+ self.stride = m.stride # model strides
162
+ self.nc = m.nc # number of classes
163
+ self.no = m.no
164
+ self.reg_max = m.reg_max
165
+ self.device = device
166
+
167
+ self.use_dfl = m.reg_max > 1
168
+
169
+ self.assigner = TaskAlignedAssigner(
170
+ topk=10, num_classes=self.nc, alpha=0.5, beta=6.0
171
+ )
172
+ self.bbox_loss = BboxLoss(m.reg_max - 1, use_dfl=self.use_dfl).to(
173
+ device
174
+ )
175
+ self.proj = torch.arange(m.reg_max, dtype=torch.float, device=device)
176
+
177
+ def preprocess(self, targets, batch_size, scale_tensor):
178
+ """Preprocesses the target counts and matches with the input batch size to output a tensor."""
179
+ if targets.shape[0] == 0:
180
+ out = torch.zeros(batch_size, 0, 5, device=self.device)
181
+ else:
182
+ i = targets[:, 0] # image index
183
+ _, counts = i.unique(return_counts=True)
184
+ counts = counts.to(dtype=torch.int32)
185
+ out = torch.zeros(batch_size, counts.max(), 5, device=self.device)
186
+ for j in range(batch_size):
187
+ matches = i == j
188
+ n = matches.sum()
189
+ if n:
190
+ out[j, :n] = targets[matches, 1:]
191
+ out[..., 1:5] = xywh2xyxy(out[..., 1:5].mul_(scale_tensor))
192
+ return out
193
+
194
+ def bbox_decode(self, anchor_points, pred_dist):
195
+ """Decode predicted object bounding box coordinates from anchor points and distribution."""
196
+ if self.use_dfl:
197
+ b, a, c = pred_dist.shape # batch, anchors, channels
198
+ pred_dist = (
199
+ pred_dist.view(b, a, 4, c // 4)
200
+ .softmax(3)
201
+ .matmul(self.proj.type(pred_dist.dtype))
202
+ )
203
+ # pred_dist = pred_dist.view(b, a, c // 4, 4).transpose(2,3).softmax(3).matmul(self.proj.type(pred_dist.dtype))
204
+ # pred_dist = (pred_dist.view(b, a, c // 4, 4).softmax(2) * self.proj.type(pred_dist.dtype).view(1, 1, -1, 1)).sum(2)
205
+ return dist2bbox(pred_dist, anchor_points, xywh=False)
206
+
207
+ def __call__(self, preds, batch):
208
+ """Calculate the sum of the loss for box, cls and dfl multiplied by batch size."""
209
+ loss = torch.zeros(3, device=self.device) # box, cls, dfl
210
+ feats = preds[1] if isinstance(preds, tuple) else preds
211
+ pred_distri, pred_scores = torch.cat(
212
+ [xi.view(feats[0].shape[0], self.no, -1) for xi in feats], 2
213
+ ).split((self.reg_max * 4, self.nc), 1)
214
+
215
+ pred_scores = pred_scores.permute(0, 2, 1).contiguous()
216
+ pred_distri = pred_distri.permute(0, 2, 1).contiguous()
217
+
218
+ dtype = pred_scores.dtype
219
+ batch_size = pred_scores.shape[0]
220
+ imgsz = (
221
+ torch.tensor(feats[0].shape[2:], device=self.device, dtype=dtype)
222
+ * self.stride[0]
223
+ ) # image size (h,w)
224
+ anchor_points, stride_tensor = make_anchors(feats, self.stride, 0.5)
225
+
226
+ # targets
227
+ targets = torch.cat(
228
+ (
229
+ batch["batch_idx"].view(-1, 1),
230
+ batch["cls"].view(-1, 1),
231
+ batch["bboxes"],
232
+ ),
233
+ 1,
234
+ )
235
+ targets = self.preprocess(
236
+ targets.to(self.device),
237
+ batch_size,
238
+ scale_tensor=imgsz[[1, 0, 1, 0]],
239
+ )
240
+ gt_labels, gt_bboxes = targets.split((1, 4), 2) # cls, xyxy
241
+ mask_gt = gt_bboxes.sum(2, keepdim=True).gt_(0)
242
+
243
+ # pboxes
244
+ pred_bboxes = self.bbox_decode(
245
+ anchor_points, pred_distri
246
+ ) # xyxy, (b, h*w, 4)
247
+
248
+ _, target_bboxes, target_scores, fg_mask, _ = self.assigner(
249
+ pred_scores.detach().sigmoid(),
250
+ (pred_bboxes.detach() * stride_tensor).type(gt_bboxes.dtype),
251
+ anchor_points * stride_tensor,
252
+ gt_labels,
253
+ gt_bboxes,
254
+ mask_gt,
255
+ )
256
+
257
+ target_scores_sum = max(target_scores.sum(), 1)
258
+
259
+ # cls loss
260
+ # loss[1] = self.varifocal_loss(pred_scores, target_scores, target_labels) / target_scores_sum # VFL way
261
+ loss[1] = (
262
+ self.bce(pred_scores, target_scores.to(dtype)).sum()
263
+ / target_scores_sum
264
+ ) # BCE
265
+
266
+ # bbox loss
267
+ if fg_mask.sum():
268
+ target_bboxes /= stride_tensor
269
+ loss[0], loss[2] = self.bbox_loss(
270
+ pred_distri,
271
+ pred_bboxes,
272
+ anchor_points,
273
+ target_bboxes,
274
+ target_scores,
275
+ target_scores_sum,
276
+ fg_mask,
277
+ )
278
+
279
+ loss[0] *= self.hyp.box # box gain
280
+ loss[1] *= self.hyp.cls # cls gain
281
+ loss[2] *= self.hyp.dfl # dfl gain
282
+
283
+ return loss.sum() * batch_size, loss.detach() # loss(box, cls, dfl)
284
+
285
+
286
+ # Criterion class for computing training losses
287
+ class v8SegmentationLoss(v8DetectionLoss):
288
+
289
+ def __init__(self, model): # model must be de-paralleled
290
+ super().__init__(model)
291
+ self.nm = model.model[-1].nm # number of masks
292
+ self.overlap = model.args.overlap_mask
293
+
294
+ def __call__(self, preds, batch):
295
+ """Calculate and return the loss for the YOLO model."""
296
+ loss = torch.zeros(4, device=self.device) # box, cls, dfl
297
+ feats, pred_masks, proto = preds if len(preds) == 3 else preds[1]
298
+ batch_size, _, mask_h, mask_w = (
299
+ proto.shape
300
+ ) # batch size, number of masks, mask height, mask width
301
+ pred_distri, pred_scores = torch.cat(
302
+ [xi.view(feats[0].shape[0], self.no, -1) for xi in feats], 2
303
+ ).split((self.reg_max * 4, self.nc), 1)
304
+
305
+ # b, grids, ..
306
+ pred_scores = pred_scores.permute(0, 2, 1).contiguous()
307
+ pred_distri = pred_distri.permute(0, 2, 1).contiguous()
308
+ pred_masks = pred_masks.permute(0, 2, 1).contiguous()
309
+
310
+ dtype = pred_scores.dtype
311
+ imgsz = (
312
+ torch.tensor(feats[0].shape[2:], device=self.device, dtype=dtype)
313
+ * self.stride[0]
314
+ ) # image size (h,w)
315
+ anchor_points, stride_tensor = make_anchors(feats, self.stride, 0.5)
316
+
317
+ # targets
318
+ try:
319
+ batch_idx = batch["batch_idx"].view(-1, 1)
320
+ targets = torch.cat(
321
+ (batch_idx, batch["cls"].view(-1, 1), batch["bboxes"]), 1
322
+ )
323
+ targets = self.preprocess(
324
+ targets.to(self.device),
325
+ batch_size,
326
+ scale_tensor=imgsz[[1, 0, 1, 0]],
327
+ )
328
+ gt_labels, gt_bboxes = targets.split((1, 4), 2) # cls, xyxy
329
+ mask_gt = gt_bboxes.sum(2, keepdim=True).gt_(0)
330
+ except RuntimeError as e:
331
+ raise TypeError(
332
+ "ERROR ❌ segment dataset incorrectly formatted or not a segment dataset.\n"
333
+ "This error can occur when incorrectly training a 'segment' model on a 'detect' dataset, "
334
+ "i.e. 'yolo train model=yolov8n-seg.pt data=coco128.yaml'.\nVerify your dataset is a "
335
+ "correctly formatted 'segment' dataset using 'data=coco128-seg.yaml' "
336
+ "as an example.\nSee https://docs.ultralytics.com/tasks/segment/ for help."
337
+ ) from e
338
+
339
+ # pboxes
340
+ pred_bboxes = self.bbox_decode(
341
+ anchor_points, pred_distri
342
+ ) # xyxy, (b, h*w, 4)
343
+
344
+ _, target_bboxes, target_scores, fg_mask, target_gt_idx = (
345
+ self.assigner(
346
+ pred_scores.detach().sigmoid(),
347
+ (pred_bboxes.detach() * stride_tensor).type(gt_bboxes.dtype),
348
+ anchor_points * stride_tensor,
349
+ gt_labels,
350
+ gt_bboxes,
351
+ mask_gt,
352
+ )
353
+ )
354
+
355
+ target_scores_sum = max(target_scores.sum(), 1)
356
+
357
+ # cls loss
358
+ # loss[1] = self.varifocal_loss(pred_scores, target_scores, target_labels) / target_scores_sum # VFL way
359
+ loss[2] = (
360
+ self.bce(pred_scores, target_scores.to(dtype)).sum()
361
+ / target_scores_sum
362
+ ) # BCE
363
+
364
+ if fg_mask.sum():
365
+ # bbox loss
366
+ loss[0], loss[3] = self.bbox_loss(
367
+ pred_distri,
368
+ pred_bboxes,
369
+ anchor_points,
370
+ target_bboxes / stride_tensor,
371
+ target_scores,
372
+ target_scores_sum,
373
+ fg_mask,
374
+ )
375
+ # masks loss
376
+ masks = batch["masks"].to(self.device).float()
377
+ if tuple(masks.shape[-2:]) != (mask_h, mask_w): # downsample
378
+ masks = F.interpolate(
379
+ masks[None], (mask_h, mask_w), mode="nearest"
380
+ )[0]
381
+
382
+ for i in range(batch_size):
383
+ if fg_mask[i].sum():
384
+ mask_idx = target_gt_idx[i][fg_mask[i]]
385
+ if self.overlap:
386
+ gt_mask = torch.where(
387
+ masks[[i]] == (mask_idx + 1).view(-1, 1, 1),
388
+ 1.0,
389
+ 0.0,
390
+ )
391
+ else:
392
+ gt_mask = masks[batch_idx.view(-1) == i][mask_idx]
393
+ xyxyn = target_bboxes[i][fg_mask[i]] / imgsz[[1, 0, 1, 0]]
394
+ marea = xyxy2xywh(xyxyn)[:, 2:].prod(1)
395
+ mxyxy = xyxyn * torch.tensor(
396
+ [mask_w, mask_h, mask_w, mask_h], device=self.device
397
+ )
398
+ loss[1] += self.single_mask_loss(
399
+ gt_mask,
400
+ pred_masks[i][fg_mask[i]],
401
+ proto[i],
402
+ mxyxy,
403
+ marea,
404
+ ) # seg
405
+
406
+ # WARNING: lines below prevents Multi-GPU DDP 'unused gradient' PyTorch errors, do not remove
407
+ else:
408
+ loss[1] += (proto * 0).sum() + (
409
+ pred_masks * 0
410
+ ).sum() # inf sums may lead to nan loss
411
+
412
+ # WARNING: lines below prevent Multi-GPU DDP 'unused gradient' PyTorch errors, do not remove
413
+ else:
414
+ loss[1] += (proto * 0).sum() + (
415
+ pred_masks * 0
416
+ ).sum() # inf sums may lead to nan loss
417
+
418
+ loss[0] *= self.hyp.box # box gain
419
+ loss[1] *= self.hyp.box / batch_size # seg gain
420
+ loss[2] *= self.hyp.cls # cls gain
421
+ loss[3] *= self.hyp.dfl # dfl gain
422
+
423
+ return loss.sum() * batch_size, loss.detach() # loss(box, cls, dfl)
424
+
425
+ def single_mask_loss(self, gt_mask, pred, proto, xyxy, area):
426
+ """Mask loss for one image."""
427
+ pred_mask = (pred @ proto.view(self.nm, -1)).view(
428
+ -1, *proto.shape[1:]
429
+ ) # (n, 32) @ (32,80,80) -> (n,80,80)
430
+ loss = F.binary_cross_entropy_with_logits(
431
+ pred_mask, gt_mask, reduction="none"
432
+ )
433
+ return (crop_mask(loss, xyxy).mean(dim=(1, 2)) / area).mean()
434
+
435
+
436
+ # Criterion class for computing training losses
437
+ class v8PoseLoss(v8DetectionLoss):
438
+
439
+ def __init__(self, model): # model must be de-paralleled
440
+ super().__init__(model)
441
+ self.kpt_shape = model.model[-1].kpt_shape
442
+ self.bce_pose = nn.BCEWithLogitsLoss()
443
+ is_pose = self.kpt_shape == [17, 3]
444
+ nkpt = self.kpt_shape[0] # number of keypoints
445
+ sigmas = (
446
+ torch.from_numpy(OKS_SIGMA).to(self.device)
447
+ if is_pose
448
+ else torch.ones(nkpt, device=self.device) / nkpt
449
+ )
450
+ self.keypoint_loss = KeypointLoss(sigmas=sigmas)
451
+
452
+ def __call__(self, preds, batch):
453
+ """Calculate the total loss and detach it."""
454
+ loss = torch.zeros(
455
+ 5, device=self.device
456
+ ) # box, cls, dfl, kpt_location, kpt_visibility
457
+ feats, pred_kpts = preds if isinstance(preds[0], list) else preds[1]
458
+ pred_distri, pred_scores = torch.cat(
459
+ [xi.view(feats[0].shape[0], self.no, -1) for xi in feats], 2
460
+ ).split((self.reg_max * 4, self.nc), 1)
461
+
462
+ # b, grids, ..
463
+ pred_scores = pred_scores.permute(0, 2, 1).contiguous()
464
+ pred_distri = pred_distri.permute(0, 2, 1).contiguous()
465
+ pred_kpts = pred_kpts.permute(0, 2, 1).contiguous()
466
+
467
+ dtype = pred_scores.dtype
468
+ imgsz = (
469
+ torch.tensor(feats[0].shape[2:], device=self.device, dtype=dtype)
470
+ * self.stride[0]
471
+ ) # image size (h,w)
472
+ anchor_points, stride_tensor = make_anchors(feats, self.stride, 0.5)
473
+
474
+ # targets
475
+ batch_size = pred_scores.shape[0]
476
+ batch_idx = batch["batch_idx"].view(-1, 1)
477
+ targets = torch.cat(
478
+ (batch_idx, batch["cls"].view(-1, 1), batch["bboxes"]), 1
479
+ )
480
+ targets = self.preprocess(
481
+ targets.to(self.device),
482
+ batch_size,
483
+ scale_tensor=imgsz[[1, 0, 1, 0]],
484
+ )
485
+ gt_labels, gt_bboxes = targets.split((1, 4), 2) # cls, xyxy
486
+ mask_gt = gt_bboxes.sum(2, keepdim=True).gt_(0)
487
+
488
+ # pboxes
489
+ pred_bboxes = self.bbox_decode(
490
+ anchor_points, pred_distri
491
+ ) # xyxy, (b, h*w, 4)
492
+ pred_kpts = self.kpts_decode(
493
+ anchor_points, pred_kpts.view(batch_size, -1, *self.kpt_shape)
494
+ ) # (b, h*w, 17, 3)
495
+
496
+ _, target_bboxes, target_scores, fg_mask, target_gt_idx = (
497
+ self.assigner(
498
+ pred_scores.detach().sigmoid(),
499
+ (pred_bboxes.detach() * stride_tensor).type(gt_bboxes.dtype),
500
+ anchor_points * stride_tensor,
501
+ gt_labels,
502
+ gt_bboxes,
503
+ mask_gt,
504
+ )
505
+ )
506
+
507
+ target_scores_sum = max(target_scores.sum(), 1)
508
+
509
+ # cls loss
510
+ # loss[1] = self.varifocal_loss(pred_scores, target_scores, target_labels) / target_scores_sum # VFL way
511
+ loss[3] = (
512
+ self.bce(pred_scores, target_scores.to(dtype)).sum()
513
+ / target_scores_sum
514
+ ) # BCE
515
+
516
+ # bbox loss
517
+ if fg_mask.sum():
518
+ target_bboxes /= stride_tensor
519
+ loss[0], loss[4] = self.bbox_loss(
520
+ pred_distri,
521
+ pred_bboxes,
522
+ anchor_points,
523
+ target_bboxes,
524
+ target_scores,
525
+ target_scores_sum,
526
+ fg_mask,
527
+ )
528
+ keypoints = batch["keypoints"].to(self.device).float().clone()
529
+ keypoints[..., 0] *= imgsz[1]
530
+ keypoints[..., 1] *= imgsz[0]
531
+ for i in range(batch_size):
532
+ if fg_mask[i].sum():
533
+ idx = target_gt_idx[i][fg_mask[i]]
534
+ gt_kpt = keypoints[batch_idx.view(-1) == i][idx] # (n, 51)
535
+ gt_kpt[..., 0] /= stride_tensor[fg_mask[i]]
536
+ gt_kpt[..., 1] /= stride_tensor[fg_mask[i]]
537
+ area = xyxy2xywh(target_bboxes[i][fg_mask[i]])[:, 2:].prod(
538
+ 1, keepdim=True
539
+ )
540
+ pred_kpt = pred_kpts[i][fg_mask[i]]
541
+ kpt_mask = gt_kpt[..., 2] != 0
542
+ loss[1] += self.keypoint_loss(
543
+ pred_kpt, gt_kpt, kpt_mask, area
544
+ ) # pose loss
545
+ # kpt_score loss
546
+ if pred_kpt.shape[-1] == 3:
547
+ loss[2] += self.bce_pose(
548
+ pred_kpt[..., 2], kpt_mask.float()
549
+ ) # keypoint obj loss
550
+
551
+ loss[0] *= self.hyp.box # box gain
552
+ loss[1] *= self.hyp.pose / batch_size # pose gain
553
+ loss[2] *= self.hyp.kobj / batch_size # kobj gain
554
+ loss[3] *= self.hyp.cls # cls gain
555
+ loss[4] *= self.hyp.dfl # dfl gain
556
+
557
+ return loss.sum() * batch_size, loss.detach() # loss(box, cls, dfl)
558
+
559
+ def kpts_decode(self, anchor_points, pred_kpts):
560
+ """Decodes predicted keypoints to image coordinates."""
561
+ y = pred_kpts.clone()
562
+ y[..., :2] *= 2.0
563
+ y[..., 0] += anchor_points[:, [0]] - 0.5
564
+ y[..., 1] += anchor_points[:, [1]] - 0.5
565
+ return y
566
+
567
+
568
+ class v8ClassificationLoss:
569
+
570
+ def __call__(self, preds, batch):
571
+ """Compute the classification loss between predictions and true labels."""
572
+ loss = (
573
+ torch.nn.functional.cross_entropy(
574
+ preds, batch["cls"], reduction="sum"
575
+ )
576
+ / 64
577
+ )
578
+ loss_items = loss.detach()
579
+ return loss, loss_items