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,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
@@ -0,0 +1,5 @@
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 .imagenet import *
@@ -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,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 .cls_run_config import *
6
+ from .cls_trainer import *
@@ -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