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,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)