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