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,52 @@
|
|
|
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 math
|
|
6
|
+
|
|
7
|
+
import torch
|
|
8
|
+
|
|
9
|
+
from ...models.utils.list import val2list
|
|
10
|
+
|
|
11
|
+
__all__ = ["CosineLRwithWarmup"]
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class CosineLRwithWarmup(torch.optim.lr_scheduler._LRScheduler):
|
|
15
|
+
def __init__(
|
|
16
|
+
self,
|
|
17
|
+
optimizer: torch.optim.Optimizer,
|
|
18
|
+
warmup_steps: int,
|
|
19
|
+
warmup_lr: float,
|
|
20
|
+
decay_steps: int or list[int],
|
|
21
|
+
last_epoch: int = -1,
|
|
22
|
+
) -> None:
|
|
23
|
+
self.warmup_steps = warmup_steps
|
|
24
|
+
self.warmup_lr = warmup_lr
|
|
25
|
+
self.decay_steps = val2list(decay_steps)
|
|
26
|
+
super().__init__(optimizer, last_epoch)
|
|
27
|
+
|
|
28
|
+
def get_lr(self) -> list[float]:
|
|
29
|
+
if self.last_epoch < self.warmup_steps:
|
|
30
|
+
return [
|
|
31
|
+
(base_lr - self.warmup_lr)
|
|
32
|
+
* (self.last_epoch + 1)
|
|
33
|
+
/ self.warmup_steps
|
|
34
|
+
+ self.warmup_lr
|
|
35
|
+
for base_lr in self.base_lrs
|
|
36
|
+
]
|
|
37
|
+
else:
|
|
38
|
+
current_steps = self.last_epoch - self.warmup_steps
|
|
39
|
+
decay_steps = [0] + self.decay_steps
|
|
40
|
+
idx = len(decay_steps) - 2
|
|
41
|
+
for i, decay_step in enumerate(decay_steps[:-1]):
|
|
42
|
+
if decay_step <= current_steps < decay_steps[i + 1]:
|
|
43
|
+
idx = i
|
|
44
|
+
break
|
|
45
|
+
current_steps -= decay_steps[idx]
|
|
46
|
+
decay_step = decay_steps[idx + 1] - decay_steps[idx]
|
|
47
|
+
return [
|
|
48
|
+
0.5
|
|
49
|
+
* base_lr
|
|
50
|
+
* (1 + math.cos(math.pi * current_steps / decay_step))
|
|
51
|
+
for base_lr in self.base_lrs
|
|
52
|
+
]
|
|
@@ -0,0 +1,43 @@
|
|
|
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 torch
|
|
6
|
+
|
|
7
|
+
from ..utils.dist import sync_tensor
|
|
8
|
+
|
|
9
|
+
__all__ = ["AverageMeter"]
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class AverageMeter:
|
|
13
|
+
"""Computes and stores the average and current value."""
|
|
14
|
+
|
|
15
|
+
def __init__(self, is_distributed=True):
|
|
16
|
+
self.is_distributed = is_distributed
|
|
17
|
+
self.sum = 0
|
|
18
|
+
self.count = 0
|
|
19
|
+
|
|
20
|
+
def _sync(
|
|
21
|
+
self, val: torch.Tensor or int or float
|
|
22
|
+
) -> torch.Tensor or int or float:
|
|
23
|
+
return sync_tensor(val, reduce="sum") if self.is_distributed else val
|
|
24
|
+
|
|
25
|
+
def update(self, val: torch.Tensor or int or float, delta_n=1):
|
|
26
|
+
self.count += self._sync(delta_n)
|
|
27
|
+
self.sum += self._sync(val * delta_n)
|
|
28
|
+
|
|
29
|
+
def get_count(self) -> torch.Tensor or int or float:
|
|
30
|
+
return (
|
|
31
|
+
self.count.item()
|
|
32
|
+
if isinstance(self.count, torch.Tensor) and self.count.numel() == 1
|
|
33
|
+
else self.count
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
@property
|
|
37
|
+
def avg(self):
|
|
38
|
+
avg = -1 if self.count == 0 else self.sum / self.count
|
|
39
|
+
return (
|
|
40
|
+
avg.item()
|
|
41
|
+
if isinstance(avg, torch.Tensor) and avg.numel() == 1
|
|
42
|
+
else avg
|
|
43
|
+
)
|
|
@@ -0,0 +1,101 @@
|
|
|
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
|
+
|
|
7
|
+
import yaml
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"parse_with_yaml",
|
|
11
|
+
"parse_unknown_args",
|
|
12
|
+
"partial_update_config",
|
|
13
|
+
"resolve_and_load_config",
|
|
14
|
+
"load_config",
|
|
15
|
+
"dump_config",
|
|
16
|
+
]
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def parse_with_yaml(config_str: str) -> str or dict:
|
|
20
|
+
try:
|
|
21
|
+
# add space manually for dict
|
|
22
|
+
if "{" in config_str and "}" in config_str and ":" in config_str:
|
|
23
|
+
out_str = config_str.replace(":", ": ")
|
|
24
|
+
else:
|
|
25
|
+
out_str = config_str
|
|
26
|
+
return yaml.safe_load(out_str)
|
|
27
|
+
except ValueError:
|
|
28
|
+
# return raw string if parsing fails
|
|
29
|
+
return config_str
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def parse_unknown_args(unknown: list) -> dict:
|
|
33
|
+
"""Parse unknown args."""
|
|
34
|
+
index = 0
|
|
35
|
+
parsed_dict = {}
|
|
36
|
+
while index < len(unknown):
|
|
37
|
+
key, val = unknown[index], unknown[index + 1]
|
|
38
|
+
index += 2
|
|
39
|
+
if not key.startswith("--"):
|
|
40
|
+
continue
|
|
41
|
+
key = key[2:]
|
|
42
|
+
|
|
43
|
+
# try parsing with either dot notation or full yaml notation
|
|
44
|
+
# Note that the vanilla case "--key value" will be parsed the same
|
|
45
|
+
if "." in key:
|
|
46
|
+
# key == a.b.c, val == val --> parsed_dict[a][b][c] = val
|
|
47
|
+
keys = key.split(".")
|
|
48
|
+
dict_to_update = parsed_dict
|
|
49
|
+
for key in keys[:-1]:
|
|
50
|
+
if not (key in dict_to_update and isinstance(dict_to_update[key], dict)):
|
|
51
|
+
dict_to_update[key] = {}
|
|
52
|
+
dict_to_update = dict_to_update[key]
|
|
53
|
+
dict_to_update[keys[-1]] = parse_with_yaml(val) # so we can parse lists, bools, etc...
|
|
54
|
+
else:
|
|
55
|
+
parsed_dict[key] = parse_with_yaml(val)
|
|
56
|
+
return parsed_dict
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def partial_update_config(config: dict, partial_config: dict) -> dict:
|
|
60
|
+
for key in partial_config:
|
|
61
|
+
if key in config and isinstance(partial_config[key], dict) and isinstance(config[key], dict):
|
|
62
|
+
partial_update_config(config[key], partial_config[key])
|
|
63
|
+
else:
|
|
64
|
+
config[key] = partial_config[key]
|
|
65
|
+
return config
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def resolve_and_load_config(path: str, config_name="config.yaml") -> dict:
|
|
69
|
+
path = os.path.realpath(os.path.expanduser(path))
|
|
70
|
+
if os.path.isdir(path):
|
|
71
|
+
config_path = os.path.join(path, config_name)
|
|
72
|
+
else:
|
|
73
|
+
config_path = path
|
|
74
|
+
if os.path.isfile(config_path):
|
|
75
|
+
pass
|
|
76
|
+
else:
|
|
77
|
+
raise Exception(f"Cannot find a valid config at {path}")
|
|
78
|
+
config = load_config(config_path)
|
|
79
|
+
return config
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
class SafeLoaderWithTuple(yaml.SafeLoader):
|
|
83
|
+
"""A yaml safe loader with python tuple loading capabilities."""
|
|
84
|
+
|
|
85
|
+
def construct_python_tuple(self, node):
|
|
86
|
+
return tuple(self.construct_sequence(node))
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
SafeLoaderWithTuple.add_constructor("tag:yaml.org,2002:python/tuple", SafeLoaderWithTuple.construct_python_tuple)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def load_config(filename: str) -> dict:
|
|
93
|
+
"""Load a yaml file."""
|
|
94
|
+
filename = os.path.realpath(os.path.expanduser(filename))
|
|
95
|
+
return yaml.load(open(filename), Loader=SafeLoaderWithTuple)
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def dump_config(config: dict, filename: str) -> None:
|
|
99
|
+
"""Dump a config file"""
|
|
100
|
+
filename = os.path.realpath(os.path.expanduser(filename))
|
|
101
|
+
yaml.dump(config, open(filename, "w"), sort_keys=False)
|
|
@@ -0,0 +1,28 @@
|
|
|
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 torch
|
|
6
|
+
|
|
7
|
+
__all__ = ["REGISTERED_OPTIMIZER_DICT", "build_optimizer"]
|
|
8
|
+
|
|
9
|
+
# register optimizer here
|
|
10
|
+
# name: optimizer, kwargs with default values
|
|
11
|
+
REGISTERED_OPTIMIZER_DICT: dict[str, tuple[type, dict[str, any]]] = {
|
|
12
|
+
"sgd": (torch.optim.SGD, {"momentum": 0.9, "nesterov": True}),
|
|
13
|
+
"adam": (torch.optim.Adam, {"betas": (0.9, 0.999), "eps": 1e-8, "amsgrad": False}),
|
|
14
|
+
"adamw": (torch.optim.AdamW, {"betas": (0.9, 0.999), "eps": 1e-8, "amsgrad": False}),
|
|
15
|
+
}
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def build_optimizer(
|
|
19
|
+
net_params, optimizer_name: str, optimizer_params: dict or None, init_lr: float
|
|
20
|
+
) -> torch.optim.Optimizer:
|
|
21
|
+
optimizer_class, default_params = REGISTERED_OPTIMIZER_DICT[optimizer_name]
|
|
22
|
+
optimizer_params = optimizer_params or {}
|
|
23
|
+
|
|
24
|
+
for key in default_params:
|
|
25
|
+
if key in optimizer_params:
|
|
26
|
+
default_params[key] = optimizer_params[key]
|
|
27
|
+
optimizer = optimizer_class(net_params, init_lr, **default_params)
|
|
28
|
+
return optimizer
|
|
@@ -0,0 +1,79 @@
|
|
|
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 efficientvit.models.efficientvit import (
|
|
6
|
+
EfficientViTCls,
|
|
7
|
+
efficientvit_cls_b0,
|
|
8
|
+
efficientvit_cls_b1,
|
|
9
|
+
efficientvit_cls_b2,
|
|
10
|
+
efficientvit_cls_b3,
|
|
11
|
+
efficientvit_cls_l1,
|
|
12
|
+
efficientvit_cls_l2,
|
|
13
|
+
efficientvit_cls_l3,
|
|
14
|
+
)
|
|
15
|
+
from efficientvit.models.nn.norm import set_norm_eps
|
|
16
|
+
from efficientvit.models.utils import load_state_dict_from_file
|
|
17
|
+
|
|
18
|
+
__all__ = ["create_cls_model"]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
REGISTERED_CLS_MODEL: dict[str, str] = {
|
|
22
|
+
"b0-r224": "assets/checkpoints/cls/b0-r224.pt",
|
|
23
|
+
###############################################
|
|
24
|
+
"b1-r224": "assets/checkpoints/cls/b1-r224.pt",
|
|
25
|
+
"b1-r256": "assets/checkpoints/cls/b1-r256.pt",
|
|
26
|
+
"b1-r288": "assets/checkpoints/cls/b1-r288.pt",
|
|
27
|
+
###############################################
|
|
28
|
+
"b2-r224": "assets/checkpoints/cls/b2-r224.pt",
|
|
29
|
+
"b2-r256": "assets/checkpoints/cls/b2-r256.pt",
|
|
30
|
+
"b2-r288": "assets/checkpoints/cls/b2-r288.pt",
|
|
31
|
+
###############################################
|
|
32
|
+
"b3-r224": "assets/checkpoints/cls/b3-r224.pt",
|
|
33
|
+
"b3-r256": "assets/checkpoints/cls/b3-r256.pt",
|
|
34
|
+
"b3-r288": "assets/checkpoints/cls/b3-r288.pt",
|
|
35
|
+
###############################################
|
|
36
|
+
"l1-r224": "assets/checkpoints/cls/l1-r224.pt",
|
|
37
|
+
###############################################
|
|
38
|
+
"l2-r224": "assets/checkpoints/cls/l2-r224.pt",
|
|
39
|
+
"l2-r256": "assets/checkpoints/cls/l2-r256.pt",
|
|
40
|
+
"l2-r288": "assets/checkpoints/cls/l2-r288.pt",
|
|
41
|
+
"l2-r320": "assets/checkpoints/cls/l2-r320.pt",
|
|
42
|
+
"l2-r384": "assets/checkpoints/cls/l2-r384.pt",
|
|
43
|
+
###############################################
|
|
44
|
+
"l3-r224": "assets/checkpoints/cls/l3-r224.pt",
|
|
45
|
+
"l3-r256": "assets/checkpoints/cls/l3-r256.pt",
|
|
46
|
+
"l3-r288": "assets/checkpoints/cls/l3-r288.pt",
|
|
47
|
+
"l3-r320": "assets/checkpoints/cls/l3-r320.pt",
|
|
48
|
+
"l3-r384": "assets/checkpoints/cls/l3-r384.pt",
|
|
49
|
+
}
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def create_cls_model(name: str, pretrained=True, weight_url: str or None = None, **kwargs) -> EfficientViTCls:
|
|
53
|
+
model_dict = {
|
|
54
|
+
"b0": efficientvit_cls_b0,
|
|
55
|
+
"b1": efficientvit_cls_b1,
|
|
56
|
+
"b2": efficientvit_cls_b2,
|
|
57
|
+
"b3": efficientvit_cls_b3,
|
|
58
|
+
#########################
|
|
59
|
+
"l1": efficientvit_cls_l1,
|
|
60
|
+
"l2": efficientvit_cls_l2,
|
|
61
|
+
"l3": efficientvit_cls_l3,
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
model_id = name.split("-")[0]
|
|
65
|
+
if model_id not in model_dict:
|
|
66
|
+
raise ValueError(f"Do not find {name} in the model zoo. List of models: {list(model_dict.keys())}")
|
|
67
|
+
else:
|
|
68
|
+
model = model_dict[model_id](**kwargs)
|
|
69
|
+
if model_id in ["l1", "l2", "l3"]:
|
|
70
|
+
set_norm_eps(model, 1e-7)
|
|
71
|
+
|
|
72
|
+
if pretrained:
|
|
73
|
+
weight_url = weight_url or REGISTERED_CLS_MODEL.get(name, None)
|
|
74
|
+
if weight_url is None:
|
|
75
|
+
raise ValueError(f"Do not find the pretrained weight of {name}.")
|
|
76
|
+
else:
|
|
77
|
+
weight = load_state_dict_from_file(weight_url)
|
|
78
|
+
model.load_state_dict(weight)
|
|
79
|
+
return model
|
|
File without changes
|
|
@@ -0,0 +1,142 @@
|
|
|
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 math
|
|
7
|
+
import os
|
|
8
|
+
|
|
9
|
+
import torchvision.transforms as transforms
|
|
10
|
+
from torchvision.datasets import ImageFolder
|
|
11
|
+
|
|
12
|
+
from ...apps.data_provider import DataProvider
|
|
13
|
+
from ...apps.data_provider.augment import RandAug
|
|
14
|
+
from ...apps.data_provider.random_resolution import (
|
|
15
|
+
MyRandomResizedCrop,
|
|
16
|
+
get_interpolate,
|
|
17
|
+
)
|
|
18
|
+
from ...apps.utils import partial_update_config
|
|
19
|
+
from ...models.utils import val2list
|
|
20
|
+
|
|
21
|
+
__all__ = ["ImageNetDataProvider"]
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class ImageNetDataProvider(DataProvider):
|
|
25
|
+
name = "imagenet"
|
|
26
|
+
|
|
27
|
+
data_dir = "/dataset/imagenet"
|
|
28
|
+
n_classes = 1000
|
|
29
|
+
_DEFAULT_RRC_CONFIG = {
|
|
30
|
+
"train_interpolate": "random",
|
|
31
|
+
"test_interpolate": "bicubic",
|
|
32
|
+
"test_crop_ratio": 1.0,
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
def __init__(
|
|
36
|
+
self,
|
|
37
|
+
data_dir: str or None = None,
|
|
38
|
+
rrc_config: dict or None = None,
|
|
39
|
+
data_aug: dict or list[dict] or None = None,
|
|
40
|
+
###########################################
|
|
41
|
+
train_batch_size=128,
|
|
42
|
+
test_batch_size=128,
|
|
43
|
+
valid_size: int or float or None = None,
|
|
44
|
+
n_worker=8,
|
|
45
|
+
image_size: int or list[int] = 224,
|
|
46
|
+
num_replicas: int or None = None,
|
|
47
|
+
rank: int or None = None,
|
|
48
|
+
train_ratio: float or None = None,
|
|
49
|
+
drop_last: bool = False,
|
|
50
|
+
):
|
|
51
|
+
self.data_dir = data_dir or self.data_dir
|
|
52
|
+
self.rrc_config = partial_update_config(
|
|
53
|
+
copy.deepcopy(self._DEFAULT_RRC_CONFIG),
|
|
54
|
+
rrc_config or {},
|
|
55
|
+
)
|
|
56
|
+
self.data_aug = data_aug
|
|
57
|
+
|
|
58
|
+
super().__init__(
|
|
59
|
+
train_batch_size,
|
|
60
|
+
test_batch_size,
|
|
61
|
+
valid_size,
|
|
62
|
+
n_worker,
|
|
63
|
+
image_size,
|
|
64
|
+
num_replicas,
|
|
65
|
+
rank,
|
|
66
|
+
train_ratio,
|
|
67
|
+
drop_last,
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
def build_valid_transform(
|
|
71
|
+
self, image_size: tuple[int, int] or None = None
|
|
72
|
+
) -> any:
|
|
73
|
+
image_size = (image_size or self.active_image_size)[0]
|
|
74
|
+
crop_size = int(
|
|
75
|
+
math.ceil(image_size / self.rrc_config["test_crop_ratio"])
|
|
76
|
+
)
|
|
77
|
+
return transforms.Compose(
|
|
78
|
+
[
|
|
79
|
+
transforms.Resize(
|
|
80
|
+
crop_size,
|
|
81
|
+
interpolation=get_interpolate(
|
|
82
|
+
self.rrc_config["test_interpolate"]
|
|
83
|
+
),
|
|
84
|
+
),
|
|
85
|
+
transforms.CenterCrop(image_size),
|
|
86
|
+
transforms.ToTensor(),
|
|
87
|
+
transforms.Normalize(**self.mean_std),
|
|
88
|
+
]
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
def build_train_transform(
|
|
92
|
+
self, image_size: tuple[int, int] or None = None
|
|
93
|
+
) -> any:
|
|
94
|
+
image_size = image_size or self.image_size
|
|
95
|
+
|
|
96
|
+
# random_resize_crop -> random_horizontal_flip
|
|
97
|
+
train_transforms = [
|
|
98
|
+
MyRandomResizedCrop(
|
|
99
|
+
interpolation=self.rrc_config["train_interpolate"]
|
|
100
|
+
),
|
|
101
|
+
transforms.RandomHorizontalFlip(),
|
|
102
|
+
]
|
|
103
|
+
|
|
104
|
+
# data augmentation
|
|
105
|
+
post_aug = []
|
|
106
|
+
if self.data_aug is not None:
|
|
107
|
+
for aug_op in val2list(self.data_aug):
|
|
108
|
+
if aug_op["name"] == "randaug":
|
|
109
|
+
data_aug = RandAug(aug_op, mean=self.mean_std["mean"])
|
|
110
|
+
elif aug_op["name"] == "erase":
|
|
111
|
+
from timm.data.random_erasing import RandomErasing
|
|
112
|
+
|
|
113
|
+
random_erase = RandomErasing(aug_op["p"], device="cpu")
|
|
114
|
+
post_aug.append(random_erase)
|
|
115
|
+
data_aug = None
|
|
116
|
+
else:
|
|
117
|
+
raise NotImplementedError
|
|
118
|
+
if data_aug is not None:
|
|
119
|
+
train_transforms.append(data_aug)
|
|
120
|
+
train_transforms = [
|
|
121
|
+
*train_transforms,
|
|
122
|
+
transforms.ToTensor(),
|
|
123
|
+
transforms.Normalize(**self.mean_std),
|
|
124
|
+
*post_aug,
|
|
125
|
+
]
|
|
126
|
+
return transforms.Compose(train_transforms)
|
|
127
|
+
|
|
128
|
+
def build_datasets(self) -> tuple[any, any, any]:
|
|
129
|
+
train_transform = self.build_train_transform()
|
|
130
|
+
valid_transform = self.build_valid_transform()
|
|
131
|
+
|
|
132
|
+
train_dataset = ImageFolder(
|
|
133
|
+
os.path.join(self.data_dir, "train"), train_transform
|
|
134
|
+
)
|
|
135
|
+
test_dataset = ImageFolder(
|
|
136
|
+
os.path.join(self.data_dir, "val"), valid_transform
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
train_dataset, val_dataset = self.sample_val_dataset(
|
|
140
|
+
train_dataset, valid_transform
|
|
141
|
+
)
|
|
142
|
+
return train_dataset, val_dataset, test_dataset
|
|
@@ -0,0 +1,18 @@
|
|
|
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 ...apps.trainer.run_config import RunConfig
|
|
6
|
+
|
|
7
|
+
__all__ = ["ClsRunConfig"]
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class ClsRunConfig(RunConfig):
|
|
11
|
+
label_smooth: float
|
|
12
|
+
mixup_config: dict # allow none to turn off mixup
|
|
13
|
+
bce: bool
|
|
14
|
+
mesa: dict
|
|
15
|
+
|
|
16
|
+
@property
|
|
17
|
+
def none_allowed(self):
|
|
18
|
+
return ["mixup_config", "mesa"] + super().none_allowed
|