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,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)