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,397 @@
1
+ # Ultralytics YOLO 🚀, AGPL-3.0 license
2
+ """
3
+ Common modules
4
+ """
5
+
6
+ from copy import copy
7
+ from pathlib import Path
8
+
9
+ import cv2
10
+ import numpy as np
11
+ import requests
12
+ import torch
13
+ import torch.nn as nn
14
+ from PIL import Image, ImageOps
15
+ from torch.cuda import amp
16
+
17
+ from ..nn.autobackend import AutoBackend
18
+ from ..yolo.data.augment import LetterBox
19
+ from ..yolo.utils import LOGGER, colorstr
20
+ from ..yolo.utils.files import increment_path
21
+ from ..yolo.utils.ops import (
22
+ Profile,
23
+ make_divisible,
24
+ non_max_suppression,
25
+ scale_boxes,
26
+ xyxy2xywh,
27
+ )
28
+ from ..yolo.utils.plotting import Annotator, colors, save_one_box
29
+ from ..yolo.utils.torch_utils import copy_attr, smart_inference_mode
30
+
31
+
32
+ class AutoShape(nn.Module):
33
+ """YOLOv8 input-robust model wrapper for passing cv2/np/PIL/torch inputs. Includes preprocessing, inference and NMS."""
34
+
35
+ conf = 0.25 # NMS confidence threshold
36
+ iou = 0.45 # NMS IoU threshold
37
+ agnostic = False # NMS class-agnostic
38
+ multi_label = False # NMS multiple labels per box
39
+ classes = None # (optional list) filter by class, i.e. = [0, 15, 16] for COCO persons, cats and dogs
40
+ max_det = 1000 # maximum number of detections per image
41
+ amp = False # Automatic Mixed Precision (AMP) inference
42
+
43
+ def __init__(self, model, verbose=True):
44
+ """Initializes object and copies attributes from model object."""
45
+ super().__init__()
46
+ if verbose:
47
+ LOGGER.info("Adding AutoShape... ")
48
+ copy_attr(
49
+ self,
50
+ model,
51
+ include=("yaml", "nc", "hyp", "names", "stride", "abc"),
52
+ exclude=(),
53
+ ) # copy attributes
54
+ self.dmb = isinstance(
55
+ model, AutoBackend
56
+ ) # DetectMultiBackend() instance
57
+ self.pt = not self.dmb or model.pt # PyTorch model
58
+ self.model = model.eval()
59
+ if self.pt:
60
+ m = (
61
+ self.model.model.model[-1]
62
+ if self.dmb
63
+ else self.model.model[-1]
64
+ ) # Detect()
65
+ m.inplace = (
66
+ False # Detect.inplace=False for safe multithread inference
67
+ )
68
+ m.export = True # do not output loss values
69
+
70
+ def _apply(self, fn):
71
+ """Apply to(), cpu(), cuda(), half() to model tensors that are not parameters or registered buffers."""
72
+ self = super()._apply(fn)
73
+ if self.pt:
74
+ m = (
75
+ self.model.model.model[-1]
76
+ if self.dmb
77
+ else self.model.model[-1]
78
+ ) # Detect()
79
+ m.stride = fn(m.stride)
80
+ m.grid = list(map(fn, m.grid))
81
+ if isinstance(m.anchor_grid, list):
82
+ m.anchor_grid = list(map(fn, m.anchor_grid))
83
+ return self
84
+
85
+ @smart_inference_mode()
86
+ def forward(self, ims, size=640, augment=False, profile=False):
87
+ """Inference from various sources. For size(height=640, width=1280), RGB images example inputs are:."""
88
+ # file: ims = 'data/images/zidane.jpg' # str or PosixPath
89
+ # URI: = 'https://ultralytics.com/images/zidane.jpg'
90
+ # OpenCV: = cv2.imread('image.jpg')[:,:,::-1] # HWC BGR to RGB x(640,1280,3)
91
+ # PIL: = Image.open('image.jpg') or ImageGrab.grab() # HWC x(640,1280,3)
92
+ # numpy: = np.zeros((640,1280,3)) # HWC
93
+ # torch: = torch.zeros(16,3,320,640) # BCHW (scaled to size=640, 0-1 values)
94
+ # multiple: = [Image.open('image1.jpg'), Image.open('image2.jpg'), ...] # list of images
95
+
96
+ dt = (Profile(), Profile(), Profile())
97
+ with dt[0]:
98
+ if isinstance(size, int): # expand
99
+ size = (size, size)
100
+ p = (
101
+ next(self.model.parameters())
102
+ if self.pt
103
+ else torch.empty(1, device=self.model.device)
104
+ ) # param
105
+ autocast = self.amp and (
106
+ p.device.type != "cpu"
107
+ ) # Automatic Mixed Precision (AMP) inference
108
+ if isinstance(ims, torch.Tensor): # torch
109
+ with amp.autocast(autocast):
110
+ return self.model(
111
+ ims.to(p.device).type_as(p), augment=augment
112
+ ) # inference
113
+
114
+ # Preprocess
115
+ n, ims = (
116
+ (len(ims), list(ims))
117
+ if isinstance(ims, (list, tuple))
118
+ else (1, [ims])
119
+ ) # number, list of images
120
+ shape0, shape1, files = (
121
+ [],
122
+ [],
123
+ [],
124
+ ) # image and inference shapes, filenames
125
+ for i, im in enumerate(ims):
126
+ f = f"image{i}" # filename
127
+ if isinstance(im, (str, Path)): # filename or uri
128
+ im, f = (
129
+ Image.open(
130
+ requests.get(im, stream=True).raw
131
+ if str(im).startswith("http")
132
+ else im
133
+ ),
134
+ im,
135
+ )
136
+ im = np.asarray(ImageOps.exif_transpose(im))
137
+ elif isinstance(im, Image.Image): # PIL Image
138
+ im, f = (
139
+ np.asarray(ImageOps.exif_transpose(im)),
140
+ getattr(im, "filename", f) or f,
141
+ )
142
+ files.append(Path(f).with_suffix(".jpg").name)
143
+ if im.shape[0] < 5: # image in CHW
144
+ im = im.transpose(
145
+ (1, 2, 0)
146
+ ) # reverse dataloader .transpose(2, 0, 1)
147
+ im = (
148
+ im[..., :3]
149
+ if im.ndim == 3
150
+ else cv2.cvtColor(im, cv2.COLOR_GRAY2BGR)
151
+ ) # enforce 3ch input
152
+ s = im.shape[:2] # HWC
153
+ shape0.append(s) # image shape
154
+ g = max(size) / max(s) # gain
155
+ shape1.append([y * g for y in s])
156
+ ims[i] = (
157
+ im if im.data.contiguous else np.ascontiguousarray(im)
158
+ ) # update
159
+ shape1 = (
160
+ [
161
+ make_divisible(x, self.stride)
162
+ for x in np.array(shape1).max(0)
163
+ ]
164
+ if self.pt
165
+ else size
166
+ ) # inf shape
167
+ x = [
168
+ LetterBox(shape1, auto=False)(image=im)["img"] for im in ims
169
+ ] # pad
170
+ x = np.ascontiguousarray(
171
+ np.array(x).transpose((0, 3, 1, 2))
172
+ ) # stack and BHWC to BCHW
173
+ x = (
174
+ torch.from_numpy(x).to(p.device).type_as(p) / 255
175
+ ) # uint8 to fp16/32
176
+
177
+ with amp.autocast(autocast):
178
+ # Inference
179
+ with dt[1]:
180
+ y = self.model(x, augment=augment) # forward
181
+
182
+ # Postprocess
183
+ with dt[2]:
184
+ y = non_max_suppression(
185
+ y if self.dmb else y[0],
186
+ self.conf,
187
+ self.iou,
188
+ self.classes,
189
+ self.agnostic,
190
+ self.multi_label,
191
+ max_det=self.max_det,
192
+ ) # NMS
193
+ for i in range(n):
194
+ scale_boxes(shape1, y[i][:, :4], shape0[i])
195
+
196
+ return Detections(ims, y, files, dt, self.names, x.shape)
197
+
198
+
199
+ class Detections:
200
+ """YOLOv8 detections class for inference results"""
201
+
202
+ def __init__(
203
+ self, ims, pred, files, times=(0, 0, 0), names=None, shape=None
204
+ ):
205
+ """Initialize object attributes for YOLO detection results."""
206
+ super().__init__()
207
+ d = pred[0].device # device
208
+ gn = [
209
+ torch.tensor(
210
+ [*(im.shape[i] for i in [1, 0, 1, 0]), 1, 1], device=d
211
+ )
212
+ for im in ims
213
+ ] # normalizations
214
+ self.ims = ims # list of images as numpy arrays
215
+ self.pred = pred # list of tensors pred[0] = (xyxy, conf, cls)
216
+ self.names = names # class names
217
+ self.files = files # image filenames
218
+ self.times = times # profiling times
219
+ self.xyxy = pred # xyxy pixels
220
+ self.xywh = [xyxy2xywh(x) for x in pred] # xywh pixels
221
+ self.xyxyn = [x / g for x, g in zip(self.xyxy, gn)] # xyxy normalized
222
+ self.xywhn = [x / g for x, g in zip(self.xywh, gn)] # xywh normalized
223
+ self.n = len(self.pred) # number of images (batch size)
224
+ self.t = tuple(x.t / self.n * 1e3 for x in times) # timestamps (ms)
225
+ self.s = tuple(shape) # inference BCHW shape
226
+
227
+ def _run(
228
+ self,
229
+ pprint=False,
230
+ show=False,
231
+ save=False,
232
+ crop=False,
233
+ render=False,
234
+ labels=True,
235
+ save_dir=Path(""),
236
+ ):
237
+ """Return performance metrics and optionally cropped/save images or results."""
238
+ s, crops = "", []
239
+ for i, (im, pred) in enumerate(zip(self.ims, self.pred)):
240
+ s += f"\nimage {i + 1}/{len(self.pred)}: {im.shape[0]}x{im.shape[1]} " # string
241
+ if pred.shape[0]:
242
+ for c in pred[:, -1].unique():
243
+ n = (pred[:, -1] == c).sum() # detections per class
244
+ s += f"{n} {self.names[int(c)]}{'s' * (n > 1)}, " # add to string
245
+ s = s.rstrip(", ")
246
+ if show or save or render or crop:
247
+ annotator = Annotator(im, example=str(self.names))
248
+ for *box, conf, cls in reversed(
249
+ pred
250
+ ): # xyxy, confidence, class
251
+ label = f"{self.names[int(cls)]} {conf:.2f}"
252
+ if crop:
253
+ file = (
254
+ save_dir
255
+ / "crops"
256
+ / self.names[int(cls)]
257
+ / self.files[i]
258
+ if save
259
+ else None
260
+ )
261
+ crops.append(
262
+ {
263
+ "box": box,
264
+ "conf": conf,
265
+ "cls": cls,
266
+ "label": label,
267
+ "im": save_one_box(
268
+ box, im, file=file, save=save
269
+ ),
270
+ }
271
+ )
272
+ else: # all others
273
+ annotator.box_label(
274
+ box, label if labels else "", color=colors(cls)
275
+ )
276
+ im = annotator.im
277
+ else:
278
+ s += "(no detections)"
279
+
280
+ im = (
281
+ Image.fromarray(im.astype(np.uint8))
282
+ if isinstance(im, np.ndarray)
283
+ else im
284
+ ) # from np
285
+ if show:
286
+ im.show(self.files[i]) # show
287
+ if save:
288
+ f = self.files[i]
289
+ im.save(save_dir / f) # save
290
+ if i == self.n - 1:
291
+ LOGGER.info(
292
+ f"Saved {self.n} image{'s' * (self.n > 1)} to {colorstr('bold', save_dir)}"
293
+ )
294
+ if render:
295
+ self.ims[i] = np.asarray(im)
296
+ if pprint:
297
+ s = s.lstrip("\n")
298
+ return (
299
+ f"{s}\nSpeed: %.1fms preprocess, %.1fms inference, %.1fms NMS per image at shape {self.s}"
300
+ % self.t
301
+ )
302
+ if crop:
303
+ if save:
304
+ LOGGER.info(f"Saved results to {save_dir}\n")
305
+ return crops
306
+
307
+ def show(self, labels=True):
308
+ """Displays YOLO results with detected bounding boxes."""
309
+ self._run(show=True, labels=labels) # show results
310
+
311
+ def save(self, labels=True, save_dir="runs/detect/exp", exist_ok=False):
312
+ """Save detection results with optional labels to specified directory."""
313
+ save_dir = increment_path(
314
+ save_dir, exist_ok, mkdir=True
315
+ ) # increment save_dir
316
+ self._run(save=True, labels=labels, save_dir=save_dir) # save results
317
+
318
+ def crop(self, save=True, save_dir="runs/detect/exp", exist_ok=False):
319
+ """Crops images into detections and saves them if 'save' is True."""
320
+ save_dir = (
321
+ increment_path(save_dir, exist_ok, mkdir=True) if save else None
322
+ )
323
+ return self._run(
324
+ crop=True, save=save, save_dir=save_dir
325
+ ) # crop results
326
+
327
+ def render(self, labels=True):
328
+ """Renders detected objects and returns images."""
329
+ self._run(render=True, labels=labels) # render results
330
+ return self.ims
331
+
332
+ def pandas(self):
333
+ """Return detections as pandas DataFrames, i.e. print(results.pandas().xyxy[0])."""
334
+ import pandas
335
+
336
+ new = copy(self) # return copy
337
+ ca = (
338
+ "xmin",
339
+ "ymin",
340
+ "xmax",
341
+ "ymax",
342
+ "confidence",
343
+ "class",
344
+ "name",
345
+ ) # xyxy columns
346
+ cb = (
347
+ "xcenter",
348
+ "ycenter",
349
+ "width",
350
+ "height",
351
+ "confidence",
352
+ "class",
353
+ "name",
354
+ ) # xywh columns
355
+ for k, c in zip(["xyxy", "xyxyn", "xywh", "xywhn"], [ca, ca, cb, cb]):
356
+ a = [
357
+ [
358
+ x[:5] + [int(x[5]), self.names[int(x[5])]]
359
+ for x in x.tolist()
360
+ ]
361
+ for x in getattr(self, k)
362
+ ] # update
363
+ setattr(new, k, [pandas.DataFrame(x, columns=c) for x in a])
364
+ return new
365
+
366
+ def tolist(self):
367
+ """Return a list of Detections objects, i.e. 'for result in results.tolist():'."""
368
+ r = range(self.n) # iterable
369
+ x = [
370
+ Detections(
371
+ [self.ims[i]],
372
+ [self.pred[i]],
373
+ [self.files[i]],
374
+ self.times,
375
+ self.names,
376
+ self.s,
377
+ )
378
+ for i in r
379
+ ]
380
+ # for d in x:
381
+ # for k in ['ims', 'pred', 'xyxy', 'xyxyn', 'xywh', 'xywhn']:
382
+ # setattr(d, k, getattr(d, k)[0]) # pop out of list
383
+ return x
384
+
385
+ def print(self):
386
+ """Print the results of the `self._run()` function."""
387
+ LOGGER.info(self.__str__())
388
+
389
+ def __len__(self): # override len(results)
390
+ return self.n
391
+
392
+ def __str__(self): # override print(results)
393
+ return self._run(pprint=True) # print results
394
+
395
+ def __repr__(self):
396
+ """Returns a printable representation of the object."""
397
+ return f"YOLOv8 {self.__class__} instance\n" + self.__str__()
@@ -0,0 +1,110 @@
1
+ # Ultralytics YOLO 🚀, AGPL-3.0 license
2
+ """
3
+ Ultralytics modules. Visualize with:
4
+
5
+ from ultralytics.nn.modules import *
6
+ import torch
7
+ import os
8
+
9
+ x = torch.ones(1, 128, 40, 40)
10
+ m = Conv(128, 128)
11
+ f = f'{m._get_name()}.onnx'
12
+ torch.onnx.export(m, x, f)
13
+ os.system(f'onnxsim {f} {f} && open {f}')
14
+ """
15
+
16
+ from .block import (
17
+ C1,
18
+ C2,
19
+ C3,
20
+ C3TR,
21
+ DFL,
22
+ SPP,
23
+ SPPF,
24
+ Bottleneck,
25
+ BottleneckCSP,
26
+ C2f,
27
+ C3Ghost,
28
+ C3x,
29
+ GhostBottleneck,
30
+ HGBlock,
31
+ HGStem,
32
+ Proto,
33
+ RepC3,
34
+ )
35
+ from .conv import (
36
+ CBAM,
37
+ ChannelAttention,
38
+ Concat,
39
+ Conv,
40
+ Conv2,
41
+ ConvTranspose,
42
+ DWConv,
43
+ DWConvTranspose2d,
44
+ Focus,
45
+ GhostConv,
46
+ LightConv,
47
+ RepConv,
48
+ SpatialAttention,
49
+ )
50
+ from .head import Classify, Detect, Pose, RTDETRDecoder, Segment
51
+ from .transformer import (
52
+ AIFI,
53
+ MLP,
54
+ DeformableTransformerDecoder,
55
+ DeformableTransformerDecoderLayer,
56
+ LayerNorm2d,
57
+ MLPBlock,
58
+ MSDeformAttn,
59
+ TransformerBlock,
60
+ TransformerEncoderLayer,
61
+ TransformerLayer,
62
+ )
63
+
64
+ __all__ = (
65
+ "Conv",
66
+ "Conv2",
67
+ "LightConv",
68
+ "RepConv",
69
+ "DWConv",
70
+ "DWConvTranspose2d",
71
+ "ConvTranspose",
72
+ "Focus",
73
+ "GhostConv",
74
+ "ChannelAttention",
75
+ "SpatialAttention",
76
+ "CBAM",
77
+ "Concat",
78
+ "TransformerLayer",
79
+ "TransformerBlock",
80
+ "MLPBlock",
81
+ "LayerNorm2d",
82
+ "DFL",
83
+ "HGBlock",
84
+ "HGStem",
85
+ "SPP",
86
+ "SPPF",
87
+ "C1",
88
+ "C2",
89
+ "C3",
90
+ "C2f",
91
+ "C3x",
92
+ "C3TR",
93
+ "C3Ghost",
94
+ "GhostBottleneck",
95
+ "Bottleneck",
96
+ "BottleneckCSP",
97
+ "Proto",
98
+ "Detect",
99
+ "Segment",
100
+ "Pose",
101
+ "Classify",
102
+ "TransformerEncoderLayer",
103
+ "RepC3",
104
+ "RTDETRDecoder",
105
+ "AIFI",
106
+ "DeformableTransformerDecoder",
107
+ "DeformableTransformerDecoderLayer",
108
+ "MSDeformAttn",
109
+ "MLP",
110
+ )