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,286 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
|
|
3
|
+
import glob
|
|
4
|
+
import math
|
|
5
|
+
import os
|
|
6
|
+
import random
|
|
7
|
+
from copy import deepcopy
|
|
8
|
+
from multiprocessing.pool import ThreadPool
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from typing import Optional
|
|
11
|
+
|
|
12
|
+
import cv2
|
|
13
|
+
import numpy as np
|
|
14
|
+
import psutil
|
|
15
|
+
from torch.utils.data import Dataset
|
|
16
|
+
from tqdm import tqdm
|
|
17
|
+
|
|
18
|
+
from ..utils import DEFAULT_CFG, LOCAL_RANK, LOGGER, NUM_THREADS, TQDM_BAR_FORMAT
|
|
19
|
+
from .utils import HELP_URL, IMG_FORMATS
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class BaseDataset(Dataset):
|
|
23
|
+
"""
|
|
24
|
+
Base dataset class for loading and processing image data.
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
img_path (str): Path to the folder containing images.
|
|
28
|
+
imgsz (int, optional): Image size. Defaults to 640.
|
|
29
|
+
cache (bool, optional): Cache images to RAM or disk during training. Defaults to False.
|
|
30
|
+
augment (bool, optional): If True, data augmentation is applied. Defaults to True.
|
|
31
|
+
hyp (dict, optional): Hyperparameters to apply data augmentation. Defaults to None.
|
|
32
|
+
prefix (str, optional): Prefix to print in log messages. Defaults to ''.
|
|
33
|
+
rect (bool, optional): If True, rectangular training is used. Defaults to False.
|
|
34
|
+
batch_size (int, optional): Size of batches. Defaults to None.
|
|
35
|
+
stride (int, optional): Stride. Defaults to 32.
|
|
36
|
+
pad (float, optional): Padding. Defaults to 0.0.
|
|
37
|
+
single_cls (bool, optional): If True, single class training is used. Defaults to False.
|
|
38
|
+
classes (list): List of included classes. Default is None.
|
|
39
|
+
fraction (float): Fraction of dataset to utilize. Default is 1.0 (use all data).
|
|
40
|
+
|
|
41
|
+
Attributes:
|
|
42
|
+
im_files (list): List of image file paths.
|
|
43
|
+
labels (list): List of label data dictionaries.
|
|
44
|
+
ni (int): Number of images in the dataset.
|
|
45
|
+
ims (list): List of loaded images.
|
|
46
|
+
npy_files (list): List of numpy file paths.
|
|
47
|
+
transforms (callable): Image transformation function.
|
|
48
|
+
"""
|
|
49
|
+
|
|
50
|
+
def __init__(self,
|
|
51
|
+
img_path,
|
|
52
|
+
imgsz=640,
|
|
53
|
+
cache=False,
|
|
54
|
+
augment=True,
|
|
55
|
+
hyp=DEFAULT_CFG,
|
|
56
|
+
prefix='',
|
|
57
|
+
rect=False,
|
|
58
|
+
batch_size=16,
|
|
59
|
+
stride=32,
|
|
60
|
+
pad=0.5,
|
|
61
|
+
single_cls=False,
|
|
62
|
+
classes=None,
|
|
63
|
+
fraction=1.0):
|
|
64
|
+
super().__init__()
|
|
65
|
+
self.img_path = img_path
|
|
66
|
+
self.imgsz = imgsz
|
|
67
|
+
self.augment = augment
|
|
68
|
+
self.single_cls = single_cls
|
|
69
|
+
self.prefix = prefix
|
|
70
|
+
self.fraction = fraction
|
|
71
|
+
self.im_files = self.get_img_files(self.img_path)
|
|
72
|
+
self.labels = self.get_labels()
|
|
73
|
+
self.update_labels(include_class=classes) # single_cls and include_class
|
|
74
|
+
self.ni = len(self.labels) # number of images
|
|
75
|
+
self.rect = rect
|
|
76
|
+
self.batch_size = batch_size
|
|
77
|
+
self.stride = stride
|
|
78
|
+
self.pad = pad
|
|
79
|
+
if self.rect:
|
|
80
|
+
assert self.batch_size is not None
|
|
81
|
+
self.set_rectangle()
|
|
82
|
+
|
|
83
|
+
# Buffer thread for mosaic images
|
|
84
|
+
self.buffer = [] # buffer size = batch size
|
|
85
|
+
self.max_buffer_length = min((self.ni, self.batch_size * 8, 1000)) if self.augment else 0
|
|
86
|
+
|
|
87
|
+
# Cache stuff
|
|
88
|
+
if cache == 'ram' and not self.check_cache_ram():
|
|
89
|
+
cache = False
|
|
90
|
+
self.ims, self.im_hw0, self.im_hw = [None] * self.ni, [None] * self.ni, [None] * self.ni
|
|
91
|
+
self.npy_files = [Path(f).with_suffix('.npy') for f in self.im_files]
|
|
92
|
+
if cache:
|
|
93
|
+
self.cache_images(cache)
|
|
94
|
+
|
|
95
|
+
# Transforms
|
|
96
|
+
self.transforms = self.build_transforms(hyp=hyp)
|
|
97
|
+
|
|
98
|
+
def get_img_files(self, img_path):
|
|
99
|
+
"""Read image files."""
|
|
100
|
+
try:
|
|
101
|
+
f = [] # image files
|
|
102
|
+
for p in img_path if isinstance(img_path, list) else [img_path]:
|
|
103
|
+
p = Path(p) # os-agnostic
|
|
104
|
+
if p.is_dir(): # dir
|
|
105
|
+
f += glob.glob(str(p / '**' / '*.*'), recursive=True)
|
|
106
|
+
# F = list(p.rglob('*.*')) # pathlib
|
|
107
|
+
elif p.is_file(): # file
|
|
108
|
+
with open(p) as t:
|
|
109
|
+
t = t.read().strip().splitlines()
|
|
110
|
+
parent = str(p.parent) + os.sep
|
|
111
|
+
f += [x.replace('./', parent) if x.startswith('./') else x for x in t] # local to global path
|
|
112
|
+
# F += [p.parent / x.lstrip(os.sep) for x in t] # local to global path (pathlib)
|
|
113
|
+
else:
|
|
114
|
+
raise FileNotFoundError(f'{self.prefix}{p} does not exist')
|
|
115
|
+
im_files = sorted(x.replace('/', os.sep) for x in f if x.split('.')[-1].lower() in IMG_FORMATS)
|
|
116
|
+
# self.img_files = sorted([x for x in f if x.suffix[1:].lower() in IMG_FORMATS]) # pathlib
|
|
117
|
+
assert im_files, f'{self.prefix}No images found'
|
|
118
|
+
except Exception as e:
|
|
119
|
+
raise FileNotFoundError(f'{self.prefix}Error loading data from {img_path}\n{HELP_URL}') from e
|
|
120
|
+
if self.fraction < 1:
|
|
121
|
+
im_files = im_files[:round(len(im_files) * self.fraction)]
|
|
122
|
+
return im_files
|
|
123
|
+
|
|
124
|
+
def update_labels(self, include_class: Optional[list]):
|
|
125
|
+
"""include_class, filter labels to include only these classes (optional)."""
|
|
126
|
+
include_class_array = np.array(include_class).reshape(1, -1)
|
|
127
|
+
for i in range(len(self.labels)):
|
|
128
|
+
if include_class is not None:
|
|
129
|
+
cls = self.labels[i]['cls']
|
|
130
|
+
bboxes = self.labels[i]['bboxes']
|
|
131
|
+
segments = self.labels[i]['segments']
|
|
132
|
+
keypoints = self.labels[i]['keypoints']
|
|
133
|
+
j = (cls == include_class_array).any(1)
|
|
134
|
+
self.labels[i]['cls'] = cls[j]
|
|
135
|
+
self.labels[i]['bboxes'] = bboxes[j]
|
|
136
|
+
if segments:
|
|
137
|
+
self.labels[i]['segments'] = [segments[si] for si, idx in enumerate(j) if idx]
|
|
138
|
+
if keypoints is not None:
|
|
139
|
+
self.labels[i]['keypoints'] = keypoints[j]
|
|
140
|
+
if self.single_cls:
|
|
141
|
+
self.labels[i]['cls'][:, 0] = 0
|
|
142
|
+
|
|
143
|
+
def load_image(self, i):
|
|
144
|
+
"""Loads 1 image from dataset index 'i', returns (im, resized hw)."""
|
|
145
|
+
im, f, fn = self.ims[i], self.im_files[i], self.npy_files[i]
|
|
146
|
+
if im is None: # not cached in RAM
|
|
147
|
+
if fn.exists(): # load npy
|
|
148
|
+
im = np.load(fn)
|
|
149
|
+
else: # read image
|
|
150
|
+
im = cv2.imread(f) # BGR
|
|
151
|
+
if im is None:
|
|
152
|
+
raise FileNotFoundError(f'Image Not Found {f}')
|
|
153
|
+
h0, w0 = im.shape[:2] # orig hw
|
|
154
|
+
r = self.imgsz / max(h0, w0) # ratio
|
|
155
|
+
if r != 1: # if sizes are not equal
|
|
156
|
+
interp = cv2.INTER_LINEAR if (self.augment or r > 1) else cv2.INTER_AREA
|
|
157
|
+
im = cv2.resize(im, (min(math.ceil(w0 * r), self.imgsz), min(math.ceil(h0 * r), self.imgsz)),
|
|
158
|
+
interpolation=interp)
|
|
159
|
+
|
|
160
|
+
# Add to buffer if training with augmentations
|
|
161
|
+
if self.augment:
|
|
162
|
+
self.ims[i], self.im_hw0[i], self.im_hw[i] = im, (h0, w0), im.shape[:2] # im, hw_original, hw_resized
|
|
163
|
+
self.buffer.append(i)
|
|
164
|
+
if len(self.buffer) >= self.max_buffer_length:
|
|
165
|
+
j = self.buffer.pop(0)
|
|
166
|
+
self.ims[j], self.im_hw0[j], self.im_hw[j] = None, None, None
|
|
167
|
+
|
|
168
|
+
return im, (h0, w0), im.shape[:2]
|
|
169
|
+
|
|
170
|
+
return self.ims[i], self.im_hw0[i], self.im_hw[i]
|
|
171
|
+
|
|
172
|
+
def cache_images(self, cache):
|
|
173
|
+
"""Cache images to memory or disk."""
|
|
174
|
+
b, gb = 0, 1 << 30 # bytes of cached images, bytes per gigabytes
|
|
175
|
+
fcn = self.cache_images_to_disk if cache == 'disk' else self.load_image
|
|
176
|
+
with ThreadPool(NUM_THREADS) as pool:
|
|
177
|
+
results = pool.imap(fcn, range(self.ni))
|
|
178
|
+
pbar = tqdm(enumerate(results), total=self.ni, bar_format=TQDM_BAR_FORMAT, disable=LOCAL_RANK > 0)
|
|
179
|
+
for i, x in pbar:
|
|
180
|
+
if cache == 'disk':
|
|
181
|
+
b += self.npy_files[i].stat().st_size
|
|
182
|
+
else: # 'ram'
|
|
183
|
+
self.ims[i], self.im_hw0[i], self.im_hw[i] = x # im, hw_orig, hw_resized = load_image(self, i)
|
|
184
|
+
b += self.ims[i].nbytes
|
|
185
|
+
pbar.desc = f'{self.prefix}Caching images ({b / gb:.1f}GB {cache})'
|
|
186
|
+
pbar.close()
|
|
187
|
+
|
|
188
|
+
def cache_images_to_disk(self, i):
|
|
189
|
+
"""Saves an image as an *.npy file for faster loading."""
|
|
190
|
+
f = self.npy_files[i]
|
|
191
|
+
if not f.exists():
|
|
192
|
+
np.save(f.as_posix(), cv2.imread(self.im_files[i]))
|
|
193
|
+
|
|
194
|
+
def check_cache_ram(self, safety_margin=0.5):
|
|
195
|
+
"""Check image caching requirements vs available memory."""
|
|
196
|
+
b, gb = 0, 1 << 30 # bytes of cached images, bytes per gigabytes
|
|
197
|
+
n = min(self.ni, 30) # extrapolate from 30 random images
|
|
198
|
+
for _ in range(n):
|
|
199
|
+
im = cv2.imread(random.choice(self.im_files)) # sample image
|
|
200
|
+
ratio = self.imgsz / max(im.shape[0], im.shape[1]) # max(h, w) # ratio
|
|
201
|
+
b += im.nbytes * ratio ** 2
|
|
202
|
+
mem_required = b * self.ni / n * (1 + safety_margin) # GB required to cache dataset into RAM
|
|
203
|
+
mem = psutil.virtual_memory()
|
|
204
|
+
cache = mem_required < mem.available # to cache or not to cache, that is the question
|
|
205
|
+
if not cache:
|
|
206
|
+
LOGGER.info(f'{self.prefix}{mem_required / gb:.1f}GB RAM required to cache images '
|
|
207
|
+
f'with {int(safety_margin * 100)}% safety margin but only '
|
|
208
|
+
f'{mem.available / gb:.1f}/{mem.total / gb:.1f}GB available, '
|
|
209
|
+
f"{'caching images ✅' if cache else 'not caching images ⚠️'}")
|
|
210
|
+
return cache
|
|
211
|
+
|
|
212
|
+
def set_rectangle(self):
|
|
213
|
+
"""Sets the shape of bounding boxes for YOLO detections as rectangles."""
|
|
214
|
+
bi = np.floor(np.arange(self.ni) / self.batch_size).astype(int) # batch index
|
|
215
|
+
nb = bi[-1] + 1 # number of batches
|
|
216
|
+
|
|
217
|
+
s = np.array([x.pop('shape') for x in self.labels]) # hw
|
|
218
|
+
ar = s[:, 0] / s[:, 1] # aspect ratio
|
|
219
|
+
irect = ar.argsort()
|
|
220
|
+
self.im_files = [self.im_files[i] for i in irect]
|
|
221
|
+
self.labels = [self.labels[i] for i in irect]
|
|
222
|
+
ar = ar[irect]
|
|
223
|
+
|
|
224
|
+
# Set training image shapes
|
|
225
|
+
shapes = [[1, 1]] * nb
|
|
226
|
+
for i in range(nb):
|
|
227
|
+
ari = ar[bi == i]
|
|
228
|
+
mini, maxi = ari.min(), ari.max()
|
|
229
|
+
if maxi < 1:
|
|
230
|
+
shapes[i] = [maxi, 1]
|
|
231
|
+
elif mini > 1:
|
|
232
|
+
shapes[i] = [1, 1 / mini]
|
|
233
|
+
|
|
234
|
+
self.batch_shapes = np.ceil(np.array(shapes) * self.imgsz / self.stride + self.pad).astype(int) * self.stride
|
|
235
|
+
self.batch = bi # batch index of image
|
|
236
|
+
|
|
237
|
+
def __getitem__(self, index):
|
|
238
|
+
"""Returns transformed label information for given index."""
|
|
239
|
+
return self.transforms(self.get_image_and_label(index))
|
|
240
|
+
|
|
241
|
+
def get_image_and_label(self, index):
|
|
242
|
+
"""Get and return label information from the dataset."""
|
|
243
|
+
label = deepcopy(self.labels[index]) # requires deepcopy() https://github.com/ultralytics/ultralytics/pull/1948
|
|
244
|
+
label.pop('shape', None) # shape is for rect, remove it
|
|
245
|
+
label['img'], label['ori_shape'], label['resized_shape'] = self.load_image(index)
|
|
246
|
+
label['ratio_pad'] = (label['resized_shape'][0] / label['ori_shape'][0],
|
|
247
|
+
label['resized_shape'][1] / label['ori_shape'][1]) # for evaluation
|
|
248
|
+
if self.rect:
|
|
249
|
+
label['rect_shape'] = self.batch_shapes[self.batch[index]]
|
|
250
|
+
return self.update_labels_info(label)
|
|
251
|
+
|
|
252
|
+
def __len__(self):
|
|
253
|
+
"""Returns the length of the labels list for the dataset."""
|
|
254
|
+
return len(self.labels)
|
|
255
|
+
|
|
256
|
+
def update_labels_info(self, label):
|
|
257
|
+
"""custom your label format here."""
|
|
258
|
+
return label
|
|
259
|
+
|
|
260
|
+
def build_transforms(self, hyp=None):
|
|
261
|
+
"""Users can custom augmentations here
|
|
262
|
+
like:
|
|
263
|
+
if self.augment:
|
|
264
|
+
# Training transforms
|
|
265
|
+
return Compose([])
|
|
266
|
+
else:
|
|
267
|
+
# Val transforms
|
|
268
|
+
return Compose([])
|
|
269
|
+
"""
|
|
270
|
+
raise NotImplementedError
|
|
271
|
+
|
|
272
|
+
def get_labels(self):
|
|
273
|
+
"""Users can custom their own format here.
|
|
274
|
+
Make sure your output is a list with each element like below:
|
|
275
|
+
dict(
|
|
276
|
+
im_file=im_file,
|
|
277
|
+
shape=shape, # format: (height, width)
|
|
278
|
+
cls=cls,
|
|
279
|
+
bboxes=bboxes, # xywh
|
|
280
|
+
segments=segments, # xy
|
|
281
|
+
keypoints=keypoints, # xy
|
|
282
|
+
normalized=True, # or False
|
|
283
|
+
bbox_format="xyxy", # or xywh, ltwh
|
|
284
|
+
)
|
|
285
|
+
"""
|
|
286
|
+
raise NotImplementedError
|
|
@@ -0,0 +1,213 @@
|
|
|
1
|
+
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import random
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import torch
|
|
9
|
+
from PIL import Image
|
|
10
|
+
from torch.utils.data import dataloader, distributed
|
|
11
|
+
|
|
12
|
+
from .dataloaders.stream_loaders import (
|
|
13
|
+
LOADERS,
|
|
14
|
+
LoadImages,
|
|
15
|
+
LoadPilAndNumpy,
|
|
16
|
+
LoadScreenshots,
|
|
17
|
+
LoadStreams,
|
|
18
|
+
LoadTensor,
|
|
19
|
+
SourceTypes,
|
|
20
|
+
autocast_list,
|
|
21
|
+
)
|
|
22
|
+
from .utils import IMG_FORMATS, VID_FORMATS
|
|
23
|
+
from ..utils.checks import check_file
|
|
24
|
+
|
|
25
|
+
from ..utils import RANK, colorstr
|
|
26
|
+
from .dataset import YOLODataset
|
|
27
|
+
from .utils import PIN_MEMORY
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class InfiniteDataLoader(dataloader.DataLoader):
|
|
31
|
+
"""Dataloader that reuses workers. Uses same syntax as vanilla DataLoader."""
|
|
32
|
+
|
|
33
|
+
def __init__(self, *args, **kwargs):
|
|
34
|
+
"""Dataloader that infinitely recycles workers, inherits from DataLoader."""
|
|
35
|
+
super().__init__(*args, **kwargs)
|
|
36
|
+
object.__setattr__(
|
|
37
|
+
self, "batch_sampler", _RepeatSampler(self.batch_sampler)
|
|
38
|
+
)
|
|
39
|
+
self.iterator = super().__iter__()
|
|
40
|
+
|
|
41
|
+
def __len__(self):
|
|
42
|
+
"""Returns the length of the batch sampler's sampler."""
|
|
43
|
+
return len(self.batch_sampler.sampler)
|
|
44
|
+
|
|
45
|
+
def __iter__(self):
|
|
46
|
+
"""Creates a sampler that repeats indefinitely."""
|
|
47
|
+
for _ in range(len(self)):
|
|
48
|
+
yield next(self.iterator)
|
|
49
|
+
|
|
50
|
+
def reset(self):
|
|
51
|
+
"""Reset iterator.
|
|
52
|
+
This is useful when we want to modify settings of dataset while training.
|
|
53
|
+
"""
|
|
54
|
+
self.iterator = self._get_iterator()
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class _RepeatSampler:
|
|
58
|
+
"""
|
|
59
|
+
Sampler that repeats forever.
|
|
60
|
+
|
|
61
|
+
Args:
|
|
62
|
+
sampler (Dataset.sampler): The sampler to repeat.
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
def __init__(self, sampler):
|
|
66
|
+
"""Initializes an object that repeats a given sampler indefinitely."""
|
|
67
|
+
self.sampler = sampler
|
|
68
|
+
|
|
69
|
+
def __iter__(self):
|
|
70
|
+
"""Iterates over the 'sampler' and yields its contents."""
|
|
71
|
+
while True:
|
|
72
|
+
yield from iter(self.sampler)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def seed_worker(worker_id): # noqa
|
|
76
|
+
# Set dataloader worker seed https://pytorch.org/docs/stable/notes/randomness.html#dataloader
|
|
77
|
+
worker_seed = torch.initial_seed() % 2**32
|
|
78
|
+
np.random.seed(worker_seed)
|
|
79
|
+
random.seed(worker_seed)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def build_yolo_dataset(
|
|
83
|
+
cfg, img_path, batch, data, mode="train", rect=False, stride=32
|
|
84
|
+
):
|
|
85
|
+
"""Build YOLO Dataset"""
|
|
86
|
+
return YOLODataset(
|
|
87
|
+
img_path=img_path,
|
|
88
|
+
imgsz=cfg.imgsz,
|
|
89
|
+
batch_size=batch,
|
|
90
|
+
augment=mode == "train", # augmentation
|
|
91
|
+
hyp=cfg, # TODO: probably add a get_hyps_from_cfg function
|
|
92
|
+
rect=cfg.rect or rect, # rectangular batches
|
|
93
|
+
cache=cfg.cache or None,
|
|
94
|
+
single_cls=cfg.single_cls or False,
|
|
95
|
+
stride=int(stride),
|
|
96
|
+
pad=0.0 if mode == "train" else 0.5,
|
|
97
|
+
prefix=colorstr(f"{mode}: "),
|
|
98
|
+
use_segments=cfg.task == "segment",
|
|
99
|
+
use_keypoints=cfg.task == "pose",
|
|
100
|
+
classes=cfg.classes,
|
|
101
|
+
data=data,
|
|
102
|
+
fraction=cfg.fraction if mode == "train" else 1.0,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def build_dataloader(dataset, batch, workers, shuffle=True, rank=-1):
|
|
107
|
+
"""Return an InfiniteDataLoader or DataLoader for training or validation set."""
|
|
108
|
+
batch = min(batch, len(dataset))
|
|
109
|
+
nd = torch.cuda.device_count() # number of CUDA devices
|
|
110
|
+
nw = min(
|
|
111
|
+
[os.cpu_count() // max(nd, 1), batch if batch > 1 else 0, workers]
|
|
112
|
+
) # number of workers
|
|
113
|
+
sampler = (
|
|
114
|
+
None
|
|
115
|
+
if rank == -1
|
|
116
|
+
else distributed.DistributedSampler(dataset, shuffle=shuffle)
|
|
117
|
+
)
|
|
118
|
+
generator = torch.Generator()
|
|
119
|
+
generator.manual_seed(6148914691236517205 + RANK)
|
|
120
|
+
return InfiniteDataLoader(
|
|
121
|
+
dataset=dataset,
|
|
122
|
+
batch_size=batch,
|
|
123
|
+
shuffle=shuffle and sampler is None,
|
|
124
|
+
num_workers=nw,
|
|
125
|
+
sampler=sampler,
|
|
126
|
+
pin_memory=PIN_MEMORY,
|
|
127
|
+
collate_fn=getattr(dataset, "collate_fn", None),
|
|
128
|
+
worker_init_fn=seed_worker,
|
|
129
|
+
generator=generator,
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def check_source(source):
|
|
134
|
+
"""Check source type and return corresponding flag values."""
|
|
135
|
+
webcam, screenshot, from_img, in_memory, tensor = (
|
|
136
|
+
False,
|
|
137
|
+
False,
|
|
138
|
+
False,
|
|
139
|
+
False,
|
|
140
|
+
False,
|
|
141
|
+
)
|
|
142
|
+
if isinstance(source, (str, int, Path)): # int for local usb camera
|
|
143
|
+
source = str(source)
|
|
144
|
+
is_file = Path(source).suffix[1:] in (IMG_FORMATS + VID_FORMATS)
|
|
145
|
+
is_url = source.lower().startswith(
|
|
146
|
+
("https://", "http://", "rtsp://", "rtmp://")
|
|
147
|
+
)
|
|
148
|
+
webcam = (
|
|
149
|
+
source.isnumeric()
|
|
150
|
+
or source.endswith(".streams")
|
|
151
|
+
or (is_url and not is_file)
|
|
152
|
+
)
|
|
153
|
+
screenshot = source.lower() == "screen"
|
|
154
|
+
if is_url and is_file:
|
|
155
|
+
source = check_file(source) # download
|
|
156
|
+
elif isinstance(source, tuple(LOADERS)):
|
|
157
|
+
in_memory = True
|
|
158
|
+
elif isinstance(source, (list, tuple)):
|
|
159
|
+
source = autocast_list(
|
|
160
|
+
source
|
|
161
|
+
) # convert all list elements to PIL or np arrays
|
|
162
|
+
from_img = True
|
|
163
|
+
elif isinstance(source, (Image.Image, np.ndarray)):
|
|
164
|
+
from_img = True
|
|
165
|
+
elif isinstance(source, torch.Tensor):
|
|
166
|
+
tensor = True
|
|
167
|
+
else:
|
|
168
|
+
raise TypeError(
|
|
169
|
+
"Unsupported image type. For supported types see https://docs.ultralytics.com/modes/predict"
|
|
170
|
+
)
|
|
171
|
+
|
|
172
|
+
return source, webcam, screenshot, from_img, in_memory, tensor
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
def load_inference_source(source=None, imgsz=640, vid_stride=1):
|
|
176
|
+
"""
|
|
177
|
+
Loads an inference source for object detection and applies necessary transformations.
|
|
178
|
+
|
|
179
|
+
Args:
|
|
180
|
+
source (str, Path, Tensor, PIL.Image, np.ndarray): The input source for inference.
|
|
181
|
+
imgsz (int, optional): The size of the image for inference. Default is 640.
|
|
182
|
+
vid_stride (int, optional): The frame interval for video sources. Default is 1.
|
|
183
|
+
|
|
184
|
+
Returns:
|
|
185
|
+
dataset (Dataset): A dataset object for the specified input source.
|
|
186
|
+
"""
|
|
187
|
+
source, webcam, screenshot, from_img, in_memory, tensor = check_source(
|
|
188
|
+
source
|
|
189
|
+
)
|
|
190
|
+
source_type = (
|
|
191
|
+
source.source_type
|
|
192
|
+
if in_memory
|
|
193
|
+
else SourceTypes(webcam, screenshot, from_img, tensor)
|
|
194
|
+
)
|
|
195
|
+
|
|
196
|
+
# Dataloader
|
|
197
|
+
if tensor:
|
|
198
|
+
dataset = LoadTensor(source)
|
|
199
|
+
elif in_memory:
|
|
200
|
+
dataset = source
|
|
201
|
+
elif webcam:
|
|
202
|
+
dataset = LoadStreams(source, imgsz=imgsz, vid_stride=vid_stride)
|
|
203
|
+
elif screenshot:
|
|
204
|
+
dataset = LoadScreenshots(source, imgsz=imgsz)
|
|
205
|
+
elif from_img:
|
|
206
|
+
dataset = LoadPilAndNumpy(source, imgsz=imgsz)
|
|
207
|
+
else:
|
|
208
|
+
dataset = LoadImages(source, imgsz=imgsz, vid_stride=vid_stride)
|
|
209
|
+
|
|
210
|
+
# Attach source types to the dataset
|
|
211
|
+
setattr(dataset, "source_type", source_type)
|
|
212
|
+
|
|
213
|
+
return dataset
|