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,461 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
"""
|
|
3
|
+
Run prediction on images, videos, directories, globs, YouTube, webcam, streams, etc.
|
|
4
|
+
|
|
5
|
+
Usage - sources:
|
|
6
|
+
$ yolo mode=predict model=yolov8n.pt source=0 # webcam
|
|
7
|
+
img.jpg # image
|
|
8
|
+
vid.mp4 # video
|
|
9
|
+
screen # screenshot
|
|
10
|
+
path/ # directory
|
|
11
|
+
list.txt # list of images
|
|
12
|
+
list.streams # list of streams
|
|
13
|
+
'path/*.jpg' # glob
|
|
14
|
+
'https://youtu.be/Zgi9g1ksQHc' # YouTube
|
|
15
|
+
'rtsp://example.com/media.mp4' # RTSP, RTMP, HTTP stream
|
|
16
|
+
|
|
17
|
+
Usage - formats:
|
|
18
|
+
$ yolo mode=predict model=yolov8n.pt # PyTorch
|
|
19
|
+
yolov8n.torchscript # TorchScript
|
|
20
|
+
yolov8n.onnx # ONNX Runtime or OpenCV DNN with dnn=True
|
|
21
|
+
yolov8n_openvino_model # OpenVINO
|
|
22
|
+
yolov8n.engine # TensorRT
|
|
23
|
+
yolov8n.mlmodel # CoreML (macOS-only)
|
|
24
|
+
yolov8n_saved_model # TensorFlow SavedModel
|
|
25
|
+
yolov8n.pb # TensorFlow GraphDef
|
|
26
|
+
yolov8n.tflite # TensorFlow Lite
|
|
27
|
+
yolov8n_edgetpu.tflite # TensorFlow Edge TPU
|
|
28
|
+
yolov8n_paddle_model # PaddlePaddle
|
|
29
|
+
"""
|
|
30
|
+
import platform
|
|
31
|
+
from pathlib import Path
|
|
32
|
+
|
|
33
|
+
import cv2
|
|
34
|
+
import numpy as np
|
|
35
|
+
import torch
|
|
36
|
+
|
|
37
|
+
from ...nn.autobackend import AutoBackend
|
|
38
|
+
from ..cfg import get_cfg
|
|
39
|
+
from ..data import load_inference_source
|
|
40
|
+
from ..data.augment import LetterBox, classify_transforms
|
|
41
|
+
from ..utils import (
|
|
42
|
+
DEFAULT_CFG,
|
|
43
|
+
LOGGER,
|
|
44
|
+
SETTINGS,
|
|
45
|
+
callbacks,
|
|
46
|
+
colorstr,
|
|
47
|
+
ops,
|
|
48
|
+
)
|
|
49
|
+
from ..utils.checks import check_imgsz, check_imshow
|
|
50
|
+
from ..utils.files import increment_path
|
|
51
|
+
from ..utils.torch_utils import (
|
|
52
|
+
select_device,
|
|
53
|
+
smart_inference_mode,
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
STREAM_WARNING = """
|
|
57
|
+
WARNING ⚠️ stream/video/webcam/dir predict source will accumulate results in RAM unless `stream=True` is passed,
|
|
58
|
+
causing potential out-of-memory errors for large sources or long-running streams/videos.
|
|
59
|
+
|
|
60
|
+
Usage:
|
|
61
|
+
results = model(source=..., stream=True) # generator of Results objects
|
|
62
|
+
for r in results:
|
|
63
|
+
boxes = r.boxes # Boxes object for bbox outputs
|
|
64
|
+
masks = r.masks # Masks object for segment masks outputs
|
|
65
|
+
probs = r.probs # Class probabilities for classification outputs
|
|
66
|
+
"""
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
class BasePredictor:
|
|
70
|
+
"""
|
|
71
|
+
BasePredictor
|
|
72
|
+
|
|
73
|
+
A base class for creating predictors.
|
|
74
|
+
|
|
75
|
+
Attributes:
|
|
76
|
+
args (SimpleNamespace): Configuration for the predictor.
|
|
77
|
+
save_dir (Path): Directory to save results.
|
|
78
|
+
done_warmup (bool): Whether the predictor has finished setup.
|
|
79
|
+
model (nn.Module): Model used for prediction.
|
|
80
|
+
data (dict): Data configuration.
|
|
81
|
+
device (torch.device): Device used for prediction.
|
|
82
|
+
dataset (Dataset): Dataset used for prediction.
|
|
83
|
+
vid_path (str): Path to video file.
|
|
84
|
+
vid_writer (cv2.VideoWriter): Video writer for saving video output.
|
|
85
|
+
data_path (str): Path to data.
|
|
86
|
+
"""
|
|
87
|
+
|
|
88
|
+
def __init__(self, cfg=DEFAULT_CFG, overrides=None, _callbacks=None):
|
|
89
|
+
"""
|
|
90
|
+
Initializes the BasePredictor class.
|
|
91
|
+
|
|
92
|
+
Args:
|
|
93
|
+
cfg (str, optional): Path to a configuration file. Defaults to DEFAULT_CFG.
|
|
94
|
+
overrides (dict, optional): Configuration overrides. Defaults to None.
|
|
95
|
+
"""
|
|
96
|
+
self.args = get_cfg(cfg, overrides)
|
|
97
|
+
self.save_dir = self.get_save_dir()
|
|
98
|
+
if self.args.conf is None:
|
|
99
|
+
self.args.conf = 0.25 # default conf=0.25
|
|
100
|
+
self.done_warmup = False
|
|
101
|
+
if self.args.show:
|
|
102
|
+
self.args.show = check_imshow(warn=True)
|
|
103
|
+
|
|
104
|
+
# Usable if setup is done
|
|
105
|
+
self.model = None
|
|
106
|
+
self.data = self.args.data # data_dict
|
|
107
|
+
self.imgsz = None
|
|
108
|
+
self.device = None
|
|
109
|
+
self.dataset = None
|
|
110
|
+
self.vid_path, self.vid_writer = None, None
|
|
111
|
+
self.plotted_img = None
|
|
112
|
+
self.data_path = None
|
|
113
|
+
self.source_type = None
|
|
114
|
+
self.batch = None
|
|
115
|
+
self.results = None
|
|
116
|
+
self.transforms = None
|
|
117
|
+
self.callbacks = _callbacks or callbacks.get_default_callbacks()
|
|
118
|
+
|
|
119
|
+
def get_save_dir(self):
|
|
120
|
+
project = (
|
|
121
|
+
self.args.project or Path(SETTINGS["runs_dir"]) / self.args.task
|
|
122
|
+
)
|
|
123
|
+
name = self.args.name or f"{self.args.mode}"
|
|
124
|
+
return increment_path(
|
|
125
|
+
Path(project) / name, exist_ok=self.args.exist_ok
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
def preprocess(self, im):
|
|
129
|
+
"""Prepares input image before inference.
|
|
130
|
+
|
|
131
|
+
Args:
|
|
132
|
+
im (torch.Tensor | List(np.ndarray)): (N, 3, h, w) for tensor, [(h, w, 3) x N] for list.
|
|
133
|
+
"""
|
|
134
|
+
if not isinstance(im, torch.Tensor):
|
|
135
|
+
im = np.stack(self.pre_transform(im))
|
|
136
|
+
im = im[..., ::-1].transpose(
|
|
137
|
+
(0, 3, 1, 2)
|
|
138
|
+
) # BGR to RGB, BHWC to BCHW, (n, 3, h, w)
|
|
139
|
+
im = np.ascontiguousarray(im) # contiguous
|
|
140
|
+
im = torch.from_numpy(im)
|
|
141
|
+
# NOTE: assuming im with (b, 3, h, w) if it's a tensor
|
|
142
|
+
img = im.to(self.device)
|
|
143
|
+
img = (
|
|
144
|
+
img.half() if self.model.fp16 else img.float()
|
|
145
|
+
) # uint8 to fp16/32
|
|
146
|
+
img /= 255 # 0 - 255 to 0.0 - 1.0
|
|
147
|
+
return img
|
|
148
|
+
|
|
149
|
+
def pre_transform(self, im):
|
|
150
|
+
"""Pre-tranform input image before inference.
|
|
151
|
+
|
|
152
|
+
Args:
|
|
153
|
+
im (List(np.ndarray)): (N, 3, h, w) for tensor, [(h, w, 3) x N] for list.
|
|
154
|
+
|
|
155
|
+
Return: A list of transformed imgs.
|
|
156
|
+
"""
|
|
157
|
+
same_shapes = all(x.shape == im[0].shape for x in im)
|
|
158
|
+
auto = same_shapes and self.model.pt
|
|
159
|
+
return [
|
|
160
|
+
LetterBox(self.imgsz, auto=auto, stride=self.model.stride)(image=x)
|
|
161
|
+
for x in im
|
|
162
|
+
]
|
|
163
|
+
|
|
164
|
+
def write_results(self, idx, results, batch):
|
|
165
|
+
"""Write inference results to a file or directory."""
|
|
166
|
+
p, im, _ = batch
|
|
167
|
+
log_string = ""
|
|
168
|
+
if len(im.shape) == 3:
|
|
169
|
+
im = im[None] # expand for batch dim
|
|
170
|
+
self.seen += 1
|
|
171
|
+
if (
|
|
172
|
+
self.source_type.webcam or self.source_type.from_img
|
|
173
|
+
): # batch_size >= 1
|
|
174
|
+
log_string += f"{idx}: "
|
|
175
|
+
frame = self.dataset.count
|
|
176
|
+
else:
|
|
177
|
+
frame = getattr(self.dataset, "frame", 0)
|
|
178
|
+
self.data_path = p
|
|
179
|
+
self.txt_path = str(self.save_dir / "labels" / p.stem) + (
|
|
180
|
+
"" if self.dataset.mode == "image" else f"_{frame}"
|
|
181
|
+
)
|
|
182
|
+
log_string += "%gx%g " % im.shape[2:] # print string
|
|
183
|
+
result = results[idx]
|
|
184
|
+
log_string += result.verbose()
|
|
185
|
+
|
|
186
|
+
if self.args.save or self.args.show: # Add bbox to image
|
|
187
|
+
plot_args = dict(
|
|
188
|
+
line_width=self.args.line_width,
|
|
189
|
+
boxes=self.args.boxes,
|
|
190
|
+
conf=self.args.show_conf,
|
|
191
|
+
labels=self.args.show_labels,
|
|
192
|
+
)
|
|
193
|
+
if not self.args.retina_masks:
|
|
194
|
+
plot_args["im_gpu"] = im[idx]
|
|
195
|
+
self.plotted_img = result.plot(**plot_args)
|
|
196
|
+
# Write
|
|
197
|
+
if self.args.save_txt:
|
|
198
|
+
result.save_txt(
|
|
199
|
+
f"{self.txt_path}.txt", save_conf=self.args.save_conf
|
|
200
|
+
)
|
|
201
|
+
if self.args.save_crop:
|
|
202
|
+
result.save_crop(
|
|
203
|
+
save_dir=self.save_dir / "crops", file_name=self.data_path.stem
|
|
204
|
+
)
|
|
205
|
+
|
|
206
|
+
return log_string
|
|
207
|
+
|
|
208
|
+
def postprocess(self, preds, img, orig_imgs):
|
|
209
|
+
"""Post-processes predictions for an image and returns them."""
|
|
210
|
+
return preds
|
|
211
|
+
|
|
212
|
+
def __call__(self, source=None, model=None, stream=False):
|
|
213
|
+
"""Performs inference on an image or stream."""
|
|
214
|
+
self.stream = stream
|
|
215
|
+
if stream:
|
|
216
|
+
return self.stream_inference(source, model)
|
|
217
|
+
else:
|
|
218
|
+
return list(
|
|
219
|
+
self.stream_inference(source, model)
|
|
220
|
+
) # merge list of Result into one
|
|
221
|
+
|
|
222
|
+
def predict_cli(self, source=None, model=None):
|
|
223
|
+
"""Method used for CLI prediction. It uses always generator as outputs as not required by CLI mode."""
|
|
224
|
+
gen = self.stream_inference(source, model)
|
|
225
|
+
for (
|
|
226
|
+
_
|
|
227
|
+
) in (
|
|
228
|
+
gen
|
|
229
|
+
): # running CLI inference without accumulating any outputs (do not modify)
|
|
230
|
+
pass
|
|
231
|
+
|
|
232
|
+
def setup_source(self, source):
|
|
233
|
+
"""Sets up source and inference mode."""
|
|
234
|
+
self.imgsz = check_imgsz(
|
|
235
|
+
self.args.imgsz, stride=self.model.stride, min_dim=2
|
|
236
|
+
) # check image size
|
|
237
|
+
self.transforms = (
|
|
238
|
+
getattr(
|
|
239
|
+
self.model.model,
|
|
240
|
+
"transforms",
|
|
241
|
+
classify_transforms(self.imgsz[0]),
|
|
242
|
+
)
|
|
243
|
+
if self.args.task == "classify"
|
|
244
|
+
else None
|
|
245
|
+
)
|
|
246
|
+
self.dataset = load_inference_source(
|
|
247
|
+
source=source, imgsz=self.imgsz, vid_stride=self.args.vid_stride
|
|
248
|
+
)
|
|
249
|
+
self.source_type = self.dataset.source_type
|
|
250
|
+
if not getattr(self, "stream", True) and (
|
|
251
|
+
self.dataset.mode == "stream"
|
|
252
|
+
or len(self.dataset) > 1000 # streams
|
|
253
|
+
or any(getattr(self.dataset, "video_flag", [False])) # images
|
|
254
|
+
): # videos
|
|
255
|
+
LOGGER.warning(STREAM_WARNING)
|
|
256
|
+
self.vid_path, self.vid_writer = [None] * self.dataset.bs, [
|
|
257
|
+
None
|
|
258
|
+
] * self.dataset.bs
|
|
259
|
+
|
|
260
|
+
@smart_inference_mode()
|
|
261
|
+
def stream_inference(self, source=None, model=None):
|
|
262
|
+
"""Streams real-time inference on camera feed and saves results to file."""
|
|
263
|
+
if self.args.verbose:
|
|
264
|
+
LOGGER.info("")
|
|
265
|
+
|
|
266
|
+
# Setup model
|
|
267
|
+
if not self.model:
|
|
268
|
+
self.setup_model(model)
|
|
269
|
+
# Setup source every time predict is called
|
|
270
|
+
self.setup_source(source if source is not None else self.args.source)
|
|
271
|
+
|
|
272
|
+
# Check if save_dir/ label file exists
|
|
273
|
+
if self.args.save or self.args.save_txt:
|
|
274
|
+
(
|
|
275
|
+
self.save_dir / "labels"
|
|
276
|
+
if self.args.save_txt
|
|
277
|
+
else self.save_dir
|
|
278
|
+
).mkdir(parents=True, exist_ok=True)
|
|
279
|
+
# Warmup model
|
|
280
|
+
if not self.done_warmup:
|
|
281
|
+
self.model.warmup(
|
|
282
|
+
imgsz=(
|
|
283
|
+
(
|
|
284
|
+
1
|
|
285
|
+
if self.model.pt or self.model.triton
|
|
286
|
+
else self.dataset.bs
|
|
287
|
+
),
|
|
288
|
+
3,
|
|
289
|
+
*self.imgsz,
|
|
290
|
+
)
|
|
291
|
+
)
|
|
292
|
+
self.done_warmup = True
|
|
293
|
+
|
|
294
|
+
self.seen, self.windows, self.batch, profilers = (
|
|
295
|
+
0,
|
|
296
|
+
[],
|
|
297
|
+
None,
|
|
298
|
+
(ops.Profile(), ops.Profile(), ops.Profile()),
|
|
299
|
+
)
|
|
300
|
+
self.run_callbacks("on_predict_start")
|
|
301
|
+
for batch in self.dataset:
|
|
302
|
+
self.run_callbacks("on_predict_batch_start")
|
|
303
|
+
self.batch = batch
|
|
304
|
+
path, im0s, vid_cap, s = batch
|
|
305
|
+
visualize = (
|
|
306
|
+
increment_path(self.save_dir / Path(path[0]).stem, mkdir=True)
|
|
307
|
+
if self.args.visualize and (not self.source_type.tensor)
|
|
308
|
+
else False
|
|
309
|
+
)
|
|
310
|
+
|
|
311
|
+
# Preprocess
|
|
312
|
+
with profilers[0]:
|
|
313
|
+
im = self.preprocess(im0s)
|
|
314
|
+
|
|
315
|
+
# Inference
|
|
316
|
+
with profilers[1]:
|
|
317
|
+
preds = self.model(
|
|
318
|
+
im, augment=self.args.augment, visualize=visualize
|
|
319
|
+
)
|
|
320
|
+
|
|
321
|
+
# Postprocess
|
|
322
|
+
with profilers[2]:
|
|
323
|
+
self.results = self.postprocess(preds, im, im0s)
|
|
324
|
+
self.run_callbacks("on_predict_postprocess_end")
|
|
325
|
+
|
|
326
|
+
# Visualize, save, write results
|
|
327
|
+
n = len(im0s)
|
|
328
|
+
for i in range(n):
|
|
329
|
+
self.results[i].speed = {
|
|
330
|
+
"preprocess": profilers[0].dt * 1e3 / n,
|
|
331
|
+
"inference": profilers[1].dt * 1e3 / n,
|
|
332
|
+
"postprocess": profilers[2].dt * 1e3 / n,
|
|
333
|
+
}
|
|
334
|
+
if (
|
|
335
|
+
self.source_type.tensor
|
|
336
|
+
): # skip write, show and plot operations if input is raw tensor
|
|
337
|
+
continue
|
|
338
|
+
p, im0 = path[i], im0s[i].copy()
|
|
339
|
+
p = Path(p)
|
|
340
|
+
|
|
341
|
+
if (
|
|
342
|
+
self.args.verbose
|
|
343
|
+
or self.args.save
|
|
344
|
+
or self.args.save_txt
|
|
345
|
+
or self.args.show
|
|
346
|
+
):
|
|
347
|
+
s += self.write_results(i, self.results, (p, im, im0))
|
|
348
|
+
if self.args.save or self.args.save_txt:
|
|
349
|
+
self.results[i].save_dir = self.save_dir.__str__()
|
|
350
|
+
if self.args.show and self.plotted_img is not None:
|
|
351
|
+
self.show(p)
|
|
352
|
+
if self.args.save and self.plotted_img is not None:
|
|
353
|
+
self.save_preds(vid_cap, i, str(self.save_dir / p.name))
|
|
354
|
+
|
|
355
|
+
self.run_callbacks("on_predict_batch_end")
|
|
356
|
+
yield from self.results
|
|
357
|
+
|
|
358
|
+
# Print time (inference-only)
|
|
359
|
+
if self.args.verbose:
|
|
360
|
+
LOGGER.info(f"{s}{profilers[1].dt * 1E3:.1f}ms")
|
|
361
|
+
|
|
362
|
+
# Release assets
|
|
363
|
+
if isinstance(self.vid_writer[-1], cv2.VideoWriter):
|
|
364
|
+
self.vid_writer[-1].release() # release final video writer
|
|
365
|
+
|
|
366
|
+
# Print results
|
|
367
|
+
if self.args.verbose and self.seen:
|
|
368
|
+
t = tuple(
|
|
369
|
+
x.t / self.seen * 1e3 for x in profilers
|
|
370
|
+
) # speeds per image
|
|
371
|
+
LOGGER.info(
|
|
372
|
+
f"Speed: %.1fms preprocess, %.1fms inference, %.1fms postprocess per image at shape "
|
|
373
|
+
f"{(1, 3, *self.imgsz)}" % t
|
|
374
|
+
)
|
|
375
|
+
if self.args.save or self.args.save_txt or self.args.save_crop:
|
|
376
|
+
nl = len(
|
|
377
|
+
list(self.save_dir.glob("labels/*.txt"))
|
|
378
|
+
) # number of labels
|
|
379
|
+
s = (
|
|
380
|
+
f"\n{nl} label{'s' * (nl > 1)} saved to {self.save_dir / 'labels'}"
|
|
381
|
+
if self.args.save_txt
|
|
382
|
+
else ""
|
|
383
|
+
)
|
|
384
|
+
LOGGER.info(
|
|
385
|
+
f"Results saved to {colorstr('bold', self.save_dir)}{s}"
|
|
386
|
+
)
|
|
387
|
+
|
|
388
|
+
self.run_callbacks("on_predict_end")
|
|
389
|
+
|
|
390
|
+
def setup_model(self, model, verbose=True):
|
|
391
|
+
"""Initialize YOLO model with given parameters and set it to evaluation mode."""
|
|
392
|
+
device = select_device(self.args.device, verbose=verbose)
|
|
393
|
+
model = model or self.args.model
|
|
394
|
+
self.args.half &= (
|
|
395
|
+
device.type != "cpu"
|
|
396
|
+
) # half precision only supported on CUDA
|
|
397
|
+
self.model = AutoBackend(
|
|
398
|
+
model,
|
|
399
|
+
device=device,
|
|
400
|
+
dnn=self.args.dnn,
|
|
401
|
+
data=self.args.data,
|
|
402
|
+
fp16=self.args.half,
|
|
403
|
+
fuse=True,
|
|
404
|
+
verbose=verbose,
|
|
405
|
+
)
|
|
406
|
+
self.device = device
|
|
407
|
+
self.model.eval()
|
|
408
|
+
|
|
409
|
+
def show(self, p):
|
|
410
|
+
"""Display an image in a window using OpenCV imshow()."""
|
|
411
|
+
im0 = self.plotted_img
|
|
412
|
+
if platform.system() == "Linux" and p not in self.windows:
|
|
413
|
+
self.windows.append(p)
|
|
414
|
+
cv2.namedWindow(
|
|
415
|
+
str(p), cv2.WINDOW_NORMAL | cv2.WINDOW_KEEPRATIO
|
|
416
|
+
) # allow window resize (Linux)
|
|
417
|
+
cv2.resizeWindow(str(p), im0.shape[1], im0.shape[0])
|
|
418
|
+
cv2.imshow(str(p), im0)
|
|
419
|
+
cv2.waitKey(
|
|
420
|
+
500 if self.batch[3].startswith("image") else 1
|
|
421
|
+
) # 1 millisecond
|
|
422
|
+
|
|
423
|
+
def save_preds(self, vid_cap, idx, save_path):
|
|
424
|
+
"""Save video predictions as mp4 at specified path."""
|
|
425
|
+
im0 = self.plotted_img
|
|
426
|
+
# Save imgs
|
|
427
|
+
if self.dataset.mode == "image":
|
|
428
|
+
cv2.imwrite(save_path, im0)
|
|
429
|
+
else: # 'video' or 'stream'
|
|
430
|
+
if self.vid_path[idx] != save_path: # new video
|
|
431
|
+
self.vid_path[idx] = save_path
|
|
432
|
+
if isinstance(self.vid_writer[idx], cv2.VideoWriter):
|
|
433
|
+
self.vid_writer[
|
|
434
|
+
idx
|
|
435
|
+
].release() # release previous video writer
|
|
436
|
+
if vid_cap: # video
|
|
437
|
+
fps = int(
|
|
438
|
+
vid_cap.get(cv2.CAP_PROP_FPS)
|
|
439
|
+
) # integer required, floats produce error in MP4 codec
|
|
440
|
+
w = int(vid_cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
|
441
|
+
h = int(vid_cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
|
442
|
+
else: # stream
|
|
443
|
+
fps, w, h = 30, im0.shape[1], im0.shape[0]
|
|
444
|
+
save_path = str(
|
|
445
|
+
Path(save_path).with_suffix(".mp4")
|
|
446
|
+
) # force *.mp4 suffix on results videos
|
|
447
|
+
self.vid_writer[idx] = cv2.VideoWriter(
|
|
448
|
+
save_path, cv2.VideoWriter_fourcc(*"mp4v"), fps, (w, h)
|
|
449
|
+
)
|
|
450
|
+
self.vid_writer[idx].write(im0)
|
|
451
|
+
|
|
452
|
+
def run_callbacks(self, event: str):
|
|
453
|
+
"""Runs all registered callbacks for a specific event."""
|
|
454
|
+
for callback in self.callbacks.get(event, []):
|
|
455
|
+
callback(self)
|
|
456
|
+
|
|
457
|
+
def add_callback(self, event: str, func):
|
|
458
|
+
"""
|
|
459
|
+
Add callback
|
|
460
|
+
"""
|
|
461
|
+
self.callbacks[event].append(func)
|