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,380 @@
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
+
8
+ from ..nn import (
9
+ ConvLayer,
10
+ DSConv,
11
+ EfficientViTBlock,
12
+ FusedMBConv,
13
+ IdentityLayer,
14
+ MBConv,
15
+ OpSequential,
16
+ ResBlock,
17
+ ResidualBlock,
18
+ )
19
+ from ..utils import build_kwargs_from_config
20
+
21
+ __all__ = [
22
+ "EfficientViTBackbone",
23
+ "efficientvit_backbone_b0",
24
+ "efficientvit_backbone_b1",
25
+ "efficientvit_backbone_b2",
26
+ "efficientvit_backbone_b3",
27
+ "EfficientViTLargeBackbone",
28
+ "efficientvit_backbone_l0",
29
+ "efficientvit_backbone_l1",
30
+ "efficientvit_backbone_l2",
31
+ "efficientvit_backbone_l3",
32
+ ]
33
+
34
+
35
+ class EfficientViTBackbone(nn.Module):
36
+ def __init__(
37
+ self,
38
+ width_list: list[int],
39
+ depth_list: list[int],
40
+ in_channels=3,
41
+ dim=32,
42
+ expand_ratio=4,
43
+ norm="bn2d",
44
+ act_func="hswish",
45
+ ) -> None:
46
+ super().__init__()
47
+
48
+ self.width_list = []
49
+ # input stem
50
+ self.input_stem = [
51
+ ConvLayer(
52
+ in_channels=3,
53
+ out_channels=width_list[0],
54
+ stride=2,
55
+ norm=norm,
56
+ act_func=act_func,
57
+ )
58
+ ]
59
+ for _ in range(depth_list[0]):
60
+ block = self.build_local_block(
61
+ in_channels=width_list[0],
62
+ out_channels=width_list[0],
63
+ stride=1,
64
+ expand_ratio=1,
65
+ norm=norm,
66
+ act_func=act_func,
67
+ )
68
+ self.input_stem.append(ResidualBlock(block, IdentityLayer()))
69
+ in_channels = width_list[0]
70
+ self.input_stem = OpSequential(self.input_stem)
71
+ self.width_list.append(in_channels)
72
+
73
+ # stages
74
+ self.stages = []
75
+ for w, d in zip(width_list[1:3], depth_list[1:3]):
76
+ stage = []
77
+ for i in range(d):
78
+ stride = 2 if i == 0 else 1
79
+ block = self.build_local_block(
80
+ in_channels=in_channels,
81
+ out_channels=w,
82
+ stride=stride,
83
+ expand_ratio=expand_ratio,
84
+ norm=norm,
85
+ act_func=act_func,
86
+ )
87
+ block = ResidualBlock(
88
+ block, IdentityLayer() if stride == 1 else None
89
+ )
90
+ stage.append(block)
91
+ in_channels = w
92
+ self.stages.append(OpSequential(stage))
93
+ self.width_list.append(in_channels)
94
+
95
+ for w, d in zip(width_list[3:], depth_list[3:]):
96
+ stage = []
97
+ block = self.build_local_block(
98
+ in_channels=in_channels,
99
+ out_channels=w,
100
+ stride=2,
101
+ expand_ratio=expand_ratio,
102
+ norm=norm,
103
+ act_func=act_func,
104
+ fewer_norm=True,
105
+ )
106
+ stage.append(ResidualBlock(block, None))
107
+ in_channels = w
108
+
109
+ for _ in range(d):
110
+ stage.append(
111
+ EfficientViTBlock(
112
+ in_channels=in_channels,
113
+ dim=dim,
114
+ expand_ratio=expand_ratio,
115
+ norm=norm,
116
+ act_func=act_func,
117
+ )
118
+ )
119
+ self.stages.append(OpSequential(stage))
120
+ self.width_list.append(in_channels)
121
+ self.stages = nn.ModuleList(self.stages)
122
+
123
+ @staticmethod
124
+ def build_local_block(
125
+ in_channels: int,
126
+ out_channels: int,
127
+ stride: int,
128
+ expand_ratio: float,
129
+ norm: str,
130
+ act_func: str,
131
+ fewer_norm: bool = False,
132
+ ) -> nn.Module:
133
+ if expand_ratio == 1:
134
+ block = DSConv(
135
+ in_channels=in_channels,
136
+ out_channels=out_channels,
137
+ stride=stride,
138
+ use_bias=(True, False) if fewer_norm else False,
139
+ norm=(None, norm) if fewer_norm else norm,
140
+ act_func=(act_func, None),
141
+ )
142
+ else:
143
+ block = MBConv(
144
+ in_channels=in_channels,
145
+ out_channels=out_channels,
146
+ stride=stride,
147
+ expand_ratio=expand_ratio,
148
+ use_bias=(True, True, False) if fewer_norm else False,
149
+ norm=(None, None, norm) if fewer_norm else norm,
150
+ act_func=(act_func, act_func, None),
151
+ )
152
+ return block
153
+
154
+ def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
155
+ output_dict = {"input": x}
156
+ output_dict["stage0"] = x = self.input_stem(x)
157
+ for stage_id, stage in enumerate(self.stages, 1):
158
+ output_dict["stage%d" % stage_id] = x = stage(x)
159
+ output_dict["stage_final"] = x
160
+ return output_dict
161
+
162
+
163
+ def efficientvit_backbone_b0(**kwargs) -> EfficientViTBackbone:
164
+ backbone = EfficientViTBackbone(
165
+ width_list=[8, 16, 32, 64, 128],
166
+ depth_list=[1, 2, 2, 2, 2],
167
+ dim=16,
168
+ **build_kwargs_from_config(kwargs, EfficientViTBackbone),
169
+ )
170
+ return backbone
171
+
172
+
173
+ def efficientvit_backbone_b1(**kwargs) -> EfficientViTBackbone:
174
+ backbone = EfficientViTBackbone(
175
+ width_list=[16, 32, 64, 128, 256],
176
+ depth_list=[1, 2, 3, 3, 4],
177
+ dim=16,
178
+ **build_kwargs_from_config(kwargs, EfficientViTBackbone),
179
+ )
180
+ return backbone
181
+
182
+
183
+ def efficientvit_backbone_b2(**kwargs) -> EfficientViTBackbone:
184
+ backbone = EfficientViTBackbone(
185
+ width_list=[24, 48, 96, 192, 384],
186
+ depth_list=[1, 3, 4, 4, 6],
187
+ dim=32,
188
+ **build_kwargs_from_config(kwargs, EfficientViTBackbone),
189
+ )
190
+ return backbone
191
+
192
+
193
+ def efficientvit_backbone_b3(**kwargs) -> EfficientViTBackbone:
194
+ backbone = EfficientViTBackbone(
195
+ width_list=[32, 64, 128, 256, 512],
196
+ depth_list=[1, 4, 6, 6, 9],
197
+ dim=32,
198
+ **build_kwargs_from_config(kwargs, EfficientViTBackbone),
199
+ )
200
+ return backbone
201
+
202
+
203
+ class EfficientViTLargeBackbone(nn.Module):
204
+ def __init__(
205
+ self,
206
+ width_list: list[int],
207
+ depth_list: list[int],
208
+ in_channels=3,
209
+ qkv_dim=32,
210
+ norm="bn2d",
211
+ act_func="gelu",
212
+ ) -> None:
213
+ super().__init__()
214
+
215
+ self.width_list = []
216
+ self.stages = []
217
+ # stage 0
218
+ stage0 = [
219
+ ConvLayer(
220
+ in_channels=3,
221
+ out_channels=width_list[0],
222
+ stride=2,
223
+ norm=norm,
224
+ act_func=act_func,
225
+ )
226
+ ]
227
+ for _ in range(depth_list[0]):
228
+ block = self.build_local_block(
229
+ stage_id=0,
230
+ in_channels=width_list[0],
231
+ out_channels=width_list[0],
232
+ stride=1,
233
+ expand_ratio=1,
234
+ norm=norm,
235
+ act_func=act_func,
236
+ )
237
+ stage0.append(ResidualBlock(block, IdentityLayer()))
238
+ in_channels = width_list[0]
239
+ self.stages.append(OpSequential(stage0))
240
+ self.width_list.append(in_channels)
241
+
242
+ for stage_id, (w, d) in enumerate(
243
+ zip(width_list[1:4], depth_list[1:4]), start=1
244
+ ):
245
+ stage = []
246
+ for i in range(d + 1):
247
+ stride = 2 if i == 0 else 1
248
+ block = self.build_local_block(
249
+ stage_id=stage_id,
250
+ in_channels=in_channels,
251
+ out_channels=w,
252
+ stride=stride,
253
+ expand_ratio=4 if stride == 1 else 16,
254
+ norm=norm,
255
+ act_func=act_func,
256
+ fewer_norm=stage_id > 2,
257
+ )
258
+ block = ResidualBlock(
259
+ block, IdentityLayer() if stride == 1 else None
260
+ )
261
+ stage.append(block)
262
+ in_channels = w
263
+ self.stages.append(OpSequential(stage))
264
+ self.width_list.append(in_channels)
265
+
266
+ for stage_id, (w, d) in enumerate(
267
+ zip(width_list[4:], depth_list[4:]), start=4
268
+ ):
269
+ stage = []
270
+ block = self.build_local_block(
271
+ stage_id=stage_id,
272
+ in_channels=in_channels,
273
+ out_channels=w,
274
+ stride=2,
275
+ expand_ratio=24,
276
+ norm=norm,
277
+ act_func=act_func,
278
+ fewer_norm=True,
279
+ )
280
+ stage.append(ResidualBlock(block, None))
281
+ in_channels = w
282
+
283
+ for _ in range(d):
284
+ stage.append(
285
+ EfficientViTBlock(
286
+ in_channels=in_channels,
287
+ dim=qkv_dim,
288
+ expand_ratio=6,
289
+ norm=norm,
290
+ act_func=act_func,
291
+ )
292
+ )
293
+ self.stages.append(OpSequential(stage))
294
+ self.width_list.append(in_channels)
295
+ self.stages = nn.ModuleList(self.stages)
296
+
297
+ @staticmethod
298
+ def build_local_block(
299
+ stage_id: int,
300
+ in_channels: int,
301
+ out_channels: int,
302
+ stride: int,
303
+ expand_ratio: float,
304
+ norm: str,
305
+ act_func: str,
306
+ fewer_norm: bool = False,
307
+ ) -> nn.Module:
308
+ if expand_ratio == 1:
309
+ block = ResBlock(
310
+ in_channels=in_channels,
311
+ out_channels=out_channels,
312
+ stride=stride,
313
+ use_bias=(True, False) if fewer_norm else False,
314
+ norm=(None, norm) if fewer_norm else norm,
315
+ act_func=(act_func, None),
316
+ )
317
+ elif stage_id <= 2:
318
+ block = FusedMBConv(
319
+ in_channels=in_channels,
320
+ out_channels=out_channels,
321
+ stride=stride,
322
+ expand_ratio=expand_ratio,
323
+ use_bias=(True, False) if fewer_norm else False,
324
+ norm=(None, norm) if fewer_norm else norm,
325
+ act_func=(act_func, None),
326
+ )
327
+ else:
328
+ block = MBConv(
329
+ in_channels=in_channels,
330
+ out_channels=out_channels,
331
+ stride=stride,
332
+ expand_ratio=expand_ratio,
333
+ use_bias=(True, True, False) if fewer_norm else False,
334
+ norm=(None, None, norm) if fewer_norm else norm,
335
+ act_func=(act_func, act_func, None),
336
+ )
337
+ return block
338
+
339
+ def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
340
+ output_dict = {"input": x}
341
+ for stage_id, stage in enumerate(self.stages):
342
+ output_dict["stage%d" % stage_id] = x = stage(x)
343
+ output_dict["stage_final"] = x
344
+ return output_dict
345
+
346
+
347
+ def efficientvit_backbone_l0(**kwargs) -> EfficientViTLargeBackbone:
348
+ backbone = EfficientViTLargeBackbone(
349
+ width_list=[32, 64, 128, 256, 512],
350
+ depth_list=[1, 1, 1, 4, 4],
351
+ **build_kwargs_from_config(kwargs, EfficientViTLargeBackbone),
352
+ )
353
+ return backbone
354
+
355
+
356
+ def efficientvit_backbone_l1(**kwargs) -> EfficientViTLargeBackbone:
357
+ backbone = EfficientViTLargeBackbone(
358
+ width_list=[32, 64, 128, 256, 512],
359
+ depth_list=[1, 1, 1, 6, 6],
360
+ **build_kwargs_from_config(kwargs, EfficientViTLargeBackbone),
361
+ )
362
+ return backbone
363
+
364
+
365
+ def efficientvit_backbone_l2(**kwargs) -> EfficientViTLargeBackbone:
366
+ backbone = EfficientViTLargeBackbone(
367
+ width_list=[32, 64, 128, 256, 512],
368
+ depth_list=[1, 2, 2, 8, 8],
369
+ **build_kwargs_from_config(kwargs, EfficientViTLargeBackbone),
370
+ )
371
+ return backbone
372
+
373
+
374
+ def efficientvit_backbone_l3(**kwargs) -> EfficientViTLargeBackbone:
375
+ backbone = EfficientViTLargeBackbone(
376
+ width_list=[64, 128, 256, 512, 1024],
377
+ depth_list=[1, 2, 2, 8, 8],
378
+ **build_kwargs_from_config(kwargs, EfficientViTLargeBackbone),
379
+ )
380
+ return backbone
@@ -0,0 +1,188 @@
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
+
8
+ from .backbone import EfficientViTBackbone, EfficientViTLargeBackbone
9
+ from ..nn import ConvLayer, LinearLayer, OpSequential
10
+ from ..utils import build_kwargs_from_config
11
+
12
+ __all__ = [
13
+ "EfficientViTCls",
14
+ ######################
15
+ "efficientvit_cls_b0",
16
+ "efficientvit_cls_b1",
17
+ "efficientvit_cls_b2",
18
+ "efficientvit_cls_b3",
19
+ ######################
20
+ "efficientvit_cls_l1",
21
+ "efficientvit_cls_l2",
22
+ "efficientvit_cls_l3",
23
+ ]
24
+
25
+
26
+ class ClsHead(OpSequential):
27
+ def __init__(
28
+ self,
29
+ in_channels: int,
30
+ width_list: list[int],
31
+ n_classes=1000,
32
+ dropout=0.0,
33
+ norm="bn2d",
34
+ act_func="hswish",
35
+ fid="stage_final",
36
+ ):
37
+ ops = [
38
+ ConvLayer(
39
+ in_channels, width_list[0], 1, norm=norm, act_func=act_func
40
+ ),
41
+ nn.AdaptiveAvgPool2d(output_size=1),
42
+ LinearLayer(
43
+ width_list[0],
44
+ width_list[1],
45
+ False,
46
+ norm="ln",
47
+ act_func=act_func,
48
+ ),
49
+ LinearLayer(width_list[1], n_classes, True, dropout, None, None),
50
+ ]
51
+ super().__init__(ops)
52
+
53
+ self.fid = fid
54
+
55
+ def forward(self, feed_dict: dict[str, torch.Tensor]) -> torch.Tensor:
56
+ x = feed_dict[self.fid]
57
+ return OpSequential.forward(self, x)
58
+
59
+
60
+ class EfficientViTCls(nn.Module):
61
+ def __init__(
62
+ self,
63
+ backbone: EfficientViTBackbone or EfficientViTLargeBackbone,
64
+ head: ClsHead,
65
+ ) -> None:
66
+ super().__init__()
67
+ self.backbone = backbone
68
+ self.head = head
69
+
70
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
71
+ feed_dict = self.backbone(x)
72
+ output = self.head(feed_dict)
73
+ return output
74
+
75
+
76
+ def efficientvit_cls_b0(**kwargs) -> EfficientViTCls:
77
+ from efficientvit.models.efficientvit.backbone import (
78
+ efficientvit_backbone_b0,
79
+ )
80
+
81
+ backbone = efficientvit_backbone_b0(**kwargs)
82
+
83
+ head = ClsHead(
84
+ in_channels=128,
85
+ width_list=[1024, 1280],
86
+ **build_kwargs_from_config(kwargs, ClsHead),
87
+ )
88
+ model = EfficientViTCls(backbone, head)
89
+ return model
90
+
91
+
92
+ def efficientvit_cls_b1(**kwargs) -> EfficientViTCls:
93
+ from efficientvit.models.efficientvit.backbone import (
94
+ efficientvit_backbone_b1,
95
+ )
96
+
97
+ backbone = efficientvit_backbone_b1(**kwargs)
98
+
99
+ head = ClsHead(
100
+ in_channels=256,
101
+ width_list=[1536, 1600],
102
+ **build_kwargs_from_config(kwargs, ClsHead),
103
+ )
104
+ model = EfficientViTCls(backbone, head)
105
+ return model
106
+
107
+
108
+ def efficientvit_cls_b2(**kwargs) -> EfficientViTCls:
109
+ from efficientvit.models.efficientvit.backbone import (
110
+ efficientvit_backbone_b2,
111
+ )
112
+
113
+ backbone = efficientvit_backbone_b2(**kwargs)
114
+
115
+ head = ClsHead(
116
+ in_channels=384,
117
+ width_list=[2304, 2560],
118
+ **build_kwargs_from_config(kwargs, ClsHead),
119
+ )
120
+ model = EfficientViTCls(backbone, head)
121
+ return model
122
+
123
+
124
+ def efficientvit_cls_b3(**kwargs) -> EfficientViTCls:
125
+ from efficientvit.models.efficientvit.backbone import (
126
+ efficientvit_backbone_b3,
127
+ )
128
+
129
+ backbone = efficientvit_backbone_b3(**kwargs)
130
+
131
+ head = ClsHead(
132
+ in_channels=512,
133
+ width_list=[2304, 2560],
134
+ **build_kwargs_from_config(kwargs, ClsHead),
135
+ )
136
+ model = EfficientViTCls(backbone, head)
137
+ return model
138
+
139
+
140
+ def efficientvit_cls_l1(**kwargs) -> EfficientViTCls:
141
+ from efficientvit.models.efficientvit.backbone import (
142
+ efficientvit_backbone_l1,
143
+ )
144
+
145
+ backbone = efficientvit_backbone_l1(**kwargs)
146
+
147
+ head = ClsHead(
148
+ in_channels=512,
149
+ width_list=[3072, 3200],
150
+ act_func="gelu",
151
+ **build_kwargs_from_config(kwargs, ClsHead),
152
+ )
153
+ model = EfficientViTCls(backbone, head)
154
+ return model
155
+
156
+
157
+ def efficientvit_cls_l2(**kwargs) -> EfficientViTCls:
158
+ from efficientvit.models.efficientvit.backbone import (
159
+ efficientvit_backbone_l2,
160
+ )
161
+
162
+ backbone = efficientvit_backbone_l2(**kwargs)
163
+
164
+ head = ClsHead(
165
+ in_channels=512,
166
+ width_list=[3072, 3200],
167
+ act_func="gelu",
168
+ **build_kwargs_from_config(kwargs, ClsHead),
169
+ )
170
+ model = EfficientViTCls(backbone, head)
171
+ return model
172
+
173
+
174
+ def efficientvit_cls_l3(**kwargs) -> EfficientViTCls:
175
+ from efficientvit.models.efficientvit.backbone import (
176
+ efficientvit_backbone_l3,
177
+ )
178
+
179
+ backbone = efficientvit_backbone_l3(**kwargs)
180
+
181
+ head = ClsHead(
182
+ in_channels=1024,
183
+ width_list=[6144, 6400],
184
+ act_func="gelu",
185
+ **build_kwargs_from_config(kwargs, ClsHead),
186
+ )
187
+ model = EfficientViTCls(backbone, head)
188
+ return model