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,78 @@
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 numpy as np
6
+ import torchvision.transforms as transforms
7
+ from PIL import Image
8
+ from timm.data.auto_augment import rand_augment_transform
9
+
10
+ __all__ = ["ColorAug", "RandAug"]
11
+
12
+
13
+ class ImageAug:
14
+ def aug_image(self, image: Image.Image) -> Image.Image:
15
+ raise NotImplementedError
16
+
17
+ def __call__(self, feed_dict: dict or np.ndarray or Image.Image) -> dict or np.ndarray or Image.Image:
18
+ if isinstance(feed_dict, dict):
19
+ output_dict = feed_dict
20
+ image = feed_dict[self.key]
21
+ else:
22
+ output_dict = None
23
+ image = feed_dict
24
+ is_ndarray = isinstance(image, np.ndarray)
25
+ if is_ndarray:
26
+ image = Image.fromarray(image)
27
+
28
+ image = self.aug_image(image)
29
+
30
+ if is_ndarray:
31
+ image = np.array(image)
32
+
33
+ if output_dict is None:
34
+ return image
35
+ else:
36
+ output_dict[self.key] = image
37
+ return output_dict
38
+
39
+
40
+ class ColorAug(transforms.ColorJitter, ImageAug):
41
+ def __init__(self, brightness=0, contrast=0, saturation=0, hue=0, key="data"):
42
+ super().__init__(
43
+ brightness=brightness,
44
+ contrast=contrast,
45
+ saturation=saturation,
46
+ hue=hue,
47
+ )
48
+ self.key = key
49
+
50
+ def aug_image(self, image: Image.Image) -> Image.Image:
51
+ return transforms.ColorJitter.forward(self, image)
52
+
53
+ def forward(self, feed_dict: dict or np.ndarray or Image.Image) -> dict or np.ndarray or Image.Image:
54
+ return ImageAug.__call__(self, feed_dict)
55
+
56
+
57
+ class RandAug(ImageAug):
58
+ def __init__(self, config: dict[str, any], mean: tuple[float, float, float], key="data"):
59
+ n = config.get("n", 2)
60
+ m = config.get("m", 9)
61
+ mstd = config.get("mstd", 1.0)
62
+ inc = config.get("inc", 1)
63
+ tpct = config.get("tpct", 0.45)
64
+ config_str = f"rand-n{n}-m{m}-mstd{mstd}-inc{inc}"
65
+
66
+ aa_params = dict(
67
+ translate_pct=tpct,
68
+ img_mean=tuple([min(255, round(255 * x)) for x in mean]),
69
+ interpolation=Image.BICUBIC,
70
+ )
71
+ self.aug_op = rand_augment_transform(config_str, aa_params)
72
+ self.key = key
73
+
74
+ def aug_image(self, image: Image.Image) -> Image.Image:
75
+ return self.aug_op(image)
76
+
77
+ def __repr__(self):
78
+ return self.aug_op.__repr__()
@@ -0,0 +1,254 @@
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
+ import warnings
7
+
8
+ import torch.utils.data
9
+ from torch.utils.data.distributed import DistributedSampler
10
+
11
+ from .random_resolution import RRSController
12
+ from ...models.utils import val2tuple
13
+
14
+ __all__ = ["parse_image_size", "random_drop_data", "DataProvider"]
15
+
16
+
17
+ def parse_image_size(size: int or str) -> tuple[int, int]:
18
+ if isinstance(size, str):
19
+ size = [int(val) for val in size.split("-")]
20
+ return size[0], size[1]
21
+ else:
22
+ return val2tuple(size, 2)
23
+
24
+
25
+ def random_drop_data(dataset, drop_size: int, seed: int, keys=("samples",)):
26
+ g = torch.Generator()
27
+ g.manual_seed(seed) # set random seed before sampling validation set
28
+ rand_indexes = torch.randperm(len(dataset), generator=g).tolist()
29
+
30
+ dropped_indexes = rand_indexes[:drop_size]
31
+ remaining_indexes = rand_indexes[drop_size:]
32
+
33
+ dropped_dataset = copy.deepcopy(dataset)
34
+ for key in keys:
35
+ setattr(
36
+ dropped_dataset,
37
+ key,
38
+ [getattr(dropped_dataset, key)[idx] for idx in dropped_indexes],
39
+ )
40
+ setattr(
41
+ dataset,
42
+ key,
43
+ [getattr(dataset, key)[idx] for idx in remaining_indexes],
44
+ )
45
+ return dataset, dropped_dataset
46
+
47
+
48
+ class DataProvider:
49
+ data_keys = ("samples",)
50
+ mean_std = {"mean": [0.485, 0.456, 0.406], "std": [0.229, 0.224, 0.225]}
51
+ SUB_SEED = 937162211 # random seed for sampling subset
52
+ VALID_SEED = 2147483647 # random seed for the validation set
53
+
54
+ name: str
55
+
56
+ def __init__(
57
+ self,
58
+ train_batch_size: int,
59
+ test_batch_size: int or None,
60
+ valid_size: int or float or None,
61
+ n_worker: int,
62
+ image_size: int or list[int] or str or list[str],
63
+ num_replicas: int or None = None,
64
+ rank: int or None = None,
65
+ train_ratio: float or None = None,
66
+ drop_last: bool = False,
67
+ ):
68
+ warnings.filterwarnings("ignore")
69
+ super().__init__()
70
+
71
+ # batch_size & valid_size
72
+ self.train_batch_size = train_batch_size
73
+ self.test_batch_size = test_batch_size or self.train_batch_size
74
+ self.valid_size = valid_size
75
+
76
+ # image size
77
+ if isinstance(image_size, list):
78
+ self.image_size = [parse_image_size(size) for size in image_size]
79
+ self.image_size.sort() # e.g., 160 -> 224
80
+ RRSController.IMAGE_SIZE_LIST = copy.deepcopy(self.image_size)
81
+ self.active_image_size = RRSController.ACTIVE_SIZE = (
82
+ self.image_size[-1]
83
+ )
84
+ else:
85
+ self.image_size = parse_image_size(image_size)
86
+ RRSController.IMAGE_SIZE_LIST = [self.image_size]
87
+ self.active_image_size = RRSController.ACTIVE_SIZE = (
88
+ self.image_size
89
+ )
90
+
91
+ # distributed configs
92
+ self.num_replicas = num_replicas
93
+ self.rank = rank
94
+
95
+ # build datasets
96
+ train_dataset, val_dataset, test_dataset = self.build_datasets()
97
+
98
+ if train_ratio is not None and train_ratio < 1.0:
99
+ assert 0 < train_ratio < 1
100
+ _, train_dataset = random_drop_data(
101
+ train_dataset,
102
+ int(train_ratio * len(train_dataset)),
103
+ self.SUB_SEED,
104
+ self.data_keys,
105
+ )
106
+
107
+ # build data loader
108
+ self.train = self.build_dataloader(
109
+ train_dataset,
110
+ train_batch_size,
111
+ n_worker,
112
+ drop_last=drop_last,
113
+ train=True,
114
+ )
115
+ self.valid = self.build_dataloader(
116
+ val_dataset,
117
+ test_batch_size,
118
+ n_worker,
119
+ drop_last=False,
120
+ train=False,
121
+ )
122
+ self.test = self.build_dataloader(
123
+ test_dataset,
124
+ test_batch_size,
125
+ n_worker,
126
+ drop_last=False,
127
+ train=False,
128
+ )
129
+ if self.valid is None:
130
+ self.valid = self.test
131
+ self.sub_train = None
132
+
133
+ @property
134
+ def data_shape(self) -> tuple[int, ...]:
135
+ return 3, self.active_image_size[0], self.active_image_size[1]
136
+
137
+ def build_valid_transform(
138
+ self, image_size: tuple[int, int] or None = None
139
+ ) -> any:
140
+ raise NotImplementedError
141
+
142
+ def build_train_transform(
143
+ self, image_size: tuple[int, int] or None = None
144
+ ) -> any:
145
+ raise NotImplementedError
146
+
147
+ def build_datasets(self) -> tuple[any, any, any]:
148
+ raise NotImplementedError
149
+
150
+ def build_dataloader(
151
+ self,
152
+ dataset: any or None,
153
+ batch_size: int,
154
+ n_worker: int,
155
+ drop_last: bool,
156
+ train: bool,
157
+ ):
158
+ if dataset is None:
159
+ return None
160
+ if isinstance(self.image_size, list) and train:
161
+ from efficientvit.apps.data_provider.random_resolution._data_loader import (
162
+ RRSDataLoader,
163
+ )
164
+
165
+ dataloader_class = RRSDataLoader
166
+ else:
167
+ dataloader_class = torch.utils.data.DataLoader
168
+ if self.num_replicas is None:
169
+ return dataloader_class(
170
+ dataset=dataset,
171
+ batch_size=batch_size,
172
+ shuffle=True,
173
+ num_workers=n_worker,
174
+ pin_memory=True,
175
+ drop_last=drop_last,
176
+ )
177
+ else:
178
+ sampler = DistributedSampler(dataset, self.num_replicas, self.rank)
179
+ return dataloader_class(
180
+ dataset=dataset,
181
+ batch_size=batch_size,
182
+ sampler=sampler,
183
+ num_workers=n_worker,
184
+ pin_memory=True,
185
+ drop_last=drop_last,
186
+ )
187
+
188
+ def set_epoch(self, epoch: int) -> None:
189
+ RRSController.set_epoch(epoch, len(self.train))
190
+ if isinstance(self.train.sampler, DistributedSampler):
191
+ self.train.sampler.set_epoch(epoch)
192
+
193
+ def assign_active_image_size(
194
+ self, new_size: int or tuple[int, int]
195
+ ) -> None:
196
+ self.active_image_size = val2tuple(new_size, 2)
197
+ new_transform = self.build_valid_transform(self.active_image_size)
198
+ # change the transform of the valid and test set
199
+ self.valid.dataset.transform = self.test.dataset.transform = (
200
+ new_transform
201
+ )
202
+
203
+ def sample_val_dataset(
204
+ self, train_dataset, valid_transform
205
+ ) -> tuple[any, any]:
206
+ if self.valid_size is not None:
207
+ if 0 < self.valid_size < 1:
208
+ valid_size = int(self.valid_size * len(train_dataset))
209
+ else:
210
+ assert self.valid_size >= 1
211
+ valid_size = int(self.valid_size)
212
+ train_dataset, val_dataset = random_drop_data(
213
+ train_dataset,
214
+ valid_size,
215
+ self.VALID_SEED,
216
+ self.data_keys,
217
+ )
218
+ val_dataset.transform = valid_transform
219
+ else:
220
+ val_dataset = None
221
+ return train_dataset, val_dataset
222
+
223
+ def build_sub_train_loader(self, n_samples: int, batch_size: int) -> any:
224
+ # used for resetting BN running statistics
225
+ if self.sub_train is None:
226
+ self.sub_train = {}
227
+ if self.active_image_size in self.sub_train:
228
+ return self.sub_train[self.active_image_size]
229
+
230
+ # construct dataset and dataloader
231
+ train_dataset = copy.deepcopy(self.train.dataset)
232
+ if n_samples < len(train_dataset):
233
+ _, train_dataset = random_drop_data(
234
+ train_dataset,
235
+ n_samples,
236
+ self.SUB_SEED,
237
+ self.data_keys,
238
+ )
239
+ RRSController.ACTIVE_SIZE = self.active_image_size
240
+ train_dataset.transform = self.build_train_transform(
241
+ image_size=self.active_image_size
242
+ )
243
+ data_loader = self.build_dataloader(
244
+ train_dataset, batch_size, self.train.num_workers, True, False
245
+ )
246
+
247
+ # pre-fetch data
248
+ self.sub_train[self.active_image_size] = [
249
+ data
250
+ for data in data_loader
251
+ for _ in range(max(1, n_samples // len(train_dataset)))
252
+ ]
253
+
254
+ return self.sub_train[self.active_image_size]
@@ -0,0 +1,6 @@
1
+ """Random resolution data loader compatible with multi-processing and distributed training.
2
+
3
+ Replace Pytorch's DataLoader with RRSDataLoader to support random resolution
4
+ at the training time, resolution sampling is controlled by RRSController
5
+ """
6
+ from .controller import *