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,641 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
|
|
3
|
+
import sys
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
from typing import Union
|
|
6
|
+
|
|
7
|
+
from ... import yolo # noqa
|
|
8
|
+
from ...nn.tasks import (
|
|
9
|
+
DetectionModel,
|
|
10
|
+
attempt_load_one_weight,
|
|
11
|
+
guess_model_task,
|
|
12
|
+
nn,
|
|
13
|
+
yaml_model_load,
|
|
14
|
+
)
|
|
15
|
+
from ..cfg import get_cfg
|
|
16
|
+
from .exporter import Exporter
|
|
17
|
+
from ..utils import (
|
|
18
|
+
DEFAULT_CFG,
|
|
19
|
+
DEFAULT_CFG_DICT,
|
|
20
|
+
DEFAULT_CFG_KEYS,
|
|
21
|
+
LOGGER,
|
|
22
|
+
NUM_THREADS,
|
|
23
|
+
RANK,
|
|
24
|
+
ROOT,
|
|
25
|
+
callbacks,
|
|
26
|
+
is_git_dir,
|
|
27
|
+
yaml_load,
|
|
28
|
+
)
|
|
29
|
+
from ..utils.checks import (
|
|
30
|
+
check_file,
|
|
31
|
+
check_imgsz,
|
|
32
|
+
check_pip_update_available,
|
|
33
|
+
check_yaml,
|
|
34
|
+
)
|
|
35
|
+
from ..utils.downloads import GITHUB_ASSET_STEMS
|
|
36
|
+
from ..utils.torch_utils import smart_inference_mode
|
|
37
|
+
|
|
38
|
+
# Map head to model, trainer, validator, and predictor classes
|
|
39
|
+
TASK_MAP = {
|
|
40
|
+
"detect": [
|
|
41
|
+
DetectionModel,
|
|
42
|
+
yolo.v8.detect.DetectionPredictor,
|
|
43
|
+
]
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class YOLO:
|
|
48
|
+
"""
|
|
49
|
+
YOLO (You Only Look Once) object detection model.
|
|
50
|
+
|
|
51
|
+
Args:
|
|
52
|
+
model (str, Path): Path to the model file to load or create.
|
|
53
|
+
task (Any, optional): Task type for the YOLO model. Defaults to None.
|
|
54
|
+
|
|
55
|
+
Attributes:
|
|
56
|
+
predictor (Any): The predictor object.
|
|
57
|
+
model (Any): The model object.
|
|
58
|
+
trainer (Any): The trainer object.
|
|
59
|
+
task (str): The type of model task.
|
|
60
|
+
ckpt (Any): The checkpoint object if the model loaded from *.pt file.
|
|
61
|
+
cfg (str): The model configuration if loaded from *.yaml file.
|
|
62
|
+
ckpt_path (str): The checkpoint file path.
|
|
63
|
+
overrides (dict): Overrides for the trainer object.
|
|
64
|
+
metrics (Any): The data for metrics.
|
|
65
|
+
|
|
66
|
+
Methods:
|
|
67
|
+
__call__(source=None, stream=False, **kwargs):
|
|
68
|
+
Alias for the predict method.
|
|
69
|
+
_new(cfg:str, verbose:bool=True) -> None:
|
|
70
|
+
Initializes a new model and infers the task type from the model definitions.
|
|
71
|
+
_load(weights:str, task:str='') -> None:
|
|
72
|
+
Initializes a new model and infers the task type from the model head.
|
|
73
|
+
_check_is_pytorch_model() -> None:
|
|
74
|
+
Raises TypeError if the model is not a PyTorch model.
|
|
75
|
+
reset() -> None:
|
|
76
|
+
Resets the model modules.
|
|
77
|
+
info(verbose:bool=False) -> None:
|
|
78
|
+
Logs the model info.
|
|
79
|
+
fuse() -> None:
|
|
80
|
+
Fuses the model for faster inference.
|
|
81
|
+
predict(source=None, stream=False, **kwargs) -> List[ultralytics.yolo.engine.results.Results]:
|
|
82
|
+
Performs prediction using the YOLO model.
|
|
83
|
+
|
|
84
|
+
Returns:
|
|
85
|
+
list(ultralytics.yolo.engine.results.Results): The prediction results.
|
|
86
|
+
"""
|
|
87
|
+
|
|
88
|
+
def __init__(
|
|
89
|
+
self, model: Union[str, Path] = "yolov8n.pt", task=None
|
|
90
|
+
) -> None:
|
|
91
|
+
"""
|
|
92
|
+
Initializes the YOLO model.
|
|
93
|
+
|
|
94
|
+
Args:
|
|
95
|
+
model (Union[str, Path], optional): Path or name of the model to load or create. Defaults to 'yolov8n.pt'.
|
|
96
|
+
task (Any, optional): Task type for the YOLO model. Defaults to None.
|
|
97
|
+
"""
|
|
98
|
+
self.callbacks = callbacks.get_default_callbacks()
|
|
99
|
+
self.predictor = None # reuse predictor
|
|
100
|
+
self.model = None # model object
|
|
101
|
+
self.trainer = None # trainer object
|
|
102
|
+
self.task = None # task type
|
|
103
|
+
self.ckpt = None # if loaded from *.pt
|
|
104
|
+
self.cfg = None # if loaded from *.yaml
|
|
105
|
+
self.ckpt_path = None
|
|
106
|
+
self.overrides = {} # overrides for trainer object
|
|
107
|
+
self.metrics = None # validation/training metrics
|
|
108
|
+
self.session = None # HUB session
|
|
109
|
+
model = str(model).strip() # strip spaces
|
|
110
|
+
|
|
111
|
+
# Check if Ultralytics HUB model from https://hub.ultralytics.com
|
|
112
|
+
if self.is_hub_model(model):
|
|
113
|
+
from ultralytics.hub.session import HUBTrainingSession
|
|
114
|
+
|
|
115
|
+
self.session = HUBTrainingSession(model)
|
|
116
|
+
model = self.session.model_file
|
|
117
|
+
|
|
118
|
+
# Load or create new YOLO model
|
|
119
|
+
suffix = Path(model).suffix
|
|
120
|
+
if not suffix and Path(model).stem in GITHUB_ASSET_STEMS:
|
|
121
|
+
model, suffix = (
|
|
122
|
+
Path(model).with_suffix(".pt"),
|
|
123
|
+
".pt",
|
|
124
|
+
) # add suffix, i.e. yolov8n -> yolov8n.pt
|
|
125
|
+
if suffix == ".yaml":
|
|
126
|
+
self._new(model, task)
|
|
127
|
+
else:
|
|
128
|
+
self._load(model, task)
|
|
129
|
+
|
|
130
|
+
def __call__(self, source=None, stream=False, **kwargs):
|
|
131
|
+
"""Calls the 'predict' function with given arguments to perform object detection."""
|
|
132
|
+
return self.predict(source, stream, **kwargs)
|
|
133
|
+
|
|
134
|
+
def __getattr__(self, attr):
|
|
135
|
+
"""Raises error if object has no requested attribute."""
|
|
136
|
+
name = self.__class__.__name__
|
|
137
|
+
raise AttributeError(
|
|
138
|
+
f"'{name}' object has no attribute '{attr}'. See valid attributes below.\n{self.__doc__}"
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
@staticmethod
|
|
142
|
+
def is_hub_model(model):
|
|
143
|
+
"""Check if the provided model is a HUB model."""
|
|
144
|
+
return any(
|
|
145
|
+
(
|
|
146
|
+
model.startswith(
|
|
147
|
+
"https://hub.ultralytics.com/models/"
|
|
148
|
+
), # i.e. https://hub.ultralytics.com/models/MODEL_ID
|
|
149
|
+
[len(x) for x in model.split("_")]
|
|
150
|
+
== [42, 20], # APIKEY_MODELID
|
|
151
|
+
len(model) == 20
|
|
152
|
+
and not Path(model).exists()
|
|
153
|
+
and all(x not in model for x in "./\\"),
|
|
154
|
+
)
|
|
155
|
+
) # MODELID
|
|
156
|
+
|
|
157
|
+
def _new(self, cfg: str, task=None, verbose=True):
|
|
158
|
+
"""
|
|
159
|
+
Initializes a new model and infers the task type from the model definitions.
|
|
160
|
+
|
|
161
|
+
Args:
|
|
162
|
+
cfg (str): model configuration file
|
|
163
|
+
task (str | None): model task
|
|
164
|
+
verbose (bool): display model info on load
|
|
165
|
+
"""
|
|
166
|
+
cfg_dict = yaml_model_load(cfg)
|
|
167
|
+
self.cfg = cfg
|
|
168
|
+
self.task = task or guess_model_task(cfg_dict)
|
|
169
|
+
self.model = TASK_MAP[self.task][0](
|
|
170
|
+
cfg_dict, verbose=verbose and RANK == -1
|
|
171
|
+
) # build model
|
|
172
|
+
self.overrides["model"] = self.cfg
|
|
173
|
+
|
|
174
|
+
# Below added to allow export from yamls
|
|
175
|
+
args = {
|
|
176
|
+
**DEFAULT_CFG_DICT,
|
|
177
|
+
**self.overrides,
|
|
178
|
+
} # combine model and default args, preferring model args
|
|
179
|
+
self.model.args = {
|
|
180
|
+
k: v for k, v in args.items() if k in DEFAULT_CFG_KEYS
|
|
181
|
+
} # attach args to model
|
|
182
|
+
self.model.task = self.task
|
|
183
|
+
|
|
184
|
+
def _load(self, weights: str, task=None):
|
|
185
|
+
"""
|
|
186
|
+
Initializes a new model and infers the task type from the model head.
|
|
187
|
+
|
|
188
|
+
Args:
|
|
189
|
+
weights (str): model checkpoint to be loaded
|
|
190
|
+
task (str | None): model task
|
|
191
|
+
"""
|
|
192
|
+
suffix = Path(weights).suffix
|
|
193
|
+
if suffix == ".pt":
|
|
194
|
+
self.model, self.ckpt = attempt_load_one_weight(weights)
|
|
195
|
+
self.task = self.model.args["task"]
|
|
196
|
+
self.overrides = self.model.args = self._reset_ckpt_args(
|
|
197
|
+
self.model.args
|
|
198
|
+
)
|
|
199
|
+
self.ckpt_path = self.model.pt_path
|
|
200
|
+
else:
|
|
201
|
+
weights = check_file(weights)
|
|
202
|
+
self.model, self.ckpt = weights, None
|
|
203
|
+
self.task = task or guess_model_task(weights)
|
|
204
|
+
self.ckpt_path = weights
|
|
205
|
+
self.overrides["model"] = weights
|
|
206
|
+
self.overrides["task"] = self.task
|
|
207
|
+
|
|
208
|
+
def _check_is_pytorch_model(self):
|
|
209
|
+
"""
|
|
210
|
+
Raises TypeError is model is not a PyTorch model
|
|
211
|
+
"""
|
|
212
|
+
pt_str = (
|
|
213
|
+
isinstance(self.model, (str, Path))
|
|
214
|
+
and Path(self.model).suffix == ".pt"
|
|
215
|
+
)
|
|
216
|
+
pt_module = isinstance(self.model, nn.Module)
|
|
217
|
+
if not (pt_module or pt_str):
|
|
218
|
+
raise TypeError(
|
|
219
|
+
f"model='{self.model}' must be a *.pt PyTorch model, but is a different type. "
|
|
220
|
+
f"PyTorch models can be used to train, val, predict and export, i.e. "
|
|
221
|
+
f"'yolo export model=yolov8n.pt', but exported formats like ONNX, TensorRT etc. only "
|
|
222
|
+
f"support 'predict' and 'val' modes, i.e. 'yolo predict model=yolov8n.onnx'."
|
|
223
|
+
)
|
|
224
|
+
|
|
225
|
+
@smart_inference_mode()
|
|
226
|
+
def reset_weights(self):
|
|
227
|
+
"""
|
|
228
|
+
Resets the model modules parameters to randomly initialized values, losing all training information.
|
|
229
|
+
"""
|
|
230
|
+
self._check_is_pytorch_model()
|
|
231
|
+
for m in self.model.modules():
|
|
232
|
+
if hasattr(m, "reset_parameters"):
|
|
233
|
+
m.reset_parameters()
|
|
234
|
+
for p in self.model.parameters():
|
|
235
|
+
p.requires_grad = True
|
|
236
|
+
return self
|
|
237
|
+
|
|
238
|
+
@smart_inference_mode()
|
|
239
|
+
def load(self, weights="yolov8n.pt"):
|
|
240
|
+
"""
|
|
241
|
+
Transfers parameters with matching names and shapes from 'weights' to model.
|
|
242
|
+
"""
|
|
243
|
+
self._check_is_pytorch_model()
|
|
244
|
+
if isinstance(weights, (str, Path)):
|
|
245
|
+
weights, self.ckpt = attempt_load_one_weight(weights)
|
|
246
|
+
self.model.load(weights)
|
|
247
|
+
return self
|
|
248
|
+
|
|
249
|
+
def info(self, detailed=False, verbose=True):
|
|
250
|
+
"""
|
|
251
|
+
Logs model info.
|
|
252
|
+
|
|
253
|
+
Args:
|
|
254
|
+
detailed (bool): Show detailed information about model.
|
|
255
|
+
verbose (bool): Controls verbosity.
|
|
256
|
+
"""
|
|
257
|
+
self._check_is_pytorch_model()
|
|
258
|
+
return self.model.info(detailed=detailed, verbose=verbose)
|
|
259
|
+
|
|
260
|
+
def fuse(self):
|
|
261
|
+
"""Fuse PyTorch Conv2d and BatchNorm2d layers."""
|
|
262
|
+
self._check_is_pytorch_model()
|
|
263
|
+
self.model.fuse()
|
|
264
|
+
|
|
265
|
+
@smart_inference_mode()
|
|
266
|
+
def predict(self, source=None, stream=False, **kwargs):
|
|
267
|
+
"""
|
|
268
|
+
Perform prediction using the YOLO model.
|
|
269
|
+
|
|
270
|
+
Args:
|
|
271
|
+
source (str | int | PIL | np.ndarray): The source of the image to make predictions on.
|
|
272
|
+
Accepts all source types accepted by the YOLO model.
|
|
273
|
+
stream (bool): Whether to stream the predictions or not. Defaults to False.
|
|
274
|
+
**kwargs : Additional keyword arguments passed to the predictor.
|
|
275
|
+
Check the 'configuration' section in the documentation for all available options.
|
|
276
|
+
|
|
277
|
+
Returns:
|
|
278
|
+
(List[ultralytics.yolo.engine.results.Results]): The prediction results.
|
|
279
|
+
"""
|
|
280
|
+
if source is None:
|
|
281
|
+
source = (
|
|
282
|
+
ROOT / "assets"
|
|
283
|
+
if is_git_dir()
|
|
284
|
+
else "https://ultralytics.com/images/bus.jpg"
|
|
285
|
+
)
|
|
286
|
+
LOGGER.warning(
|
|
287
|
+
f"WARNING ⚠️ 'source' is missing. Using 'source={source}'."
|
|
288
|
+
)
|
|
289
|
+
is_cli = (
|
|
290
|
+
sys.argv[0].endswith("yolo") or sys.argv[0].endswith("ultralytics")
|
|
291
|
+
) and any(
|
|
292
|
+
x in sys.argv
|
|
293
|
+
for x in ("predict", "track", "mode=predict", "mode=track")
|
|
294
|
+
)
|
|
295
|
+
overrides = self.overrides.copy()
|
|
296
|
+
overrides["conf"] = 0.25
|
|
297
|
+
overrides.update(kwargs) # prefer kwargs
|
|
298
|
+
overrides["mode"] = kwargs.get("mode", "predict")
|
|
299
|
+
assert overrides["mode"] in ["track", "predict"]
|
|
300
|
+
if not is_cli:
|
|
301
|
+
overrides["save"] = kwargs.get(
|
|
302
|
+
"save", False
|
|
303
|
+
) # do not save by default if called in Python
|
|
304
|
+
if not self.predictor:
|
|
305
|
+
self.task = overrides.get("task") or self.task
|
|
306
|
+
self.predictor = TASK_MAP[self.task][3](
|
|
307
|
+
overrides=overrides, _callbacks=self.callbacks
|
|
308
|
+
)
|
|
309
|
+
self.predictor.setup_model(model=self.model, verbose=is_cli)
|
|
310
|
+
else: # only update args if predictor is already setup
|
|
311
|
+
self.predictor.args = get_cfg(self.predictor.args, overrides)
|
|
312
|
+
if "project" in overrides or "name" in overrides:
|
|
313
|
+
self.predictor.save_dir = self.predictor.get_save_dir()
|
|
314
|
+
return (
|
|
315
|
+
self.predictor.predict_cli(source=source)
|
|
316
|
+
if is_cli
|
|
317
|
+
else self.predictor(source=source, stream=stream)
|
|
318
|
+
)
|
|
319
|
+
|
|
320
|
+
def track(self, source=None, stream=False, persist=False, **kwargs):
|
|
321
|
+
"""
|
|
322
|
+
Perform object tracking on the input source using the registered trackers.
|
|
323
|
+
|
|
324
|
+
Args:
|
|
325
|
+
source (str, optional): The input source for object tracking. Can be a file path or a video stream.
|
|
326
|
+
stream (bool, optional): Whether the input source is a video stream. Defaults to False.
|
|
327
|
+
persist (bool, optional): Whether to persist the trackers if they already exist. Defaults to False.
|
|
328
|
+
**kwargs (optional): Additional keyword arguments for the tracking process.
|
|
329
|
+
|
|
330
|
+
Returns:
|
|
331
|
+
(List[ultralytics.yolo.engine.results.Results]): The tracking results.
|
|
332
|
+
|
|
333
|
+
"""
|
|
334
|
+
if not hasattr(self.predictor, "trackers"):
|
|
335
|
+
from ultralytics.tracker import register_tracker
|
|
336
|
+
|
|
337
|
+
register_tracker(self, persist)
|
|
338
|
+
# ByteTrack-based method needs low confidence predictions as input
|
|
339
|
+
conf = kwargs.get("conf") or 0.1
|
|
340
|
+
kwargs["conf"] = conf
|
|
341
|
+
kwargs["mode"] = "track"
|
|
342
|
+
return self.predict(source=source, stream=stream, **kwargs)
|
|
343
|
+
|
|
344
|
+
@smart_inference_mode()
|
|
345
|
+
def val(self, data=None, **kwargs):
|
|
346
|
+
"""
|
|
347
|
+
Validate a model on a given dataset.
|
|
348
|
+
|
|
349
|
+
Args:
|
|
350
|
+
data (str): The dataset to validate on. Accepts all formats accepted by yolo
|
|
351
|
+
**kwargs : Any other args accepted by the validators. To see all args check 'configuration' section in docs
|
|
352
|
+
"""
|
|
353
|
+
overrides = self.overrides.copy()
|
|
354
|
+
overrides["rect"] = True # rect batches as default
|
|
355
|
+
overrides.update(kwargs)
|
|
356
|
+
overrides["mode"] = "val"
|
|
357
|
+
args = get_cfg(cfg=DEFAULT_CFG, overrides=overrides)
|
|
358
|
+
args.data = data or args.data
|
|
359
|
+
if "task" in overrides:
|
|
360
|
+
self.task = args.task
|
|
361
|
+
else:
|
|
362
|
+
args.task = self.task
|
|
363
|
+
if args.imgsz == DEFAULT_CFG.imgsz and not isinstance(
|
|
364
|
+
self.model, (str, Path)
|
|
365
|
+
):
|
|
366
|
+
args.imgsz = self.model.args[
|
|
367
|
+
"imgsz"
|
|
368
|
+
] # use trained imgsz unless custom value is passed
|
|
369
|
+
args.imgsz = check_imgsz(args.imgsz, max_dim=1)
|
|
370
|
+
|
|
371
|
+
validator = TASK_MAP[self.task][2](
|
|
372
|
+
args=args, _callbacks=self.callbacks
|
|
373
|
+
)
|
|
374
|
+
validator(model=self.model)
|
|
375
|
+
self.metrics = validator.metrics
|
|
376
|
+
|
|
377
|
+
return validator.metrics
|
|
378
|
+
|
|
379
|
+
@smart_inference_mode()
|
|
380
|
+
def benchmark(self, **kwargs):
|
|
381
|
+
"""
|
|
382
|
+
Benchmark a model on all export formats.
|
|
383
|
+
|
|
384
|
+
Args:
|
|
385
|
+
**kwargs : Any other args accepted by the validators. To see all args check 'configuration' section in docs
|
|
386
|
+
"""
|
|
387
|
+
self._check_is_pytorch_model()
|
|
388
|
+
from ultralytics.yolo.utils.benchmarks import benchmark
|
|
389
|
+
|
|
390
|
+
overrides = self.model.args.copy()
|
|
391
|
+
overrides.update(kwargs)
|
|
392
|
+
overrides["mode"] = "benchmark"
|
|
393
|
+
overrides = {
|
|
394
|
+
**DEFAULT_CFG_DICT,
|
|
395
|
+
**overrides,
|
|
396
|
+
} # fill in missing overrides keys with defaults
|
|
397
|
+
return benchmark(
|
|
398
|
+
model=self,
|
|
399
|
+
imgsz=overrides["imgsz"],
|
|
400
|
+
half=overrides["half"],
|
|
401
|
+
device=overrides["device"],
|
|
402
|
+
)
|
|
403
|
+
|
|
404
|
+
def export(self, **kwargs):
|
|
405
|
+
"""
|
|
406
|
+
Export model.
|
|
407
|
+
|
|
408
|
+
Args:
|
|
409
|
+
**kwargs : Any other args accepted by the predictors. To see all args check 'configuration' section in docs
|
|
410
|
+
"""
|
|
411
|
+
self._check_is_pytorch_model()
|
|
412
|
+
overrides = self.overrides.copy()
|
|
413
|
+
overrides.update(kwargs)
|
|
414
|
+
overrides["mode"] = "export"
|
|
415
|
+
if overrides.get("imgsz") is None:
|
|
416
|
+
overrides["imgsz"] = self.model.args[
|
|
417
|
+
"imgsz"
|
|
418
|
+
] # use trained imgsz unless custom value is passed
|
|
419
|
+
if "batch" not in kwargs:
|
|
420
|
+
overrides["batch"] = 1 # default to 1 if not modified
|
|
421
|
+
args = get_cfg(cfg=DEFAULT_CFG, overrides=overrides)
|
|
422
|
+
args.task = self.task
|
|
423
|
+
return Exporter(overrides=args, _callbacks=self.callbacks)(
|
|
424
|
+
model=self.model
|
|
425
|
+
)
|
|
426
|
+
|
|
427
|
+
def train(self, **kwargs):
|
|
428
|
+
"""
|
|
429
|
+
Trains the model on a given dataset.
|
|
430
|
+
|
|
431
|
+
Args:
|
|
432
|
+
**kwargs (Any): Any number of arguments representing the training configuration.
|
|
433
|
+
"""
|
|
434
|
+
self._check_is_pytorch_model()
|
|
435
|
+
if self.session: # Ultralytics HUB session
|
|
436
|
+
if any(kwargs):
|
|
437
|
+
LOGGER.warning(
|
|
438
|
+
"WARNING ⚠️ using HUB training arguments, ignoring local training arguments."
|
|
439
|
+
)
|
|
440
|
+
kwargs = self.session.train_args
|
|
441
|
+
check_pip_update_available()
|
|
442
|
+
overrides = self.overrides.copy()
|
|
443
|
+
if kwargs.get("cfg"):
|
|
444
|
+
LOGGER.info(
|
|
445
|
+
f"cfg file passed. Overriding default params with {kwargs['cfg']}."
|
|
446
|
+
)
|
|
447
|
+
overrides = yaml_load(check_yaml(kwargs["cfg"]))
|
|
448
|
+
overrides.update(kwargs)
|
|
449
|
+
overrides["mode"] = "train"
|
|
450
|
+
if not overrides.get("data"):
|
|
451
|
+
raise AttributeError(
|
|
452
|
+
"Dataset required but missing, i.e. pass 'data=coco128.yaml'"
|
|
453
|
+
)
|
|
454
|
+
if overrides.get("resume"):
|
|
455
|
+
overrides["resume"] = self.ckpt_path
|
|
456
|
+
self.task = overrides.get("task") or self.task
|
|
457
|
+
self.trainer = TASK_MAP[self.task][1](
|
|
458
|
+
overrides=overrides, _callbacks=self.callbacks
|
|
459
|
+
)
|
|
460
|
+
if not overrides.get(
|
|
461
|
+
"resume"
|
|
462
|
+
): # manually set model only if not resuming
|
|
463
|
+
self.trainer.model = self.trainer.get_model(
|
|
464
|
+
weights=self.model if self.ckpt else None, cfg=self.model.yaml
|
|
465
|
+
)
|
|
466
|
+
self.model = self.trainer.model
|
|
467
|
+
self.trainer.hub_session = self.session # attach optional HUB session
|
|
468
|
+
self.trainer.train()
|
|
469
|
+
# Update model and cfg after training
|
|
470
|
+
if RANK in (-1, 0):
|
|
471
|
+
self.model, _ = attempt_load_one_weight(str(self.trainer.best))
|
|
472
|
+
self.overrides = self.model.args
|
|
473
|
+
self.metrics = getattr(
|
|
474
|
+
self.trainer.validator, "metrics", None
|
|
475
|
+
) # TODO: no metrics returned by DDP
|
|
476
|
+
|
|
477
|
+
def to(self, device):
|
|
478
|
+
"""
|
|
479
|
+
Sends the model to the given device.
|
|
480
|
+
|
|
481
|
+
Args:
|
|
482
|
+
device (str): device
|
|
483
|
+
"""
|
|
484
|
+
self._check_is_pytorch_model()
|
|
485
|
+
self.model.to(device)
|
|
486
|
+
|
|
487
|
+
def tune(
|
|
488
|
+
self,
|
|
489
|
+
data: str,
|
|
490
|
+
space: dict = None,
|
|
491
|
+
grace_period: int = 10,
|
|
492
|
+
gpu_per_trial: int = None,
|
|
493
|
+
max_samples: int = 10,
|
|
494
|
+
train_args: dict = None,
|
|
495
|
+
):
|
|
496
|
+
"""
|
|
497
|
+
Runs hyperparameter tuning using Ray Tune.
|
|
498
|
+
|
|
499
|
+
Args:
|
|
500
|
+
data (str): The dataset to run the tuner on.
|
|
501
|
+
space (dict, optional): The hyperparameter search space. Defaults to None.
|
|
502
|
+
grace_period (int, optional): The grace period in epochs of the ASHA scheduler. Defaults to 10.
|
|
503
|
+
gpu_per_trial (int, optional): The number of GPUs to allocate per trial. Defaults to None.
|
|
504
|
+
max_samples (int, optional): The maximum number of trials to run. Defaults to 10.
|
|
505
|
+
train_args (dict, optional): Additional arguments to pass to the `train()` method. Defaults to {}.
|
|
506
|
+
|
|
507
|
+
Returns:
|
|
508
|
+
(dict): A dictionary containing the results of the hyperparameter search.
|
|
509
|
+
|
|
510
|
+
Raises:
|
|
511
|
+
ModuleNotFoundError: If Ray Tune is not installed.
|
|
512
|
+
"""
|
|
513
|
+
if train_args is None:
|
|
514
|
+
train_args = {}
|
|
515
|
+
|
|
516
|
+
try:
|
|
517
|
+
from ultralytics.yolo.utils.tuner import (
|
|
518
|
+
ASHAScheduler,
|
|
519
|
+
RunConfig,
|
|
520
|
+
WandbLoggerCallback,
|
|
521
|
+
default_space,
|
|
522
|
+
task_metric_map,
|
|
523
|
+
tune,
|
|
524
|
+
)
|
|
525
|
+
except ImportError:
|
|
526
|
+
raise ModuleNotFoundError(
|
|
527
|
+
"Install Ray Tune: `pip install 'ray[tune]'`"
|
|
528
|
+
)
|
|
529
|
+
|
|
530
|
+
try:
|
|
531
|
+
import wandb
|
|
532
|
+
from wandb import __version__ # noqa
|
|
533
|
+
except ImportError:
|
|
534
|
+
wandb = False
|
|
535
|
+
|
|
536
|
+
def _tune(config):
|
|
537
|
+
"""
|
|
538
|
+
Trains the YOLO model with the specified hyperparameters and additional arguments.
|
|
539
|
+
|
|
540
|
+
Args:
|
|
541
|
+
config (dict): A dictionary of hyperparameters to use for training.
|
|
542
|
+
|
|
543
|
+
Returns:
|
|
544
|
+
None.
|
|
545
|
+
"""
|
|
546
|
+
self._reset_callbacks()
|
|
547
|
+
config.update(train_args)
|
|
548
|
+
self.train(**config)
|
|
549
|
+
|
|
550
|
+
if not space:
|
|
551
|
+
LOGGER.warning(
|
|
552
|
+
"WARNING: search space not provided. Using default search space"
|
|
553
|
+
)
|
|
554
|
+
space = default_space
|
|
555
|
+
|
|
556
|
+
space["data"] = data
|
|
557
|
+
|
|
558
|
+
# Define the trainable function with allocated resources
|
|
559
|
+
trainable_with_resources = tune.with_resources(
|
|
560
|
+
_tune, {"cpu": NUM_THREADS, "gpu": gpu_per_trial or 0}
|
|
561
|
+
)
|
|
562
|
+
|
|
563
|
+
# Define the ASHA scheduler for hyperparameter search
|
|
564
|
+
asha_scheduler = ASHAScheduler(
|
|
565
|
+
time_attr="epoch",
|
|
566
|
+
metric=task_metric_map[self.task],
|
|
567
|
+
mode="max",
|
|
568
|
+
max_t=train_args.get("epochs") or 100,
|
|
569
|
+
grace_period=grace_period,
|
|
570
|
+
reduction_factor=3,
|
|
571
|
+
)
|
|
572
|
+
|
|
573
|
+
# Define the callbacks for the hyperparameter search
|
|
574
|
+
tuner_callbacks = (
|
|
575
|
+
[WandbLoggerCallback(project="YOLOv8-tune")] if wandb else []
|
|
576
|
+
)
|
|
577
|
+
|
|
578
|
+
# Create the Ray Tune hyperparameter search tuner
|
|
579
|
+
tuner = tune.Tuner(
|
|
580
|
+
trainable_with_resources,
|
|
581
|
+
param_space=space,
|
|
582
|
+
tune_config=tune.TuneConfig(
|
|
583
|
+
scheduler=asha_scheduler, num_samples=max_samples
|
|
584
|
+
),
|
|
585
|
+
run_config=RunConfig(
|
|
586
|
+
callbacks=tuner_callbacks, local_dir="./runs"
|
|
587
|
+
),
|
|
588
|
+
)
|
|
589
|
+
|
|
590
|
+
# Run the hyperparameter search
|
|
591
|
+
tuner.fit()
|
|
592
|
+
|
|
593
|
+
# Return the results of the hyperparameter search
|
|
594
|
+
return tuner.get_results()
|
|
595
|
+
|
|
596
|
+
@property
|
|
597
|
+
def names(self):
|
|
598
|
+
"""Returns class names of the loaded model."""
|
|
599
|
+
return self.model.names if hasattr(self.model, "names") else None
|
|
600
|
+
|
|
601
|
+
@property
|
|
602
|
+
def device(self):
|
|
603
|
+
"""Returns device if PyTorch model."""
|
|
604
|
+
return (
|
|
605
|
+
next(self.model.parameters()).device
|
|
606
|
+
if isinstance(self.model, nn.Module)
|
|
607
|
+
else None
|
|
608
|
+
)
|
|
609
|
+
|
|
610
|
+
@property
|
|
611
|
+
def transforms(self):
|
|
612
|
+
"""Returns transform of the loaded model."""
|
|
613
|
+
return (
|
|
614
|
+
self.model.transforms
|
|
615
|
+
if hasattr(self.model, "transforms")
|
|
616
|
+
else None
|
|
617
|
+
)
|
|
618
|
+
|
|
619
|
+
def add_callback(self, event: str, func):
|
|
620
|
+
"""Add a callback."""
|
|
621
|
+
self.callbacks[event].append(func)
|
|
622
|
+
|
|
623
|
+
def clear_callback(self, event: str):
|
|
624
|
+
"""Clear all event callbacks."""
|
|
625
|
+
self.callbacks[event] = []
|
|
626
|
+
|
|
627
|
+
@staticmethod
|
|
628
|
+
def _reset_ckpt_args(args):
|
|
629
|
+
"""Reset arguments when loading a PyTorch model."""
|
|
630
|
+
include = {
|
|
631
|
+
"imgsz",
|
|
632
|
+
"data",
|
|
633
|
+
"task",
|
|
634
|
+
"single_cls",
|
|
635
|
+
} # only remember these arguments when loading a PyTorch model
|
|
636
|
+
return {k: v for k, v in args.items() if k in include}
|
|
637
|
+
|
|
638
|
+
def _reset_callbacks(self):
|
|
639
|
+
"""Reset all registered callbacks."""
|
|
640
|
+
for event in callbacks.default_callbacks.keys():
|
|
641
|
+
self.callbacks[event] = [callbacks.default_callbacks[event][0]]
|