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,318 @@
|
|
|
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 torch
|
|
8
|
+
import torch.nn as nn
|
|
9
|
+
import torchpack.distributed as dist
|
|
10
|
+
|
|
11
|
+
from ..data_provider import DataProvider, parse_image_size
|
|
12
|
+
from .run_config import RunConfig
|
|
13
|
+
from ..utils import EMA
|
|
14
|
+
from ...models.nn.norm import reset_bn
|
|
15
|
+
from ...models.utils import is_parallel, load_state_dict_from_file
|
|
16
|
+
|
|
17
|
+
__all__ = ["Trainer"]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Trainer:
|
|
21
|
+
def __init__(
|
|
22
|
+
self, path: str, model: nn.Module, data_provider: DataProvider
|
|
23
|
+
):
|
|
24
|
+
self.path = os.path.realpath(os.path.expanduser(path))
|
|
25
|
+
self.model = model.cuda()
|
|
26
|
+
self.data_provider = data_provider
|
|
27
|
+
|
|
28
|
+
self.ema = None
|
|
29
|
+
|
|
30
|
+
self.checkpoint_path = os.path.join(self.path, "checkpoint")
|
|
31
|
+
self.logs_path = os.path.join(self.path, "logs")
|
|
32
|
+
for path in [self.path, self.checkpoint_path, self.logs_path]:
|
|
33
|
+
os.makedirs(path, exist_ok=True)
|
|
34
|
+
|
|
35
|
+
self.best_val = 0.0
|
|
36
|
+
self.start_epoch = 0
|
|
37
|
+
|
|
38
|
+
@property
|
|
39
|
+
def network(self) -> nn.Module:
|
|
40
|
+
return self.model.module if is_parallel(self.model) else self.model
|
|
41
|
+
|
|
42
|
+
@property
|
|
43
|
+
def eval_network(self) -> nn.Module:
|
|
44
|
+
if self.ema is None:
|
|
45
|
+
model = self.model
|
|
46
|
+
else:
|
|
47
|
+
model = self.ema.shadows
|
|
48
|
+
model = model.module if is_parallel(model) else model
|
|
49
|
+
return model
|
|
50
|
+
|
|
51
|
+
def write_log(
|
|
52
|
+
self, log_str, prefix="valid", print_log=True, mode="a"
|
|
53
|
+
) -> None:
|
|
54
|
+
if dist.is_master():
|
|
55
|
+
fout = open(os.path.join(self.logs_path, f"{prefix}.log"), mode)
|
|
56
|
+
fout.write(log_str + "\n")
|
|
57
|
+
fout.flush()
|
|
58
|
+
fout.close()
|
|
59
|
+
if print_log:
|
|
60
|
+
print(log_str)
|
|
61
|
+
|
|
62
|
+
def save_model(
|
|
63
|
+
self,
|
|
64
|
+
checkpoint=None,
|
|
65
|
+
only_state_dict=True,
|
|
66
|
+
epoch=0,
|
|
67
|
+
model_name=None,
|
|
68
|
+
) -> None:
|
|
69
|
+
if dist.is_master():
|
|
70
|
+
if checkpoint is None:
|
|
71
|
+
if only_state_dict:
|
|
72
|
+
checkpoint = {"state_dict": self.network.state_dict()}
|
|
73
|
+
else:
|
|
74
|
+
checkpoint = {
|
|
75
|
+
"state_dict": self.network.state_dict(),
|
|
76
|
+
"epoch": epoch,
|
|
77
|
+
"best_val": self.best_val,
|
|
78
|
+
"optimizer": self.optimizer.state_dict(),
|
|
79
|
+
"lr_scheduler": self.lr_scheduler.state_dict(),
|
|
80
|
+
"ema": (
|
|
81
|
+
self.ema.state_dict()
|
|
82
|
+
if self.ema is not None
|
|
83
|
+
else None
|
|
84
|
+
),
|
|
85
|
+
"scaler": (
|
|
86
|
+
self.scaler.state_dict() if self.fp16 else None
|
|
87
|
+
),
|
|
88
|
+
}
|
|
89
|
+
|
|
90
|
+
model_name = model_name or "checkpoint.pt"
|
|
91
|
+
|
|
92
|
+
latest_fname = os.path.join(self.checkpoint_path, "latest.txt")
|
|
93
|
+
model_path = os.path.join(self.checkpoint_path, model_name)
|
|
94
|
+
with open(latest_fname, "w") as _fout:
|
|
95
|
+
_fout.write(model_path + "\n")
|
|
96
|
+
torch.save(checkpoint, model_path)
|
|
97
|
+
|
|
98
|
+
def load_model(self, model_fname=None) -> None:
|
|
99
|
+
latest_fname = os.path.join(self.checkpoint_path, "latest.txt")
|
|
100
|
+
if model_fname is None and os.path.exists(latest_fname):
|
|
101
|
+
with open(latest_fname, "r") as fin:
|
|
102
|
+
model_fname = fin.readline()
|
|
103
|
+
if len(model_fname) > 0 and model_fname[-1] == "\n":
|
|
104
|
+
model_fname = model_fname[:-1]
|
|
105
|
+
try:
|
|
106
|
+
if model_fname is None:
|
|
107
|
+
model_fname = f"{self.checkpoint_path}/checkpoint.pt"
|
|
108
|
+
elif not os.path.exists(model_fname):
|
|
109
|
+
model_fname = (
|
|
110
|
+
f"{self.checkpoint_path}/{os.path.basename(model_fname)}"
|
|
111
|
+
)
|
|
112
|
+
if not os.path.exists(model_fname):
|
|
113
|
+
model_fname = f"{self.checkpoint_path}/checkpoint.pt"
|
|
114
|
+
print(f"=> loading checkpoint {model_fname}")
|
|
115
|
+
checkpoint = load_state_dict_from_file(model_fname, False)
|
|
116
|
+
except Exception:
|
|
117
|
+
self.write_log(
|
|
118
|
+
f"fail to load checkpoint from {self.checkpoint_path}"
|
|
119
|
+
)
|
|
120
|
+
return
|
|
121
|
+
|
|
122
|
+
# load checkpoint
|
|
123
|
+
self.network.load_state_dict(checkpoint["state_dict"], strict=False)
|
|
124
|
+
log = []
|
|
125
|
+
if "epoch" in checkpoint:
|
|
126
|
+
self.start_epoch = checkpoint["epoch"] + 1
|
|
127
|
+
self.run_config.update_global_step(self.start_epoch)
|
|
128
|
+
log.append(f"epoch={self.start_epoch - 1}")
|
|
129
|
+
if "best_val" in checkpoint:
|
|
130
|
+
self.best_val = checkpoint["best_val"]
|
|
131
|
+
log.append(f"best_val={self.best_val:.2f}")
|
|
132
|
+
if "optimizer" in checkpoint:
|
|
133
|
+
self.optimizer.load_state_dict(checkpoint["optimizer"])
|
|
134
|
+
log.append("optimizer")
|
|
135
|
+
if "lr_scheduler" in checkpoint:
|
|
136
|
+
self.lr_scheduler.load_state_dict(checkpoint["lr_scheduler"])
|
|
137
|
+
log.append("lr_scheduler")
|
|
138
|
+
if "ema" in checkpoint and self.ema is not None:
|
|
139
|
+
self.ema.load_state_dict(checkpoint["ema"])
|
|
140
|
+
log.append("ema")
|
|
141
|
+
if "scaler" in checkpoint and self.fp16:
|
|
142
|
+
self.scaler.load_state_dict(checkpoint["scaler"])
|
|
143
|
+
log.append("scaler")
|
|
144
|
+
self.write_log("Loaded: " + ", ".join(log))
|
|
145
|
+
|
|
146
|
+
""" validate """
|
|
147
|
+
|
|
148
|
+
def reset_bn(
|
|
149
|
+
self,
|
|
150
|
+
network: nn.Module or None = None,
|
|
151
|
+
subset_size: int = 16000,
|
|
152
|
+
subset_batch_size: int = 100,
|
|
153
|
+
data_loader=None,
|
|
154
|
+
progress_bar=False,
|
|
155
|
+
) -> None:
|
|
156
|
+
network = network or self.network
|
|
157
|
+
if data_loader is None:
|
|
158
|
+
data_loader = []
|
|
159
|
+
for data in self.data_provider.build_sub_train_loader(
|
|
160
|
+
subset_size, subset_batch_size
|
|
161
|
+
):
|
|
162
|
+
if isinstance(data, list):
|
|
163
|
+
data_loader.append(data[0])
|
|
164
|
+
elif isinstance(data, dict):
|
|
165
|
+
data_loader.append(data["data"])
|
|
166
|
+
elif isinstance(data, torch.Tensor):
|
|
167
|
+
data_loader.append(data)
|
|
168
|
+
else:
|
|
169
|
+
raise NotImplementedError
|
|
170
|
+
|
|
171
|
+
network.eval()
|
|
172
|
+
reset_bn(
|
|
173
|
+
network,
|
|
174
|
+
data_loader,
|
|
175
|
+
sync=True,
|
|
176
|
+
progress_bar=progress_bar,
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
def _validate(self, model, data_loader, epoch) -> dict[str, any]:
|
|
180
|
+
raise NotImplementedError
|
|
181
|
+
|
|
182
|
+
def validate(
|
|
183
|
+
self, model=None, data_loader=None, is_test=True, epoch=0
|
|
184
|
+
) -> dict[str, any]:
|
|
185
|
+
model = model or self.eval_network
|
|
186
|
+
if data_loader is None:
|
|
187
|
+
if is_test:
|
|
188
|
+
data_loader = self.data_provider.test
|
|
189
|
+
else:
|
|
190
|
+
data_loader = self.data_provider.valid
|
|
191
|
+
|
|
192
|
+
model.eval()
|
|
193
|
+
return self._validate(model, data_loader, epoch)
|
|
194
|
+
|
|
195
|
+
def multires_validate(
|
|
196
|
+
self,
|
|
197
|
+
model=None,
|
|
198
|
+
data_loader=None,
|
|
199
|
+
is_test=True,
|
|
200
|
+
epoch=0,
|
|
201
|
+
eval_image_size=None,
|
|
202
|
+
) -> dict[str, dict[str, any]]:
|
|
203
|
+
eval_image_size = eval_image_size or self.run_config.eval_image_size
|
|
204
|
+
eval_image_size = eval_image_size or self.data_provider.image_size
|
|
205
|
+
model = model or self.eval_network
|
|
206
|
+
|
|
207
|
+
if not isinstance(eval_image_size, list):
|
|
208
|
+
eval_image_size = [eval_image_size]
|
|
209
|
+
|
|
210
|
+
output_dict = {}
|
|
211
|
+
for r in eval_image_size:
|
|
212
|
+
self.data_provider.assign_active_image_size(parse_image_size(r))
|
|
213
|
+
if self.run_config.reset_bn:
|
|
214
|
+
self.reset_bn(
|
|
215
|
+
network=model,
|
|
216
|
+
subset_size=self.run_config.reset_bn_size,
|
|
217
|
+
subset_batch_size=self.run_config.reset_bn_batch_size,
|
|
218
|
+
progress_bar=True,
|
|
219
|
+
)
|
|
220
|
+
output_dict[f"r{r}"] = self.validate(
|
|
221
|
+
model, data_loader, is_test, epoch
|
|
222
|
+
)
|
|
223
|
+
return output_dict
|
|
224
|
+
|
|
225
|
+
""" training """
|
|
226
|
+
|
|
227
|
+
def prep_for_training(
|
|
228
|
+
self,
|
|
229
|
+
run_config: RunConfig,
|
|
230
|
+
ema_decay: float or None = None,
|
|
231
|
+
fp16=False,
|
|
232
|
+
) -> None:
|
|
233
|
+
self.run_config = run_config
|
|
234
|
+
self.model = nn.parallel.DistributedDataParallel(
|
|
235
|
+
self.model.cuda(),
|
|
236
|
+
device_ids=[dist.local_rank()],
|
|
237
|
+
static_graph=True,
|
|
238
|
+
)
|
|
239
|
+
|
|
240
|
+
self.run_config.global_step = 0
|
|
241
|
+
self.run_config.batch_per_epoch = len(self.data_provider.train)
|
|
242
|
+
assert self.run_config.batch_per_epoch > 0, "Training set is empty"
|
|
243
|
+
|
|
244
|
+
# build optimizer
|
|
245
|
+
self.optimizer, self.lr_scheduler = self.run_config.build_optimizer(
|
|
246
|
+
self.model
|
|
247
|
+
)
|
|
248
|
+
|
|
249
|
+
if ema_decay is not None:
|
|
250
|
+
self.ema = EMA(self.network, ema_decay)
|
|
251
|
+
|
|
252
|
+
# fp16
|
|
253
|
+
self.fp16 = fp16
|
|
254
|
+
self.scaler = torch.cuda.amp.GradScaler(enabled=self.fp16)
|
|
255
|
+
|
|
256
|
+
def sync_model(self):
|
|
257
|
+
print("Sync model")
|
|
258
|
+
self.save_model(model_name="sync.pt")
|
|
259
|
+
dist.barrier()
|
|
260
|
+
checkpoint = torch.load(
|
|
261
|
+
os.path.join(self.checkpoint_path, "sync.pt"), map_location="cpu"
|
|
262
|
+
)
|
|
263
|
+
dist.barrier()
|
|
264
|
+
if dist.is_master():
|
|
265
|
+
os.remove(os.path.join(self.checkpoint_path, "sync.pt"))
|
|
266
|
+
dist.barrier()
|
|
267
|
+
|
|
268
|
+
# load checkpoint
|
|
269
|
+
self.network.load_state_dict(checkpoint["state_dict"], strict=False)
|
|
270
|
+
if "optimizer" in checkpoint:
|
|
271
|
+
self.optimizer.load_state_dict(checkpoint["optimizer"])
|
|
272
|
+
if "lr_scheduler" in checkpoint:
|
|
273
|
+
self.lr_scheduler.load_state_dict(checkpoint["lr_scheduler"])
|
|
274
|
+
if "ema" in checkpoint and self.ema is not None:
|
|
275
|
+
self.ema.load_state_dict(checkpoint["ema"])
|
|
276
|
+
if "scaler" in checkpoint and self.fp16:
|
|
277
|
+
self.scaler.load_state_dict(checkpoint["scaler"])
|
|
278
|
+
|
|
279
|
+
def before_step(self, feed_dict: dict[str, any]) -> dict[str, any]:
|
|
280
|
+
for key in feed_dict:
|
|
281
|
+
if isinstance(feed_dict[key], torch.Tensor):
|
|
282
|
+
feed_dict[key] = feed_dict[key].cuda()
|
|
283
|
+
return feed_dict
|
|
284
|
+
|
|
285
|
+
def run_step(self, feed_dict: dict[str, any]) -> dict[str, any]:
|
|
286
|
+
raise NotImplementedError
|
|
287
|
+
|
|
288
|
+
def after_step(self) -> None:
|
|
289
|
+
self.scaler.unscale_(self.optimizer)
|
|
290
|
+
# gradient clip
|
|
291
|
+
if self.run_config.grad_clip is not None:
|
|
292
|
+
torch.nn.utils.clip_grad_value_(
|
|
293
|
+
self.model.parameters(), self.run_config.grad_clip
|
|
294
|
+
)
|
|
295
|
+
# update
|
|
296
|
+
self.scaler.step(self.optimizer)
|
|
297
|
+
self.scaler.update()
|
|
298
|
+
|
|
299
|
+
self.lr_scheduler.step()
|
|
300
|
+
self.run_config.step()
|
|
301
|
+
# update ema
|
|
302
|
+
if self.ema is not None:
|
|
303
|
+
self.ema.step(self.network, self.run_config.global_step)
|
|
304
|
+
|
|
305
|
+
def _train_one_epoch(self, epoch: int) -> dict[str, any]:
|
|
306
|
+
raise NotImplementedError
|
|
307
|
+
|
|
308
|
+
def train_one_epoch(self, epoch: int) -> dict[str, any]:
|
|
309
|
+
self.model.train()
|
|
310
|
+
|
|
311
|
+
self.data_provider.set_epoch(epoch)
|
|
312
|
+
|
|
313
|
+
train_info_dict = self._train_one_epoch(epoch)
|
|
314
|
+
|
|
315
|
+
return train_info_dict
|
|
316
|
+
|
|
317
|
+
def train(self) -> None:
|
|
318
|
+
raise NotImplementedError
|
|
@@ -0,0 +1,129 @@
|
|
|
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 json
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import torch.nn as nn
|
|
9
|
+
|
|
10
|
+
from ..utils import CosineLRwithWarmup, build_optimizer
|
|
11
|
+
|
|
12
|
+
__all__ = ["Scheduler", "RunConfig"]
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class Scheduler:
|
|
16
|
+
PROGRESS = 0
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class RunConfig:
|
|
20
|
+
n_epochs: int
|
|
21
|
+
init_lr: float
|
|
22
|
+
warmup_epochs: int
|
|
23
|
+
warmup_lr: float
|
|
24
|
+
lr_schedule_name: str
|
|
25
|
+
lr_schedule_param: dict
|
|
26
|
+
optimizer_name: str
|
|
27
|
+
optimizer_params: dict
|
|
28
|
+
weight_decay: float
|
|
29
|
+
no_wd_keys: list
|
|
30
|
+
grad_clip: float # allow none to turn off grad clipping
|
|
31
|
+
reset_bn: bool
|
|
32
|
+
reset_bn_size: int
|
|
33
|
+
reset_bn_batch_size: int
|
|
34
|
+
eval_image_size: list # allow none to use image_size in data_provider
|
|
35
|
+
|
|
36
|
+
@property
|
|
37
|
+
def none_allowed(self):
|
|
38
|
+
return ["grad_clip", "eval_image_size"]
|
|
39
|
+
|
|
40
|
+
def __init__(self, **kwargs): # arguments must be passed as kwargs
|
|
41
|
+
for k, val in kwargs.items():
|
|
42
|
+
setattr(self, k, val)
|
|
43
|
+
|
|
44
|
+
# check that all relevant configs are there
|
|
45
|
+
annotations = {}
|
|
46
|
+
for clas in type(self).mro():
|
|
47
|
+
if hasattr(clas, "__annotations__"):
|
|
48
|
+
annotations.update(clas.__annotations__)
|
|
49
|
+
for k, k_type in annotations.items():
|
|
50
|
+
assert hasattr(
|
|
51
|
+
self, k
|
|
52
|
+
), f"Key {k} with type {k_type} required for initialization."
|
|
53
|
+
attr = getattr(self, k)
|
|
54
|
+
if k in self.none_allowed:
|
|
55
|
+
k_type = (k_type, type(None))
|
|
56
|
+
assert isinstance(
|
|
57
|
+
attr, k_type
|
|
58
|
+
), f"Key {k} must be type {k_type}, provided={attr}."
|
|
59
|
+
|
|
60
|
+
self.global_step = 0
|
|
61
|
+
self.batch_per_epoch = 1
|
|
62
|
+
|
|
63
|
+
def build_optimizer(self, network: nn.Module) -> tuple[any, any]:
|
|
64
|
+
r"""require setting 'batch_per_epoch' before building optimizer & lr_scheduler"""
|
|
65
|
+
param_dict = {}
|
|
66
|
+
for name, param in network.named_parameters():
|
|
67
|
+
if param.requires_grad:
|
|
68
|
+
opt_config = [self.weight_decay, self.init_lr]
|
|
69
|
+
if self.no_wd_keys is not None and len(self.no_wd_keys) > 0:
|
|
70
|
+
if np.any([key in name for key in self.no_wd_keys]):
|
|
71
|
+
opt_config[0] = 0
|
|
72
|
+
opt_key = json.dumps(opt_config)
|
|
73
|
+
param_dict[opt_key] = param_dict.get(opt_key, []) + [param]
|
|
74
|
+
|
|
75
|
+
net_params = []
|
|
76
|
+
for opt_key, param_list in param_dict.items():
|
|
77
|
+
wd, lr = json.loads(opt_key)
|
|
78
|
+
net_params.append(
|
|
79
|
+
{"params": param_list, "weight_decay": wd, "lr": lr}
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
optimizer = build_optimizer(
|
|
83
|
+
net_params,
|
|
84
|
+
self.optimizer_name,
|
|
85
|
+
self.optimizer_params,
|
|
86
|
+
self.init_lr,
|
|
87
|
+
)
|
|
88
|
+
# build lr scheduler
|
|
89
|
+
if self.lr_schedule_name == "cosine":
|
|
90
|
+
decay_steps = []
|
|
91
|
+
for epoch in self.lr_schedule_param.get("step", []):
|
|
92
|
+
decay_steps.append(epoch * self.batch_per_epoch)
|
|
93
|
+
decay_steps.append(self.n_epochs * self.batch_per_epoch)
|
|
94
|
+
decay_steps.sort()
|
|
95
|
+
lr_scheduler = CosineLRwithWarmup(
|
|
96
|
+
optimizer,
|
|
97
|
+
self.warmup_epochs * self.batch_per_epoch,
|
|
98
|
+
self.warmup_lr,
|
|
99
|
+
decay_steps,
|
|
100
|
+
)
|
|
101
|
+
else:
|
|
102
|
+
raise NotImplementedError
|
|
103
|
+
return optimizer, lr_scheduler
|
|
104
|
+
|
|
105
|
+
def update_global_step(self, epoch, batch_id=0) -> None:
|
|
106
|
+
self.global_step = epoch * self.batch_per_epoch + batch_id
|
|
107
|
+
Scheduler.PROGRESS = self.progress
|
|
108
|
+
|
|
109
|
+
@property
|
|
110
|
+
def progress(self) -> float:
|
|
111
|
+
warmup_steps = self.warmup_epochs * self.batch_per_epoch
|
|
112
|
+
steps = max(0, self.global_step - warmup_steps)
|
|
113
|
+
return steps / (self.n_epochs * self.batch_per_epoch)
|
|
114
|
+
|
|
115
|
+
def step(self) -> None:
|
|
116
|
+
self.global_step += 1
|
|
117
|
+
Scheduler.PROGRESS = self.progress
|
|
118
|
+
|
|
119
|
+
def get_remaining_epoch(self, epoch, post=True) -> int:
|
|
120
|
+
return self.n_epochs + self.warmup_epochs - epoch - int(post)
|
|
121
|
+
|
|
122
|
+
def epoch_format(self, epoch: int) -> str:
|
|
123
|
+
epoch_format = f"%.{len(str(self.n_epochs))}d"
|
|
124
|
+
epoch_format = f"[{epoch_format}/{epoch_format}]"
|
|
125
|
+
epoch_format = epoch_format % (
|
|
126
|
+
epoch + 1 - self.warmup_epochs,
|
|
127
|
+
self.n_epochs,
|
|
128
|
+
)
|
|
129
|
+
return epoch_format
|
|
@@ -0,0 +1,12 @@
|
|
|
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 .dist import *
|
|
6
|
+
from .ema import *
|
|
7
|
+
from .export import *
|
|
8
|
+
from .init import *
|
|
9
|
+
from .lr import *
|
|
10
|
+
from .metric import *
|
|
11
|
+
from .misc import *
|
|
12
|
+
from .opt import *
|
|
@@ -0,0 +1,32 @@
|
|
|
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
|
+
import torch.distributed
|
|
7
|
+
from torchpack import distributed
|
|
8
|
+
|
|
9
|
+
from ...models.utils.list import list_mean, list_sum
|
|
10
|
+
|
|
11
|
+
__all__ = ["sync_tensor"]
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def sync_tensor(
|
|
15
|
+
tensor: torch.Tensor or float, reduce="mean"
|
|
16
|
+
) -> torch.Tensor or list[torch.Tensor]:
|
|
17
|
+
if not isinstance(tensor, torch.Tensor):
|
|
18
|
+
tensor = torch.Tensor(1).fill_(tensor).cuda()
|
|
19
|
+
tensor_list = [torch.empty_like(tensor) for _ in range(distributed.size())]
|
|
20
|
+
torch.distributed.all_gather(
|
|
21
|
+
tensor_list, tensor.contiguous(), async_op=False
|
|
22
|
+
)
|
|
23
|
+
if reduce == "mean":
|
|
24
|
+
return list_mean(tensor_list)
|
|
25
|
+
elif reduce == "sum":
|
|
26
|
+
return list_sum(tensor_list)
|
|
27
|
+
elif reduce == "cat":
|
|
28
|
+
return torch.cat(tensor_list, dim=0)
|
|
29
|
+
elif reduce == "root":
|
|
30
|
+
return tensor_list[0]
|
|
31
|
+
else:
|
|
32
|
+
return tensor_list
|
|
@@ -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 copy
|
|
6
|
+
import math
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
import torch.nn as nn
|
|
10
|
+
|
|
11
|
+
from ...models.utils import is_parallel
|
|
12
|
+
|
|
13
|
+
__all__ = ["EMA"]
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def update_ema(
|
|
17
|
+
ema: nn.Module, new_state_dict: dict[str, torch.Tensor], decay: float
|
|
18
|
+
) -> None:
|
|
19
|
+
for k, v in ema.state_dict().items():
|
|
20
|
+
if v.dtype.is_floating_point:
|
|
21
|
+
v -= (1.0 - decay) * (v - new_state_dict[k].detach())
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class EMA:
|
|
25
|
+
def __init__(self, model: nn.Module, decay: float, warmup_steps=2000):
|
|
26
|
+
self.shadows = copy.deepcopy(
|
|
27
|
+
model.module if is_parallel(model) else model
|
|
28
|
+
).eval()
|
|
29
|
+
self.decay = decay
|
|
30
|
+
self.warmup_steps = warmup_steps
|
|
31
|
+
|
|
32
|
+
for p in self.shadows.parameters():
|
|
33
|
+
p.requires_grad = False
|
|
34
|
+
|
|
35
|
+
def step(self, model: nn.Module, global_step: int) -> None:
|
|
36
|
+
with torch.no_grad():
|
|
37
|
+
msd = (model.module if is_parallel(model) else model).state_dict()
|
|
38
|
+
update_ema(
|
|
39
|
+
self.shadows,
|
|
40
|
+
msd,
|
|
41
|
+
self.decay * (1 - math.exp(-global_step / self.warmup_steps)),
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
def state_dict(self) -> dict[float, dict[str, torch.Tensor]]:
|
|
45
|
+
return {self.decay: self.shadows.state_dict()}
|
|
46
|
+
|
|
47
|
+
def load_state_dict(
|
|
48
|
+
self, state_dict: dict[float, dict[str, torch.Tensor]]
|
|
49
|
+
) -> None:
|
|
50
|
+
for decay in state_dict:
|
|
51
|
+
if decay == self.decay:
|
|
52
|
+
self.shadows.load_state_dict(state_dict[decay])
|
|
@@ -0,0 +1,45 @@
|
|
|
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 io
|
|
6
|
+
import os
|
|
7
|
+
|
|
8
|
+
import onnx
|
|
9
|
+
import torch
|
|
10
|
+
import torch.nn as nn
|
|
11
|
+
from onnxsim import simplify as simplify_func
|
|
12
|
+
|
|
13
|
+
__all__ = ["export_onnx"]
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def export_onnx(model: nn.Module, export_path: str, sample_inputs: any, simplify=True, opset=11) -> None:
|
|
17
|
+
"""Export a model to a platform-specific onnx format.
|
|
18
|
+
|
|
19
|
+
Args:
|
|
20
|
+
model: a torch.nn.Module object.
|
|
21
|
+
export_path: export location.
|
|
22
|
+
sample_inputs: Any.
|
|
23
|
+
simplify: a flag to turn on onnx-simplifier
|
|
24
|
+
opset: int
|
|
25
|
+
"""
|
|
26
|
+
model.eval()
|
|
27
|
+
|
|
28
|
+
buffer = io.BytesIO()
|
|
29
|
+
with torch.no_grad():
|
|
30
|
+
torch.onnx.export(model, sample_inputs, buffer, opset_version=opset)
|
|
31
|
+
buffer.seek(0, 0)
|
|
32
|
+
if simplify:
|
|
33
|
+
onnx_model = onnx.load_model(buffer)
|
|
34
|
+
onnx_model, success = simplify_func(onnx_model)
|
|
35
|
+
assert success
|
|
36
|
+
new_buffer = io.BytesIO()
|
|
37
|
+
onnx.save(onnx_model, new_buffer)
|
|
38
|
+
buffer = new_buffer
|
|
39
|
+
buffer.seek(0, 0)
|
|
40
|
+
|
|
41
|
+
if buffer.getbuffer().nbytes > 0:
|
|
42
|
+
save_dir = os.path.dirname(export_path)
|
|
43
|
+
os.makedirs(save_dir, exist_ok=True)
|
|
44
|
+
with open(export_path, "wb") as f:
|
|
45
|
+
f.write(buffer.read())
|
|
@@ -0,0 +1,66 @@
|
|
|
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
|
+
import torch.nn as nn
|
|
7
|
+
from torch.nn.modules.batchnorm import _BatchNorm
|
|
8
|
+
|
|
9
|
+
__all__ = ["init_modules", "zero_last_gamma"]
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def init_modules(model: nn.Module or list[nn.Module], init_type="trunc_normal") -> None:
|
|
13
|
+
_DEFAULT_INIT_PARAM = {"trunc_normal": 0.02}
|
|
14
|
+
|
|
15
|
+
if isinstance(model, list):
|
|
16
|
+
for sub_module in model:
|
|
17
|
+
init_modules(sub_module, init_type)
|
|
18
|
+
else:
|
|
19
|
+
init_params = init_type.split("@")
|
|
20
|
+
init_params = float(init_params[1]) if len(init_params) > 1 else None
|
|
21
|
+
|
|
22
|
+
if init_type.startswith("trunc_normal"):
|
|
23
|
+
init_func = lambda param: nn.init.trunc_normal_(
|
|
24
|
+
param, std=(init_params or _DEFAULT_INIT_PARAM["trunc_normal"])
|
|
25
|
+
)
|
|
26
|
+
else:
|
|
27
|
+
raise NotImplementedError
|
|
28
|
+
|
|
29
|
+
for m in model.modules():
|
|
30
|
+
if isinstance(m, (nn.Conv2d, nn.Linear, nn.ConvTranspose2d)):
|
|
31
|
+
init_func(m.weight)
|
|
32
|
+
if m.bias is not None:
|
|
33
|
+
m.bias.data.zero_()
|
|
34
|
+
elif isinstance(m, nn.Embedding):
|
|
35
|
+
init_func(m.weight)
|
|
36
|
+
elif isinstance(m, (_BatchNorm, nn.GroupNorm, nn.LayerNorm)):
|
|
37
|
+
m.weight.data.fill_(1)
|
|
38
|
+
m.bias.data.zero_()
|
|
39
|
+
else:
|
|
40
|
+
weight = getattr(m, "weight", None)
|
|
41
|
+
bias = getattr(m, "bias", None)
|
|
42
|
+
if isinstance(weight, torch.nn.Parameter):
|
|
43
|
+
init_func(weight)
|
|
44
|
+
if isinstance(bias, torch.nn.Parameter):
|
|
45
|
+
bias.data.zero_()
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def zero_last_gamma(model: nn.Module, init_val=0) -> None:
|
|
49
|
+
import efficientvit.models.nn.ops as ops
|
|
50
|
+
|
|
51
|
+
for m in model.modules():
|
|
52
|
+
if isinstance(m, ops.ResidualBlock) and isinstance(m.shortcut, ops.IdentityLayer):
|
|
53
|
+
if isinstance(m.main, (ops.DSConv, ops.MBConv, ops.FusedMBConv)):
|
|
54
|
+
parent_module = m.main.point_conv
|
|
55
|
+
elif isinstance(m.main, ops.ResBlock):
|
|
56
|
+
parent_module = m.main.conv2
|
|
57
|
+
elif isinstance(m.main, ops.ConvLayer):
|
|
58
|
+
parent_module = m.main
|
|
59
|
+
elif isinstance(m.main, (ops.LiteMLA)):
|
|
60
|
+
parent_module = m.main.proj
|
|
61
|
+
else:
|
|
62
|
+
parent_module = None
|
|
63
|
+
if parent_module is not None:
|
|
64
|
+
norm = getattr(parent_module, "norm", None)
|
|
65
|
+
if norm is not None:
|
|
66
|
+
nn.init.constant_(norm.weight, init_val)
|