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