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