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,265 @@
|
|
|
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
|
+
import sys
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
import torch
|
|
10
|
+
import torch.nn as nn
|
|
11
|
+
import torch.nn.functional as F
|
|
12
|
+
import torchpack.distributed as dist
|
|
13
|
+
from tqdm import tqdm
|
|
14
|
+
|
|
15
|
+
from ...apps.trainer import Trainer
|
|
16
|
+
from ...apps.utils import AverageMeter, sync_tensor
|
|
17
|
+
from .utils import accuracy, apply_mixup, label_smooth
|
|
18
|
+
from ...models.utils import list_join, list_mean, torch_random_choices
|
|
19
|
+
|
|
20
|
+
__all__ = ["ClsTrainer"]
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class ClsTrainer(Trainer):
|
|
24
|
+
def __init__(
|
|
25
|
+
self,
|
|
26
|
+
path: str,
|
|
27
|
+
model: nn.Module,
|
|
28
|
+
data_provider,
|
|
29
|
+
auto_restart_thresh: float or None = None,
|
|
30
|
+
) -> None:
|
|
31
|
+
super().__init__(
|
|
32
|
+
path=path,
|
|
33
|
+
model=model,
|
|
34
|
+
data_provider=data_provider,
|
|
35
|
+
)
|
|
36
|
+
self.auto_restart_thresh = auto_restart_thresh
|
|
37
|
+
self.test_criterion = nn.CrossEntropyLoss()
|
|
38
|
+
|
|
39
|
+
def _validate(self, model, data_loader, epoch) -> dict[str, any]:
|
|
40
|
+
val_loss = AverageMeter()
|
|
41
|
+
val_top1 = AverageMeter()
|
|
42
|
+
val_top5 = AverageMeter()
|
|
43
|
+
|
|
44
|
+
with torch.no_grad():
|
|
45
|
+
with tqdm(
|
|
46
|
+
total=len(data_loader),
|
|
47
|
+
desc=f"Validate Epoch #{epoch + 1}",
|
|
48
|
+
disable=not dist.is_master(),
|
|
49
|
+
file=sys.stdout,
|
|
50
|
+
) as t:
|
|
51
|
+
for images, labels in data_loader:
|
|
52
|
+
images, labels = images.cuda(), labels.cuda()
|
|
53
|
+
# compute output
|
|
54
|
+
output = model(images)
|
|
55
|
+
loss = self.test_criterion(output, labels)
|
|
56
|
+
val_loss.update(loss, images.shape[0])
|
|
57
|
+
if self.data_provider.n_classes >= 100:
|
|
58
|
+
acc1, acc5 = accuracy(output, labels, topk=(1, 5))
|
|
59
|
+
val_top5.update(acc5[0], images.shape[0])
|
|
60
|
+
else:
|
|
61
|
+
acc1 = accuracy(output, labels, topk=(1,))[0]
|
|
62
|
+
val_top1.update(acc1[0], images.shape[0])
|
|
63
|
+
|
|
64
|
+
t.set_postfix(
|
|
65
|
+
{
|
|
66
|
+
"loss": val_loss.avg,
|
|
67
|
+
"top1": val_top1.avg,
|
|
68
|
+
"top5": val_top5.avg,
|
|
69
|
+
"#samples": val_top1.get_count(),
|
|
70
|
+
"bs": images.shape[0],
|
|
71
|
+
"res": images.shape[2],
|
|
72
|
+
}
|
|
73
|
+
)
|
|
74
|
+
t.update()
|
|
75
|
+
return {
|
|
76
|
+
"val_top1": val_top1.avg,
|
|
77
|
+
"val_loss": val_loss.avg,
|
|
78
|
+
**({"val_top5": val_top5.avg} if val_top5.count > 0 else {}),
|
|
79
|
+
}
|
|
80
|
+
|
|
81
|
+
def before_step(self, feed_dict: dict[str, any]) -> dict[str, any]:
|
|
82
|
+
images = feed_dict["data"].cuda()
|
|
83
|
+
labels = feed_dict["label"].cuda()
|
|
84
|
+
|
|
85
|
+
# label smooth
|
|
86
|
+
labels = label_smooth(
|
|
87
|
+
labels, self.data_provider.n_classes, self.run_config.label_smooth
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
# mixup
|
|
91
|
+
if self.run_config.mixup_config is not None:
|
|
92
|
+
# choose active mixup config
|
|
93
|
+
mix_weight_list = [
|
|
94
|
+
mix_list[2] for mix_list in self.run_config.mixup_config["op"]
|
|
95
|
+
]
|
|
96
|
+
active_id = torch_random_choices(
|
|
97
|
+
list(range(len(self.run_config.mixup_config["op"]))),
|
|
98
|
+
weight_list=mix_weight_list,
|
|
99
|
+
)
|
|
100
|
+
active_id = int(sync_tensor(active_id, reduce="root"))
|
|
101
|
+
active_mixup_config = self.run_config.mixup_config["op"][active_id]
|
|
102
|
+
mixup_type, mixup_alpha = active_mixup_config[:2]
|
|
103
|
+
|
|
104
|
+
lam = float(
|
|
105
|
+
torch.distributions.beta.Beta(
|
|
106
|
+
mixup_alpha, mixup_alpha
|
|
107
|
+
).sample()
|
|
108
|
+
)
|
|
109
|
+
lam = float(np.clip(lam, 0, 1))
|
|
110
|
+
lam = float(sync_tensor(lam, reduce="root"))
|
|
111
|
+
|
|
112
|
+
images, labels = apply_mixup(images, labels, lam, mixup_type)
|
|
113
|
+
|
|
114
|
+
return {
|
|
115
|
+
"data": images,
|
|
116
|
+
"label": labels,
|
|
117
|
+
}
|
|
118
|
+
|
|
119
|
+
def run_step(self, feed_dict: dict[str, any]) -> dict[str, any]:
|
|
120
|
+
images = feed_dict["data"]
|
|
121
|
+
labels = feed_dict["label"]
|
|
122
|
+
|
|
123
|
+
# setup mesa
|
|
124
|
+
if (
|
|
125
|
+
self.run_config.mesa is not None
|
|
126
|
+
and self.run_config.mesa["thresh"] <= self.run_config.progress
|
|
127
|
+
):
|
|
128
|
+
ema_model = self.ema.shadows
|
|
129
|
+
with torch.inference_mode():
|
|
130
|
+
ema_output = ema_model(images).detach()
|
|
131
|
+
ema_output = torch.clone(ema_output)
|
|
132
|
+
ema_output = F.sigmoid(ema_output).detach()
|
|
133
|
+
else:
|
|
134
|
+
ema_output = None
|
|
135
|
+
|
|
136
|
+
with torch.autocast(
|
|
137
|
+
device_type="cuda", dtype=torch.float16, enabled=self.fp16
|
|
138
|
+
):
|
|
139
|
+
output = self.model(images)
|
|
140
|
+
loss = self.train_criterion(output, labels)
|
|
141
|
+
# mesa loss
|
|
142
|
+
if ema_output is not None:
|
|
143
|
+
mesa_loss = self.train_criterion(output, ema_output)
|
|
144
|
+
loss = loss + self.run_config.mesa["ratio"] * mesa_loss
|
|
145
|
+
self.scaler.scale(loss).backward()
|
|
146
|
+
|
|
147
|
+
# calc train top1 acc
|
|
148
|
+
if self.run_config.mixup_config is None:
|
|
149
|
+
top1 = accuracy(output, torch.argmax(labels, dim=1), topk=(1,))[0][
|
|
150
|
+
0
|
|
151
|
+
]
|
|
152
|
+
else:
|
|
153
|
+
top1 = None
|
|
154
|
+
|
|
155
|
+
return {
|
|
156
|
+
"loss": loss,
|
|
157
|
+
"top1": top1,
|
|
158
|
+
}
|
|
159
|
+
|
|
160
|
+
def _train_one_epoch(self, epoch: int) -> dict[str, any]:
|
|
161
|
+
train_loss = AverageMeter()
|
|
162
|
+
train_top1 = AverageMeter()
|
|
163
|
+
|
|
164
|
+
with tqdm(
|
|
165
|
+
total=len(self.data_provider.train),
|
|
166
|
+
desc="Train Epoch #{}".format(epoch + 1),
|
|
167
|
+
disable=not dist.is_master(),
|
|
168
|
+
file=sys.stdout,
|
|
169
|
+
) as t:
|
|
170
|
+
for images, labels in self.data_provider.train:
|
|
171
|
+
feed_dict = {"data": images, "label": labels}
|
|
172
|
+
|
|
173
|
+
# preprocessing
|
|
174
|
+
feed_dict = self.before_step(feed_dict)
|
|
175
|
+
# clear gradient
|
|
176
|
+
self.optimizer.zero_grad()
|
|
177
|
+
# forward & backward
|
|
178
|
+
output_dict = self.run_step(feed_dict)
|
|
179
|
+
# update: optimizer, lr_scheduler
|
|
180
|
+
self.after_step()
|
|
181
|
+
|
|
182
|
+
# update train metrics
|
|
183
|
+
train_loss.update(output_dict["loss"], images.shape[0])
|
|
184
|
+
if output_dict["top1"] is not None:
|
|
185
|
+
train_top1.update(output_dict["top1"], images.shape[0])
|
|
186
|
+
|
|
187
|
+
# tqdm
|
|
188
|
+
postfix_dict = {
|
|
189
|
+
"loss": train_loss.avg,
|
|
190
|
+
"top1": train_top1.avg,
|
|
191
|
+
"bs": images.shape[0],
|
|
192
|
+
"res": images.shape[2],
|
|
193
|
+
"lr": list_join(
|
|
194
|
+
sorted(
|
|
195
|
+
set(
|
|
196
|
+
[
|
|
197
|
+
group["lr"]
|
|
198
|
+
for group in self.optimizer.param_groups
|
|
199
|
+
]
|
|
200
|
+
)
|
|
201
|
+
),
|
|
202
|
+
"#",
|
|
203
|
+
"%.1E",
|
|
204
|
+
),
|
|
205
|
+
"progress": self.run_config.progress,
|
|
206
|
+
}
|
|
207
|
+
t.set_postfix(postfix_dict)
|
|
208
|
+
t.update()
|
|
209
|
+
return {
|
|
210
|
+
**({"train_top1": train_top1.avg} if train_top1.count > 0 else {}),
|
|
211
|
+
"train_loss": train_loss.avg,
|
|
212
|
+
}
|
|
213
|
+
|
|
214
|
+
def train(self, trials=0, save_freq=1) -> None:
|
|
215
|
+
if self.run_config.bce:
|
|
216
|
+
self.train_criterion = nn.BCEWithLogitsLoss()
|
|
217
|
+
else:
|
|
218
|
+
self.train_criterion = nn.CrossEntropyLoss()
|
|
219
|
+
|
|
220
|
+
for epoch in range(
|
|
221
|
+
self.start_epoch,
|
|
222
|
+
self.run_config.n_epochs + self.run_config.warmup_epochs,
|
|
223
|
+
):
|
|
224
|
+
train_info_dict = self.train_one_epoch(epoch)
|
|
225
|
+
# eval
|
|
226
|
+
val_info_dict = self.multires_validate(epoch=epoch)
|
|
227
|
+
avg_top1 = list_mean(
|
|
228
|
+
[info_dict["val_top1"] for info_dict in val_info_dict.values()]
|
|
229
|
+
)
|
|
230
|
+
is_best = avg_top1 > self.best_val
|
|
231
|
+
self.best_val = max(avg_top1, self.best_val)
|
|
232
|
+
|
|
233
|
+
if self.auto_restart_thresh is not None:
|
|
234
|
+
if self.best_val - avg_top1 > self.auto_restart_thresh:
|
|
235
|
+
self.write_log(
|
|
236
|
+
f"Abnormal accuracy drop: {self.best_val} -> {avg_top1}"
|
|
237
|
+
)
|
|
238
|
+
self.load_model(
|
|
239
|
+
os.path.join(self.checkpoint_path, "model_best.pt")
|
|
240
|
+
)
|
|
241
|
+
return self.train(trials + 1, save_freq)
|
|
242
|
+
|
|
243
|
+
# log
|
|
244
|
+
val_log = self.run_config.epoch_format(epoch)
|
|
245
|
+
val_log += f"\tval_top1={avg_top1:.2f}({self.best_val:.2f})"
|
|
246
|
+
val_log += "\tVal("
|
|
247
|
+
for key in list(val_info_dict.values())[0]:
|
|
248
|
+
if key == "val_top1":
|
|
249
|
+
continue
|
|
250
|
+
val_log += f"{key}={list_mean([info_dict[key] for info_dict in val_info_dict.values()]):.2f},"
|
|
251
|
+
val_log += ")\tTrain("
|
|
252
|
+
for key, val in train_info_dict.items():
|
|
253
|
+
val_log += f"{key}={val:.2E},"
|
|
254
|
+
val_log += f'lr={list_join(sorted(set([group["lr"] for group in self.optimizer.param_groups])), "#", "%.1E")})'
|
|
255
|
+
self.write_log(val_log, prefix="valid", print_log=False)
|
|
256
|
+
|
|
257
|
+
# save model
|
|
258
|
+
if (epoch + 1) % save_freq == 0 or (
|
|
259
|
+
is_best and self.run_config.progress > 0.8
|
|
260
|
+
):
|
|
261
|
+
self.save_model(
|
|
262
|
+
only_state_dict=False,
|
|
263
|
+
epoch=epoch,
|
|
264
|
+
model_name="model_best.pt" if is_best else "checkpoint.pt",
|
|
265
|
+
)
|
|
@@ -0,0 +1,7 @@
|
|
|
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 .label_smooth import *
|
|
6
|
+
from .metric import *
|
|
7
|
+
from .mixup 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
|
+
import torch
|
|
6
|
+
|
|
7
|
+
__all__ = ["label_smooth"]
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def label_smooth(target: torch.Tensor, n_classes: int, smooth_factor=0.1) -> torch.Tensor:
|
|
11
|
+
# convert to one-hot
|
|
12
|
+
batch_size = target.shape[0]
|
|
13
|
+
target = torch.unsqueeze(target, 1)
|
|
14
|
+
soft_target = torch.zeros((batch_size, n_classes), device=target.device)
|
|
15
|
+
soft_target.scatter_(1, target, 1)
|
|
16
|
+
# label smoothing
|
|
17
|
+
soft_target = torch.add(soft_target * (1 - smooth_factor), smooth_factor / n_classes)
|
|
18
|
+
return soft_target
|
|
@@ -0,0 +1,23 @@
|
|
|
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__ = ["accuracy"]
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def accuracy(output: torch.Tensor, target: torch.Tensor, topk=(1,)) -> list[torch.Tensor]:
|
|
11
|
+
"""Computes the precision@k for the specified values of k."""
|
|
12
|
+
maxk = max(topk)
|
|
13
|
+
batch_size = target.shape[0]
|
|
14
|
+
|
|
15
|
+
_, pred = output.topk(maxk, 1, True, True)
|
|
16
|
+
pred = pred.t()
|
|
17
|
+
correct = pred.eq(target.reshape(1, -1).expand_as(pred))
|
|
18
|
+
|
|
19
|
+
res = []
|
|
20
|
+
for k in topk:
|
|
21
|
+
correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)
|
|
22
|
+
res.append(correct_k.mul_(100.0 / batch_size))
|
|
23
|
+
return res
|
|
@@ -0,0 +1,67 @@
|
|
|
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.distributions
|
|
6
|
+
|
|
7
|
+
from ....apps.data_provider.augment import rand_bbox
|
|
8
|
+
from ....models.utils.random import torch_randint, torch_shuffle
|
|
9
|
+
|
|
10
|
+
__all__ = ["apply_mixup", "mixup", "cutmix"]
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def apply_mixup(
|
|
14
|
+
images: torch.Tensor,
|
|
15
|
+
labels: torch.Tensor,
|
|
16
|
+
lam: float,
|
|
17
|
+
mix_type="mixup",
|
|
18
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
19
|
+
if mix_type == "mixup":
|
|
20
|
+
return mixup(images, labels, lam)
|
|
21
|
+
elif mix_type == "cutmix":
|
|
22
|
+
return cutmix(images, labels, lam)
|
|
23
|
+
else:
|
|
24
|
+
raise NotImplementedError
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def mixup(
|
|
28
|
+
images: torch.Tensor,
|
|
29
|
+
target: torch.Tensor,
|
|
30
|
+
lam: float,
|
|
31
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
32
|
+
rand_index = torch_shuffle(list(range(0, images.shape[0])))
|
|
33
|
+
|
|
34
|
+
flipped_images = images[rand_index]
|
|
35
|
+
flipped_target = target[rand_index]
|
|
36
|
+
|
|
37
|
+
return (
|
|
38
|
+
lam * images + (1 - lam) * flipped_images,
|
|
39
|
+
lam * target + (1 - lam) * flipped_target,
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def cutmix(
|
|
44
|
+
images: torch.Tensor,
|
|
45
|
+
target: torch.Tensor,
|
|
46
|
+
lam: float,
|
|
47
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
48
|
+
rand_index = torch_shuffle(list(range(0, images.shape[0])))
|
|
49
|
+
|
|
50
|
+
flipped_images = images[rand_index]
|
|
51
|
+
flipped_target = target[rand_index]
|
|
52
|
+
|
|
53
|
+
b, _, h, w = images.shape
|
|
54
|
+
lam_list = []
|
|
55
|
+
for i in range(b):
|
|
56
|
+
bbx1, bby1, bbx2, bby2 = rand_bbox(
|
|
57
|
+
h=h,
|
|
58
|
+
w=w,
|
|
59
|
+
lam=lam,
|
|
60
|
+
rand_func=torch_randint,
|
|
61
|
+
)
|
|
62
|
+
images[i, :, bby1:bby2, bbx1:bbx2] = flipped_images[
|
|
63
|
+
i, :, bby1:bby2, bbx1:bbx2
|
|
64
|
+
]
|
|
65
|
+
lam_list.append(1 - ((bbx2 - bbx1) * (bby2 - bby1) / (h * w)))
|
|
66
|
+
lam = torch.Tensor(lam_list).to(images.device).view(b, 1)
|
|
67
|
+
return images, lam * target + (1 - lam) * flipped_target
|
|
File without changes
|
|
@@ -0,0 +1,8 @@
|
|
|
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 .backbone import *
|
|
6
|
+
from .cls import *
|
|
7
|
+
from .sam import *
|
|
8
|
+
from .seg import *
|