segment-everything 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- segment_everything/__init__.py +5 -0
- segment_everything/augmentation/albumentations_helper.py +0 -0
- segment_everything/detect_and_segment.py +131 -0
- segment_everything/napari_helper.py +15 -0
- segment_everything/prompt_generator.py +188 -0
- segment_everything/py.typed +5 -0
- segment_everything/stacked_label_dataset.py +113 -0
- segment_everything/stacked_labels.py +428 -0
- segment_everything/vendored/PromptGuidedDecoder/Prompt_guided_Mask_Decoder.pt +0 -0
- segment_everything/vendored/__init__.py +5 -0
- segment_everything/vendored/dice.py +158 -0
- segment_everything/vendored/efficientvit/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/__init__.py +0 -0
- segment_everything/vendored/efficientvit/apps/data_provider/__init__.py +7 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/bbox.py +30 -0
- segment_everything/vendored/efficientvit/apps/data_provider/augment/color_aug.py +78 -0
- segment_everything/vendored/efficientvit/apps/data_provider/base.py +254 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_loader.py +1538 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/_data_worker.py +357 -0
- segment_everything/vendored/efficientvit/apps/data_provider/random_resolution/controller.py +100 -0
- segment_everything/vendored/efficientvit/apps/setup.py +150 -0
- segment_everything/vendored/efficientvit/apps/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/apps/trainer/base.py +318 -0
- segment_everything/vendored/efficientvit/apps/trainer/run_config.py +129 -0
- segment_everything/vendored/efficientvit/apps/utils/__init__.py +12 -0
- segment_everything/vendored/efficientvit/apps/utils/dist.py +32 -0
- segment_everything/vendored/efficientvit/apps/utils/ema.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/export.py +45 -0
- segment_everything/vendored/efficientvit/apps/utils/init.py +66 -0
- segment_everything/vendored/efficientvit/apps/utils/lr.py +52 -0
- segment_everything/vendored/efficientvit/apps/utils/metric.py +43 -0
- segment_everything/vendored/efficientvit/apps/utils/misc.py +101 -0
- segment_everything/vendored/efficientvit/apps/utils/opt.py +28 -0
- segment_everything/vendored/efficientvit/cls_model_zoo.py +79 -0
- segment_everything/vendored/efficientvit/clscore/__init__.py +0 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/__init__.py +5 -0
- segment_everything/vendored/efficientvit/clscore/data_provider/imagenet.py +142 -0
- segment_everything/vendored/efficientvit/clscore/trainer/__init__.py +6 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_run_config.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/cls_trainer.py +265 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/label_smooth.py +18 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/metric.py +23 -0
- segment_everything/vendored/efficientvit/clscore/trainer/utils/mixup.py +67 -0
- segment_everything/vendored/efficientvit/models/__init__.py +0 -0
- segment_everything/vendored/efficientvit/models/efficientvit/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/efficientvit/backbone.py +380 -0
- segment_everything/vendored/efficientvit/models/efficientvit/cls.py +188 -0
- segment_everything/vendored/efficientvit/models/efficientvit/sam.py +181 -0
- segment_everything/vendored/efficientvit/models/efficientvit/seg.py +373 -0
- segment_everything/vendored/efficientvit/models/nn/__init__.py +8 -0
- segment_everything/vendored/efficientvit/models/nn/act.py +30 -0
- segment_everything/vendored/efficientvit/models/nn/drop.py +104 -0
- segment_everything/vendored/efficientvit/models/nn/norm.py +164 -0
- segment_everything/vendored/efficientvit/models/nn/ops.py +597 -0
- segment_everything/vendored/efficientvit/models/utils/__init__.py +7 -0
- segment_everything/vendored/efficientvit/models/utils/list.py +53 -0
- segment_everything/vendored/efficientvit/models/utils/network.py +73 -0
- segment_everything/vendored/efficientvit/models/utils/random.py +65 -0
- segment_everything/vendored/efficientvit/sam_model_zoo.py +45 -0
- segment_everything/vendored/efficientvit/seg_model_zoo.py +70 -0
- segment_everything/vendored/get_object_aware.py +26 -0
- segment_everything/vendored/mobilesamv2/__init__.py +16 -0
- segment_everything/vendored/mobilesamv2/automatic_mask_generator.py +415 -0
- segment_everything/vendored/mobilesamv2/build_sam.py +246 -0
- segment_everything/vendored/mobilesamv2/modeling/__init__.py +11 -0
- segment_everything/vendored/mobilesamv2/modeling/common.py +43 -0
- segment_everything/vendored/mobilesamv2/modeling/image_encoder.py +394 -0
- segment_everything/vendored/mobilesamv2/modeling/mask_decoder.py +213 -0
- segment_everything/vendored/mobilesamv2/modeling/prompt_encoder.py +217 -0
- segment_everything/vendored/mobilesamv2/modeling/sam.py +203 -0
- segment_everything/vendored/mobilesamv2/modeling/transformer.py +240 -0
- segment_everything/vendored/mobilesamv2/predictor.py +384 -0
- segment_everything/vendored/mobilesamv2/utils/__init__.py +5 -0
- segment_everything/vendored/mobilesamv2/utils/amg.py +347 -0
- segment_everything/vendored/mobilesamv2/utils/onnx.py +144 -0
- segment_everything/vendored/mobilesamv2/utils/transforms.py +103 -0
- segment_everything/vendored/object_detection/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/nn/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autobackend.py +658 -0
- segment_everything/vendored/object_detection/ultralytics/nn/autoshape.py +397 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/__init__.py +110 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/block.py +304 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/conv.py +297 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/head.py +468 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/transformer.py +378 -0
- segment_everything/vendored/object_detection/ultralytics/nn/modules/utils.py +78 -0
- segment_everything/vendored/object_detection/ultralytics/nn/tasks.py +1049 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/__init__.py +6 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/model.py +104 -0
- segment_everything/vendored/object_detection/ultralytics/prompt_mobilesamv2/predict.py +95 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/__init__.py +588 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/cfg/default.yaml +117 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/__init__.py +9 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/annotator.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/augment.py +899 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/base.py +286 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/build.py +213 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/converter.py +358 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataloaders/stream_loaders.py +459 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset.py +274 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/dataset_wrappers.py +53 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/download_weights.sh +18 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco.sh +60 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_coco128.sh +17 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/scripts/get_imagenet.sh +51 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/data/utils.py +716 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/__init__.py +0 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/exporter.py +1214 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/model.py +641 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/predictor.py +461 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/engine/results.py +741 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/__init__.py +893 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/autobatch.py +108 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/callbacks/base.py +212 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/checks.py +547 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/dist.py +67 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/downloads.py +353 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/errors.py +12 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/files.py +100 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/instance.py +391 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/loss.py +579 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/metrics.py +1189 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/ops.py +870 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/patches.py +45 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/plotting.py +767 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tal.py +276 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/torch_utils.py +684 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/utils/tuner.py +54 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/__init__.py +5 -0
- segment_everything/vendored/object_detection/ultralytics/yolo/v8/detect/predict.py +69 -0
- segment_everything/vendored/tinyvit/__init__.py +2 -0
- segment_everything/vendored/tinyvit/tiny_vit.py +867 -0
- segment_everything/weights_helper.py +124 -0
- segment_everything-0.1.0.dist-info/METADATA +53 -0
- segment_everything-0.1.0.dist-info/RECORD +145 -0
- segment_everything-0.1.0.dist-info/WHEEL +4 -0
- segment_everything-0.1.0.dist-info/licenses/LICENSE +28 -0
|
@@ -0,0 +1,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
|
+
)
|