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,357 @@
1
+ r""""This file is based on torch/utils/data/_utils/worker.py
2
+
3
+ Contains definitions of the methods used by the _BaseDataLoaderIter workers.
4
+ These **needs** to be in global scope since Py2 doesn't support serializing
5
+ static methods.
6
+ """
7
+
8
+ import os
9
+ import queue
10
+ import random
11
+ from dataclasses import dataclass
12
+ from typing import TYPE_CHECKING, Optional, Union
13
+
14
+ import torch
15
+ from torch._utils import ExceptionWrapper
16
+ from torch.utils.data._utils import HAS_NUMPY, IS_WINDOWS, MP_STATUS_CHECK_INTERVAL, signal_handling
17
+
18
+ if TYPE_CHECKING:
19
+ from torch.utils.data import Dataset
20
+
21
+ from .controller import RRSController
22
+
23
+ if IS_WINDOWS:
24
+ import ctypes
25
+ from ctypes.wintypes import BOOL, DWORD, HANDLE
26
+
27
+ # On Windows, the parent ID of the worker process remains unchanged when the manager process
28
+ # is gone, and the only way to check it through OS is to let the worker have a process handle
29
+ # of the manager and ask if the process status has changed.
30
+ class ManagerWatchdog:
31
+ def __init__(self):
32
+ self.manager_pid = os.getppid()
33
+
34
+ # mypy cannot detect this code is windows only
35
+ self.kernel32 = ctypes.WinDLL("kernel32", use_last_error=True) # type: ignore[attr-defined]
36
+ self.kernel32.OpenProcess.argtypes = (DWORD, BOOL, DWORD)
37
+ self.kernel32.OpenProcess.restype = HANDLE
38
+ self.kernel32.WaitForSingleObject.argtypes = (HANDLE, DWORD)
39
+ self.kernel32.WaitForSingleObject.restype = DWORD
40
+
41
+ # Value obtained from https://msdn.microsoft.com/en-us/library/ms684880.aspx
42
+ SYNCHRONIZE = 0x00100000
43
+ self.manager_handle = self.kernel32.OpenProcess(SYNCHRONIZE, 0, self.manager_pid)
44
+
45
+ if not self.manager_handle:
46
+ raise ctypes.WinError(ctypes.get_last_error()) # type: ignore[attr-defined]
47
+
48
+ self.manager_dead = False
49
+
50
+ def is_alive(self):
51
+ if not self.manager_dead:
52
+ # Value obtained from https://msdn.microsoft.com/en-us/library/windows/desktop/ms687032.aspx
53
+ self.manager_dead = self.kernel32.WaitForSingleObject(self.manager_handle, 0) == 0
54
+ return not self.manager_dead
55
+
56
+ else:
57
+
58
+ class ManagerWatchdog: # type: ignore[no-redef]
59
+ def __init__(self):
60
+ self.manager_pid = os.getppid()
61
+ self.manager_dead = False
62
+
63
+ def is_alive(self):
64
+ if not self.manager_dead:
65
+ self.manager_dead = os.getppid() != self.manager_pid
66
+ return not self.manager_dead
67
+
68
+
69
+ _worker_info = None
70
+
71
+
72
+ class WorkerInfo:
73
+ id: int
74
+ num_workers: int
75
+ seed: int
76
+ dataset: "Dataset"
77
+ __initialized = False
78
+
79
+ def __init__(self, **kwargs):
80
+ for k, v in kwargs.items():
81
+ setattr(self, k, v)
82
+ self.__keys = tuple(kwargs.keys())
83
+ self.__initialized = True
84
+
85
+ def __setattr__(self, key, val):
86
+ if self.__initialized:
87
+ raise RuntimeError("Cannot assign attributes to {} objects".format(self.__class__.__name__))
88
+ return super().__setattr__(key, val)
89
+
90
+ def __repr__(self):
91
+ items = []
92
+ for k in self.__keys:
93
+ items.append("{}={}".format(k, getattr(self, k)))
94
+ return "{}({})".format(self.__class__.__name__, ", ".join(items))
95
+
96
+
97
+ def get_worker_info() -> Optional[WorkerInfo]:
98
+ r"""Returns the information about the current
99
+ :class:`~torch.utils.data.DataLoader` iterator worker process.
100
+
101
+ When called in a worker, this returns an object guaranteed to have the
102
+ following attributes:
103
+
104
+ * :attr:`id`: the current worker id.
105
+ * :attr:`num_workers`: the total number of workers.
106
+ * :attr:`seed`: the random seed set for the current worker. This value is
107
+ determined by main process RNG and the worker id. See
108
+ :class:`~torch.utils.data.DataLoader`'s documentation for more details.
109
+ * :attr:`dataset`: the copy of the dataset object in **this** process. Note
110
+ that this will be a different object in a different process than the one
111
+ in the main process.
112
+
113
+ When called in the main process, this returns ``None``.
114
+
115
+ .. note::
116
+ When used in a :attr:`worker_init_fn` passed over to
117
+ :class:`~torch.utils.data.DataLoader`, this method can be useful to
118
+ set up each worker process differently, for instance, using ``worker_id``
119
+ to configure the ``dataset`` object to only read a specific fraction of a
120
+ sharded dataset, or use ``seed`` to seed other libraries used in dataset
121
+ code.
122
+ """
123
+ return _worker_info
124
+
125
+
126
+ r"""Dummy class used to signal the end of an IterableDataset"""
127
+
128
+
129
+ @dataclass(frozen=True)
130
+ class _IterableDatasetStopIteration:
131
+ worker_id: int
132
+
133
+
134
+ r"""Dummy class used to resume the fetching when worker reuse is enabled"""
135
+
136
+
137
+ @dataclass(frozen=True)
138
+ class _ResumeIteration:
139
+ seed: Optional[int] = None
140
+
141
+
142
+ # The function `_generate_state` is adapted from `numpy.random.SeedSequence`
143
+ # from https://github.com/numpy/numpy/blob/main/numpy/random/bit_generator.pyx
144
+ # It's MIT licensed, here is the copyright:
145
+
146
+ # Copyright (c) 2015 Melissa E. O'Neill
147
+ # Copyright (c) 2019 NumPy Developers
148
+ #
149
+ # Permission is hereby granted, free of charge, to any person obtaining a copy
150
+ # of this software and associated documentation files (the "Software"), to deal
151
+ # in the Software without restriction, including without limitation the rights
152
+ # to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
153
+ # copies of the Software, and to permit persons to whom the Software is
154
+ # furnished to do so, subject to the following conditions:
155
+ #
156
+ # The above copyright notice and this permission notice shall be included in
157
+ # all copies or substantial portions of the Software.
158
+ #
159
+ # THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
160
+ # IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
161
+ # FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
162
+ # AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
163
+ # LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
164
+ # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
165
+ # SOFTWARE.
166
+
167
+ # This function generates an array of int32 as the seed for
168
+ # `numpy.random`, in order to prevent state collision due to same
169
+ # seed and algorithm for `numpy.random` and `random` modules.
170
+ # TODO: Implement `SeedSequence` like object for `torch.random`
171
+ def _generate_state(base_seed, worker_id):
172
+ INIT_A = 0x43B0D7E5
173
+ MULT_A = 0x931E8875
174
+ INIT_B = 0x8B51F9DD
175
+ MULT_B = 0x58F38DED
176
+ MIX_MULT_L = 0xCA01F9DD
177
+ MIX_MULT_R = 0x4973F715
178
+ XSHIFT = 4 * 8 // 2
179
+ MASK32 = 0xFFFFFFFF
180
+
181
+ entropy = [worker_id, base_seed & MASK32, base_seed >> 32, 0]
182
+ pool = [0] * 4
183
+
184
+ hash_const_A = INIT_A
185
+
186
+ def hash(value):
187
+ nonlocal hash_const_A
188
+ value = (value ^ hash_const_A) & MASK32
189
+ hash_const_A = (hash_const_A * MULT_A) & MASK32
190
+ value = (value * hash_const_A) & MASK32
191
+ value = (value ^ (value >> XSHIFT)) & MASK32
192
+ return value
193
+
194
+ def mix(x, y):
195
+ result_x = (MIX_MULT_L * x) & MASK32
196
+ result_y = (MIX_MULT_R * y) & MASK32
197
+ result = (result_x - result_y) & MASK32
198
+ result = (result ^ (result >> XSHIFT)) & MASK32
199
+ return result
200
+
201
+ # Add in the entropy to the pool.
202
+ for i in range(len(pool)):
203
+ pool[i] = hash(entropy[i])
204
+
205
+ # Mix all bits together so late bits can affect earlier bits.
206
+ for i_src in range(len(pool)):
207
+ for i_dst in range(len(pool)):
208
+ if i_src != i_dst:
209
+ pool[i_dst] = mix(pool[i_dst], hash(pool[i_src]))
210
+
211
+ hash_const_B = INIT_B
212
+ state = []
213
+ for i_dst in range(4):
214
+ data_val = pool[i_dst]
215
+ data_val = (data_val ^ hash_const_B) & MASK32
216
+ hash_const_B = (hash_const_B * MULT_B) & MASK32
217
+ data_val = (data_val * hash_const_B) & MASK32
218
+ data_val = (data_val ^ (data_val >> XSHIFT)) & MASK32
219
+ state.append(data_val)
220
+ return state
221
+
222
+
223
+ def _worker_loop(
224
+ dataset_kind,
225
+ dataset,
226
+ index_queue,
227
+ data_queue,
228
+ done_event,
229
+ auto_collation,
230
+ collate_fn,
231
+ drop_last,
232
+ base_seed,
233
+ init_fn,
234
+ worker_id,
235
+ num_workers,
236
+ persistent_workers,
237
+ shared_seed,
238
+ ):
239
+ # See NOTE [ Data Loader Multiprocessing Shutdown Logic ] for details on the
240
+ # logic of this function.
241
+
242
+ try:
243
+ # Initialize C side signal handlers for SIGBUS and SIGSEGV. Python signal
244
+ # module's handlers are executed after Python returns from C low-level
245
+ # handlers, likely when the same fatal signal had already happened
246
+ # again.
247
+ # https://docs.python.org/3/library/signal.html#execution-of-python-signal-handlers
248
+ signal_handling._set_worker_signal_handlers()
249
+
250
+ torch.set_num_threads(1)
251
+ seed = base_seed + worker_id
252
+ random.seed(seed)
253
+ torch.manual_seed(seed)
254
+ if HAS_NUMPY:
255
+ np_seed = _generate_state(base_seed, worker_id)
256
+ import numpy as np
257
+
258
+ np.random.seed(np_seed)
259
+
260
+ from torch.utils.data import IterDataPipe
261
+ from torch.utils.data.graph_settings import apply_random_seed
262
+
263
+ shared_rng = torch.Generator()
264
+ if isinstance(dataset, IterDataPipe):
265
+ assert shared_seed is not None
266
+ shared_rng.manual_seed(shared_seed)
267
+ dataset = apply_random_seed(dataset, shared_rng)
268
+
269
+ global _worker_info
270
+ _worker_info = WorkerInfo(id=worker_id, num_workers=num_workers, seed=seed, dataset=dataset)
271
+
272
+ from torch.utils.data import _DatasetKind
273
+
274
+ init_exception = None
275
+
276
+ try:
277
+ if init_fn is not None:
278
+ init_fn(worker_id)
279
+
280
+ fetcher = _DatasetKind.create_fetcher(dataset_kind, dataset, auto_collation, collate_fn, drop_last)
281
+ except Exception:
282
+ init_exception = ExceptionWrapper(where="in DataLoader worker process {}".format(worker_id))
283
+
284
+ # When using Iterable mode, some worker can exit earlier than others due
285
+ # to the IterableDataset behaving differently for different workers.
286
+ # When such things happen, an `_IterableDatasetStopIteration` object is
287
+ # sent over to the main process with the ID of this worker, so that the
288
+ # main process won't send more tasks to this worker, and will send
289
+ # `None` to this worker to properly exit it.
290
+ #
291
+ # Note that we cannot set `done_event` from a worker as it is shared
292
+ # among all processes. Instead, we set the `iteration_end` flag to
293
+ # signify that the iterator is exhausted. When either `done_event` or
294
+ # `iteration_end` is set, we skip all processing step and just wait for
295
+ # `None`.
296
+ iteration_end = False
297
+
298
+ watchdog = ManagerWatchdog()
299
+
300
+ while watchdog.is_alive():
301
+ try:
302
+ r = index_queue.get(timeout=MP_STATUS_CHECK_INTERVAL)
303
+ except queue.Empty:
304
+ continue
305
+ if isinstance(r, _ResumeIteration):
306
+ # Acknowledge the main process
307
+ data_queue.put((r, None))
308
+ iteration_end = False
309
+
310
+ if isinstance(dataset, IterDataPipe):
311
+ assert r.seed is not None
312
+ shared_rng.manual_seed(r.seed)
313
+ dataset = apply_random_seed(dataset, shared_rng)
314
+
315
+ # Recreate the fetcher for worker-reuse policy
316
+ fetcher = _DatasetKind.create_fetcher(dataset_kind, dataset, auto_collation, collate_fn, drop_last)
317
+ continue
318
+ elif r is None:
319
+ # Received the final signal
320
+ assert done_event.is_set() or iteration_end
321
+ break
322
+ elif done_event.is_set() or iteration_end:
323
+ # `done_event` is set. But I haven't received the final signal
324
+ # (None) yet. I will keep continuing until get it, and skip the
325
+ # processing steps.
326
+ continue
327
+ idx, index = r
328
+ """ Added """
329
+ RRSController.sample_resolution(batch_id=idx)
330
+ """ Added """
331
+ data: Union[_IterableDatasetStopIteration, ExceptionWrapper]
332
+ if init_exception is not None:
333
+ data = init_exception
334
+ init_exception = None
335
+ else:
336
+ try:
337
+ data = fetcher.fetch(index)
338
+ except Exception as e:
339
+ if isinstance(e, StopIteration) and dataset_kind == _DatasetKind.Iterable:
340
+ data = _IterableDatasetStopIteration(worker_id)
341
+ # Set `iteration_end`
342
+ # (1) to save future `next(...)` calls, and
343
+ # (2) to avoid sending multiple `_IterableDatasetStopIteration`s.
344
+ iteration_end = True
345
+ else:
346
+ # It is important that we don't store exc_info in a variable.
347
+ # `ExceptionWrapper` does the correct thing.
348
+ # See NOTE [ Python Traceback Reference Cycle Problem ]
349
+ data = ExceptionWrapper(where="in DataLoader worker process {}".format(worker_id))
350
+ data_queue.put((idx, data))
351
+ del data, idx, index, r # save memory
352
+ except KeyboardInterrupt:
353
+ # Main process will raise KeyboardInterrupt anyways.
354
+ pass
355
+ if done_event.is_set():
356
+ data_queue.cancel_join_thread()
357
+ data_queue.close()
@@ -0,0 +1,100 @@
1
+ # EfficientViT: Multi-Scale Linear Attention for High-Resolution Dense Prediction
2
+ # Han Cai, Junyan Li, Muyan Hu, Chuang Gan, Song Han
3
+ # International Conference on Computer Vision (ICCV), 2023
4
+
5
+ import copy
6
+
7
+ import torch
8
+ import torchvision.transforms as transforms
9
+ import torchvision.transforms.functional as F
10
+
11
+ from ....models.utils import torch_random_choices
12
+
13
+ __all__ = [
14
+ "RRSController",
15
+ "get_interpolate",
16
+ "MyRandomResizedCrop",
17
+ ]
18
+
19
+
20
+ class RRSController:
21
+ ACTIVE_SIZE = (224, 224)
22
+ IMAGE_SIZE_LIST = [(224, 224)]
23
+
24
+ CHOICE_LIST = None
25
+
26
+ @staticmethod
27
+ def get_candidates() -> list[tuple[int, int]]:
28
+ return copy.deepcopy(RRSController.IMAGE_SIZE_LIST)
29
+
30
+ @staticmethod
31
+ def sample_resolution(batch_id: int) -> None:
32
+ RRSController.ACTIVE_SIZE = RRSController.CHOICE_LIST[batch_id]
33
+
34
+ @staticmethod
35
+ def set_epoch(epoch: int, batch_per_epoch: int) -> None:
36
+ g = torch.Generator()
37
+ g.manual_seed(epoch)
38
+ RRSController.CHOICE_LIST = torch_random_choices(
39
+ RRSController.get_candidates(),
40
+ g,
41
+ batch_per_epoch,
42
+ )
43
+
44
+
45
+ def get_interpolate(name: str) -> F.InterpolationMode:
46
+ mapping = {
47
+ "nearest": F.InterpolationMode.NEAREST,
48
+ "bilinear": F.InterpolationMode.BILINEAR,
49
+ "bicubic": F.InterpolationMode.BICUBIC,
50
+ "box": F.InterpolationMode.BOX,
51
+ "hamming": F.InterpolationMode.HAMMING,
52
+ "lanczos": F.InterpolationMode.LANCZOS,
53
+ }
54
+ if name in mapping:
55
+ return mapping[name]
56
+ elif name == "random":
57
+ return torch_random_choices(
58
+ [
59
+ F.InterpolationMode.NEAREST,
60
+ F.InterpolationMode.BILINEAR,
61
+ F.InterpolationMode.BICUBIC,
62
+ F.InterpolationMode.BOX,
63
+ F.InterpolationMode.HAMMING,
64
+ F.InterpolationMode.LANCZOS,
65
+ ],
66
+ )
67
+ else:
68
+ raise NotImplementedError
69
+
70
+
71
+ class MyRandomResizedCrop(transforms.RandomResizedCrop):
72
+ def __init__(
73
+ self,
74
+ scale=(0.08, 1.0),
75
+ ratio=(3.0 / 4.0, 4.0 / 3.0),
76
+ interpolation: str = "random",
77
+ ):
78
+ super(MyRandomResizedCrop, self).__init__(224, scale, ratio)
79
+ self.interpolation = interpolation
80
+
81
+ def forward(self, img: torch.Tensor) -> torch.Tensor:
82
+ i, j, h, w = self.get_params(img, list(self.scale), list(self.ratio))
83
+ target_size = RRSController.ACTIVE_SIZE
84
+ return F.resized_crop(
85
+ img,
86
+ i,
87
+ j,
88
+ h,
89
+ w,
90
+ list(target_size),
91
+ get_interpolate(self.interpolation),
92
+ )
93
+
94
+ def __repr__(self) -> str:
95
+ format_string = self.__class__.__name__
96
+ format_string += f"(\n\tsize={RRSController.get_candidates()},\n"
97
+ format_string += f"\tscale={tuple(round(s, 4) for s in self.scale)},\n"
98
+ format_string += f"\tratio={tuple(round(r, 4) for r in self.ratio)},\n"
99
+ format_string += f"\tinterpolation={self.interpolation})"
100
+ return format_string
@@ -0,0 +1,150 @@
1
+ # EfficientViT: Multi-Scale Linear Attention for High-Resolution Dense Prediction
2
+ # Han Cai, Junyan Li, Muyan Hu, Chuang Gan, Song Han
3
+ # International Conference on Computer Vision (ICCV), 2023
4
+
5
+ import os
6
+ import time
7
+ from copy import deepcopy
8
+
9
+ import torch.backends.cudnn
10
+ import torch.distributed
11
+ import torch.nn as nn
12
+ from torchpack import distributed as dist
13
+
14
+ from .data_provider import DataProvider
15
+ from .trainer.run_config import RunConfig
16
+ from .utils import (
17
+ dump_config,
18
+ init_modules,
19
+ load_config,
20
+ partial_update_config,
21
+ zero_last_gamma,
22
+ )
23
+ from ..models.utils import build_kwargs_from_config, load_state_dict_from_file
24
+
25
+ __all__ = [
26
+ "save_exp_config",
27
+ "setup_dist_env",
28
+ "setup_seed",
29
+ "setup_exp_config",
30
+ "setup_data_provider",
31
+ "setup_run_config",
32
+ "init_model",
33
+ ]
34
+
35
+
36
+ def save_exp_config(exp_config: dict, path: str, name="config.yaml") -> None:
37
+ if not dist.is_master():
38
+ return
39
+ dump_config(exp_config, os.path.join(path, name))
40
+
41
+
42
+ def setup_dist_env(gpu: str or None = None) -> None:
43
+ if gpu is not None:
44
+ os.environ["CUDA_VISIBLE_DEVICES"] = gpu
45
+ if not torch.distributed.is_initialized():
46
+ dist.init()
47
+ torch.backends.cudnn.benchmark = True
48
+ torch.cuda.set_device(dist.local_rank())
49
+
50
+
51
+ def setup_seed(manual_seed: int, resume: bool) -> None:
52
+ if resume:
53
+ manual_seed = int(time.time())
54
+ manual_seed = dist.rank() + manual_seed
55
+ torch.manual_seed(manual_seed)
56
+ torch.cuda.manual_seed_all(manual_seed)
57
+
58
+
59
+ def setup_exp_config(
60
+ config_path: str, recursive=True, opt_args: dict or None = None
61
+ ) -> dict:
62
+ # load config
63
+ if not os.path.isfile(config_path):
64
+ raise ValueError(config_path)
65
+
66
+ fpaths = [config_path]
67
+ if recursive:
68
+ extension = os.path.splitext(config_path)[1]
69
+ while os.path.dirname(config_path) != config_path:
70
+ config_path = os.path.dirname(config_path)
71
+ fpath = os.path.join(config_path, "default" + extension)
72
+ if os.path.isfile(fpath):
73
+ fpaths.append(fpath)
74
+ fpaths = fpaths[::-1]
75
+
76
+ default_config = load_config(fpaths[0])
77
+ exp_config = deepcopy(default_config)
78
+ for fpath in fpaths[1:]:
79
+ partial_update_config(exp_config, load_config(fpath))
80
+ # update config via args
81
+ if opt_args is not None:
82
+ partial_update_config(exp_config, opt_args)
83
+
84
+ return exp_config
85
+
86
+
87
+ def setup_data_provider(
88
+ exp_config: dict,
89
+ data_provider_classes: list[type[DataProvider]],
90
+ is_distributed: bool = True,
91
+ ) -> DataProvider:
92
+ dp_config = exp_config["data_provider"]
93
+ dp_config["num_replicas"] = dist.size() if is_distributed else None
94
+ dp_config["rank"] = dist.rank() if is_distributed else None
95
+ dp_config["test_batch_size"] = (
96
+ dp_config.get("test_batch_size", None)
97
+ or dp_config["base_batch_size"] * 2
98
+ )
99
+ dp_config["batch_size"] = dp_config["train_batch_size"] = dp_config[
100
+ "base_batch_size"
101
+ ]
102
+
103
+ data_provider_lookup = {
104
+ provider.name: provider for provider in data_provider_classes
105
+ }
106
+ data_provider_class = data_provider_lookup[dp_config["dataset"]]
107
+
108
+ data_provider_kwargs = build_kwargs_from_config(
109
+ dp_config, data_provider_class
110
+ )
111
+ data_provider = data_provider_class(**data_provider_kwargs)
112
+ return data_provider
113
+
114
+
115
+ def setup_run_config(
116
+ exp_config: dict, run_config_cls: type[RunConfig]
117
+ ) -> RunConfig:
118
+ exp_config["run_config"]["init_lr"] = (
119
+ exp_config["run_config"]["base_lr"] * dist.size()
120
+ )
121
+
122
+ run_config = run_config_cls(**exp_config["run_config"])
123
+
124
+ return run_config
125
+
126
+
127
+ def init_model(
128
+ network: nn.Module,
129
+ init_from: str or None = None,
130
+ backbone_init_from: str or None = None,
131
+ rand_init="trunc_normal",
132
+ last_gamma=None,
133
+ ) -> None:
134
+ # initialization
135
+ init_modules(network, init_type=rand_init)
136
+ # zero gamma of last bn in each block
137
+ if last_gamma is not None:
138
+ zero_last_gamma(network, last_gamma)
139
+
140
+ # load weight
141
+ if init_from is not None and os.path.isfile(init_from):
142
+ network.load_state_dict(load_state_dict_from_file(init_from))
143
+ print(f"Loaded init from {init_from}")
144
+ elif backbone_init_from is not None and os.path.isfile(backbone_init_from):
145
+ network.backbone.load_state_dict(
146
+ load_state_dict_from_file(backbone_init_from)
147
+ )
148
+ print(f"Loaded backbone init from {backbone_init_from}")
149
+ else:
150
+ print(f"Random init ({rand_init}) with last gamma {last_gamma}")
@@ -0,0 +1,6 @@
1
+ # EfficientViT: Multi-Scale Linear Attention for High-Resolution Dense Prediction
2
+ # Han Cai, Junyan Li, Muyan Hu, Chuang Gan, Song Han
3
+ # International Conference on Computer Vision (ICCV), 2023
4
+
5
+ from .base import *
6
+ from .run_config import *