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,867 @@
1
+ # --------------------------------------------------------
2
+ # TinyViT Model Architecture
3
+ # Copyright (c) 2022 Microsoft
4
+ # Adapted from LeViT and Swin Transformer
5
+ # LeViT: (https://github.com/facebookresearch/levit)
6
+ # Swin: (https://github.com/microsoft/swin-transformer)
7
+ # Build the TinyViT Model
8
+ # --------------------------------------------------------
9
+
10
+ import itertools
11
+ import torch
12
+ import torch.nn as nn
13
+ import torch.nn.functional as F
14
+ import torch.utils.checkpoint as checkpoint
15
+ from timm.models.layers import (
16
+ DropPath as TimmDropPath,
17
+ to_2tuple,
18
+ trunc_normal_,
19
+ )
20
+ from timm.models.registry import register_model, list_models
21
+ from typing import Tuple
22
+
23
+
24
+ class Conv2d_BN(torch.nn.Sequential):
25
+ def __init__(
26
+ self,
27
+ a,
28
+ b,
29
+ ks=1,
30
+ stride=1,
31
+ pad=0,
32
+ dilation=1,
33
+ groups=1,
34
+ bn_weight_init=1,
35
+ ):
36
+ super().__init__()
37
+ self.add_module(
38
+ "c",
39
+ torch.nn.Conv2d(
40
+ a, b, ks, stride, pad, dilation, groups, bias=False
41
+ ),
42
+ )
43
+ bn = torch.nn.BatchNorm2d(b)
44
+ torch.nn.init.constant_(bn.weight, bn_weight_init)
45
+ torch.nn.init.constant_(bn.bias, 0)
46
+ self.add_module("bn", bn)
47
+
48
+ @torch.no_grad()
49
+ def fuse(self):
50
+ c, bn = self._modules.values()
51
+ w = bn.weight / (bn.running_var + bn.eps) ** 0.5
52
+ w = c.weight * w[:, None, None, None]
53
+ b = (
54
+ bn.bias
55
+ - bn.running_mean * bn.weight / (bn.running_var + bn.eps) ** 0.5
56
+ )
57
+ m = torch.nn.Conv2d(
58
+ w.size(1) * self.c.groups,
59
+ w.size(0),
60
+ w.shape[2:],
61
+ stride=self.c.stride,
62
+ padding=self.c.padding,
63
+ dilation=self.c.dilation,
64
+ groups=self.c.groups,
65
+ )
66
+ m.weight.data.copy_(w)
67
+ m.bias.data.copy_(b)
68
+ return m
69
+
70
+
71
+ class DropPath(TimmDropPath):
72
+ def __init__(self, drop_prob=None):
73
+ super().__init__(drop_prob=drop_prob)
74
+ self.drop_prob = drop_prob
75
+
76
+ def __repr__(self):
77
+ msg = super().__repr__()
78
+ msg += f"(drop_prob={self.drop_prob})"
79
+ return msg
80
+
81
+
82
+ class PatchEmbed(nn.Module):
83
+ def __init__(self, in_chans, embed_dim, resolution, activation):
84
+ super().__init__()
85
+ img_size: Tuple[int, int] = to_2tuple(resolution)
86
+ self.patches_resolution = (img_size[0] // 4, img_size[1] // 4)
87
+ self.num_patches = (
88
+ self.patches_resolution[0] * self.patches_resolution[1]
89
+ )
90
+ self.in_chans = in_chans
91
+ self.embed_dim = embed_dim
92
+ n = embed_dim
93
+ self.seq = nn.Sequential(
94
+ Conv2d_BN(in_chans, n // 2, 3, 2, 1),
95
+ activation(),
96
+ Conv2d_BN(n // 2, n, 3, 2, 1),
97
+ )
98
+
99
+ def forward(self, x):
100
+ return self.seq(x)
101
+
102
+
103
+ class MBConv(nn.Module):
104
+ def __init__(
105
+ self, in_chans, out_chans, expand_ratio, activation, drop_path
106
+ ):
107
+ super().__init__()
108
+ self.in_chans = in_chans
109
+ self.hidden_chans = int(in_chans * expand_ratio)
110
+ self.out_chans = out_chans
111
+
112
+ self.conv1 = Conv2d_BN(in_chans, self.hidden_chans, ks=1)
113
+ self.act1 = activation()
114
+
115
+ self.conv2 = Conv2d_BN(
116
+ self.hidden_chans,
117
+ self.hidden_chans,
118
+ ks=3,
119
+ stride=1,
120
+ pad=1,
121
+ groups=self.hidden_chans,
122
+ )
123
+ self.act2 = activation()
124
+
125
+ self.conv3 = Conv2d_BN(
126
+ self.hidden_chans, out_chans, ks=1, bn_weight_init=0.0
127
+ )
128
+ self.act3 = activation()
129
+
130
+ self.drop_path = (
131
+ DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
132
+ )
133
+
134
+ def forward(self, x):
135
+ shortcut = x
136
+
137
+ x = self.conv1(x)
138
+ x = self.act1(x)
139
+
140
+ x = self.conv2(x)
141
+ x = self.act2(x)
142
+
143
+ x = self.conv3(x)
144
+
145
+ x = self.drop_path(x)
146
+
147
+ x += shortcut
148
+ x = self.act3(x)
149
+
150
+ return x
151
+
152
+
153
+ class PatchMerging(nn.Module):
154
+ def __init__(self, input_resolution, dim, out_dim, activation):
155
+ super().__init__()
156
+
157
+ self.input_resolution = input_resolution
158
+ self.dim = dim
159
+ self.out_dim = out_dim
160
+ self.act = activation()
161
+ self.conv1 = Conv2d_BN(dim, out_dim, 1, 1, 0)
162
+ stride_c = 2
163
+ if out_dim == 320 or out_dim == 448 or out_dim == 576:
164
+ stride_c = 1
165
+ self.conv2 = Conv2d_BN(
166
+ out_dim, out_dim, 3, stride_c, 1, groups=out_dim
167
+ )
168
+ self.conv3 = Conv2d_BN(out_dim, out_dim, 1, 1, 0)
169
+
170
+ def forward(self, x):
171
+ if x.ndim == 3:
172
+ H, W = self.input_resolution
173
+ B = len(x)
174
+ # (B, C, H, W)
175
+ x = x.view(B, H, W, -1).permute(0, 3, 1, 2)
176
+
177
+ x = self.conv1(x)
178
+ x = self.act(x)
179
+
180
+ x = self.conv2(x)
181
+ x = self.act(x)
182
+ x = self.conv3(x)
183
+ x = x.flatten(2).transpose(1, 2)
184
+ return x
185
+
186
+
187
+ class ConvLayer(nn.Module):
188
+ def __init__(
189
+ self,
190
+ dim,
191
+ input_resolution,
192
+ depth,
193
+ activation,
194
+ drop_path=0.0,
195
+ downsample=None,
196
+ use_checkpoint=False,
197
+ out_dim=None,
198
+ conv_expand_ratio=4.0,
199
+ ):
200
+
201
+ super().__init__()
202
+ self.dim = dim
203
+ self.input_resolution = input_resolution
204
+ self.depth = depth
205
+ self.use_checkpoint = use_checkpoint
206
+
207
+ # build blocks
208
+ self.blocks = nn.ModuleList(
209
+ [
210
+ MBConv(
211
+ dim,
212
+ dim,
213
+ conv_expand_ratio,
214
+ activation,
215
+ drop_path[i] if isinstance(drop_path, list) else drop_path,
216
+ )
217
+ for i in range(depth)
218
+ ]
219
+ )
220
+
221
+ # patch merging layer
222
+ if downsample is not None:
223
+ self.downsample = downsample(
224
+ input_resolution,
225
+ dim=dim,
226
+ out_dim=out_dim,
227
+ activation=activation,
228
+ )
229
+ else:
230
+ self.downsample = None
231
+
232
+ def forward(self, x):
233
+ for blk in self.blocks:
234
+ if self.use_checkpoint:
235
+ x = checkpoint.checkpoint(blk, x)
236
+ else:
237
+ x = blk(x)
238
+ if self.downsample is not None:
239
+ x = self.downsample(x)
240
+ return x
241
+
242
+
243
+ class Mlp(nn.Module):
244
+ def __init__(
245
+ self,
246
+ in_features,
247
+ hidden_features=None,
248
+ out_features=None,
249
+ act_layer=nn.GELU,
250
+ drop=0.0,
251
+ ):
252
+ super().__init__()
253
+ out_features = out_features or in_features
254
+ hidden_features = hidden_features or in_features
255
+ self.norm = nn.LayerNorm(in_features)
256
+ self.fc1 = nn.Linear(in_features, hidden_features)
257
+ self.fc2 = nn.Linear(hidden_features, out_features)
258
+ self.act = act_layer()
259
+ self.drop = nn.Dropout(drop)
260
+
261
+ def forward(self, x):
262
+ x = self.norm(x)
263
+
264
+ x = self.fc1(x)
265
+ x = self.act(x)
266
+ x = self.drop(x)
267
+ x = self.fc2(x)
268
+ x = self.drop(x)
269
+ return x
270
+
271
+
272
+ class Attention(torch.nn.Module):
273
+ def __init__(
274
+ self,
275
+ dim,
276
+ key_dim,
277
+ num_heads=8,
278
+ attn_ratio=4,
279
+ resolution=(14, 14),
280
+ ):
281
+ super().__init__()
282
+ # (h, w)
283
+ assert isinstance(resolution, tuple) and len(resolution) == 2
284
+ self.num_heads = num_heads
285
+ self.scale = key_dim**-0.5
286
+ self.key_dim = key_dim
287
+ self.nh_kd = nh_kd = key_dim * num_heads
288
+ self.d = int(attn_ratio * key_dim)
289
+ self.dh = int(attn_ratio * key_dim) * num_heads
290
+ self.attn_ratio = attn_ratio
291
+ h = self.dh + nh_kd * 2
292
+
293
+ self.norm = nn.LayerNorm(dim)
294
+ self.qkv = nn.Linear(dim, h)
295
+ self.proj = nn.Linear(self.dh, dim)
296
+
297
+ points = list(
298
+ itertools.product(range(resolution[0]), range(resolution[1]))
299
+ )
300
+ N = len(points)
301
+ attention_offsets = {}
302
+ idxs = []
303
+ for p1 in points:
304
+ for p2 in points:
305
+ offset = (abs(p1[0] - p2[0]), abs(p1[1] - p2[1]))
306
+ if offset not in attention_offsets:
307
+ attention_offsets[offset] = len(attention_offsets)
308
+ idxs.append(attention_offsets[offset])
309
+ self.attention_biases = torch.nn.Parameter(
310
+ torch.zeros(num_heads, len(attention_offsets))
311
+ )
312
+ self.register_buffer(
313
+ "attention_bias_idxs",
314
+ torch.LongTensor(idxs).view(N, N),
315
+ persistent=False,
316
+ )
317
+
318
+ @torch.no_grad()
319
+ def train(self, mode=True):
320
+ super().train(mode)
321
+ if mode and hasattr(self, "ab"):
322
+ del self.ab
323
+ else:
324
+ self.ab = self.attention_biases[:, self.attention_bias_idxs]
325
+ # self.register_buffer('ab',
326
+ # self.attention_biases[:, self.attention_bias_idxs],
327
+ # persistent=False)
328
+
329
+ def forward(self, x): # x (B,N,C)
330
+ B, N, _ = x.shape
331
+
332
+ # Normalization
333
+ x = self.norm(x)
334
+
335
+ qkv = self.qkv(x)
336
+ # (B, N, num_heads, d)
337
+ q, k, v = qkv.view(B, N, self.num_heads, -1).split(
338
+ [self.key_dim, self.key_dim, self.d], dim=3
339
+ )
340
+ # (B, num_heads, N, d)
341
+ q = q.permute(0, 2, 1, 3)
342
+ k = k.permute(0, 2, 1, 3)
343
+ v = v.permute(0, 2, 1, 3)
344
+
345
+ attn = (q @ k.transpose(-2, -1)) * self.scale + (
346
+ self.attention_biases[:, self.attention_bias_idxs]
347
+ if self.training
348
+ else self.ab
349
+ )
350
+ attn = attn.softmax(dim=-1)
351
+ x = (attn @ v).transpose(1, 2).reshape(B, N, self.dh)
352
+ x = self.proj(x)
353
+ return x
354
+
355
+
356
+ class TinyViTBlock(nn.Module):
357
+ r"""TinyViT Block.
358
+
359
+ Args:
360
+ dim (int): Number of input channels.
361
+ input_resolution (tuple[int, int]): Input resulotion.
362
+ num_heads (int): Number of attention heads.
363
+ window_size (int): Window size.
364
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
365
+ drop (float, optional): Dropout rate. Default: 0.0
366
+ drop_path (float, optional): Stochastic depth rate. Default: 0.0
367
+ local_conv_size (int): the kernel size of the convolution between
368
+ Attention and MLP. Default: 3
369
+ activation: the activation function. Default: nn.GELU
370
+ """
371
+
372
+ def __init__(
373
+ self,
374
+ dim,
375
+ input_resolution,
376
+ num_heads,
377
+ window_size=7,
378
+ mlp_ratio=4.0,
379
+ drop=0.0,
380
+ drop_path=0.0,
381
+ local_conv_size=3,
382
+ activation=nn.GELU,
383
+ ):
384
+ super().__init__()
385
+ self.dim = dim
386
+ self.input_resolution = input_resolution
387
+ self.num_heads = num_heads
388
+ assert window_size > 0, "window_size must be greater than 0"
389
+ self.window_size = window_size
390
+ self.mlp_ratio = mlp_ratio
391
+
392
+ self.drop_path = (
393
+ DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
394
+ )
395
+
396
+ assert dim % num_heads == 0, "dim must be divisible by num_heads"
397
+ head_dim = dim // num_heads
398
+
399
+ window_resolution = (window_size, window_size)
400
+ self.attn = Attention(
401
+ dim,
402
+ head_dim,
403
+ num_heads,
404
+ attn_ratio=1,
405
+ resolution=window_resolution,
406
+ )
407
+
408
+ mlp_hidden_dim = int(dim * mlp_ratio)
409
+ mlp_activation = activation
410
+ self.mlp = Mlp(
411
+ in_features=dim,
412
+ hidden_features=mlp_hidden_dim,
413
+ act_layer=mlp_activation,
414
+ drop=drop,
415
+ )
416
+
417
+ pad = local_conv_size // 2
418
+ self.local_conv = Conv2d_BN(
419
+ dim, dim, ks=local_conv_size, stride=1, pad=pad, groups=dim
420
+ )
421
+
422
+ def forward(self, x):
423
+ H, W = self.input_resolution
424
+ B, L, C = x.shape
425
+ assert L == H * W, "input feature has wrong size"
426
+ res_x = x
427
+ if H == self.window_size and W == self.window_size:
428
+ x = self.attn(x)
429
+ else:
430
+ x = x.view(B, H, W, C)
431
+ pad_b = (
432
+ self.window_size - H % self.window_size
433
+ ) % self.window_size
434
+ pad_r = (
435
+ self.window_size - W % self.window_size
436
+ ) % self.window_size
437
+ padding = pad_b > 0 or pad_r > 0
438
+
439
+ if padding:
440
+ x = F.pad(x, (0, 0, 0, pad_r, 0, pad_b))
441
+
442
+ pH, pW = H + pad_b, W + pad_r
443
+ nH = pH // self.window_size
444
+ nW = pW // self.window_size
445
+ # window partition
446
+ x = (
447
+ x.view(B, nH, self.window_size, nW, self.window_size, C)
448
+ .transpose(2, 3)
449
+ .reshape(B * nH * nW, self.window_size * self.window_size, C)
450
+ )
451
+ x = self.attn(x)
452
+ # window reverse
453
+ x = (
454
+ x.view(B, nH, nW, self.window_size, self.window_size, C)
455
+ .transpose(2, 3)
456
+ .reshape(B, pH, pW, C)
457
+ )
458
+
459
+ if padding:
460
+ x = x[:, :H, :W].contiguous()
461
+
462
+ x = x.view(B, L, C)
463
+
464
+ x = res_x + self.drop_path(x)
465
+
466
+ x = x.transpose(1, 2).reshape(B, C, H, W)
467
+ x = self.local_conv(x)
468
+ x = x.view(B, C, L).transpose(1, 2)
469
+
470
+ x = x + self.drop_path(self.mlp(x))
471
+ return x
472
+
473
+ def extra_repr(self) -> str:
474
+ return (
475
+ f"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, "
476
+ f"window_size={self.window_size}, mlp_ratio={self.mlp_ratio}"
477
+ )
478
+
479
+
480
+ class BasicLayer(nn.Module):
481
+ """A basic TinyViT layer for one stage.
482
+
483
+ Args:
484
+ dim (int): Number of input channels.
485
+ input_resolution (tuple[int]): Input resolution.
486
+ depth (int): Number of blocks.
487
+ num_heads (int): Number of attention heads.
488
+ window_size (int): Local window size.
489
+ mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
490
+ drop (float, optional): Dropout rate. Default: 0.0
491
+ drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
492
+ downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
493
+ use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
494
+ local_conv_size: the kernel size of the depthwise convolution between attention and MLP. Default: 3
495
+ activation: the activation function. Default: nn.GELU
496
+ out_dim: the output dimension of the layer. Default: dim
497
+ """
498
+
499
+ def __init__(
500
+ self,
501
+ dim,
502
+ input_resolution,
503
+ depth,
504
+ num_heads,
505
+ window_size,
506
+ mlp_ratio=4.0,
507
+ drop=0.0,
508
+ drop_path=0.0,
509
+ downsample=None,
510
+ use_checkpoint=False,
511
+ local_conv_size=3,
512
+ activation=nn.GELU,
513
+ out_dim=None,
514
+ ):
515
+
516
+ super().__init__()
517
+ self.dim = dim
518
+ self.input_resolution = input_resolution
519
+ self.depth = depth
520
+ self.use_checkpoint = use_checkpoint
521
+
522
+ # build blocks
523
+ self.blocks = nn.ModuleList(
524
+ [
525
+ TinyViTBlock(
526
+ dim=dim,
527
+ input_resolution=input_resolution,
528
+ num_heads=num_heads,
529
+ window_size=window_size,
530
+ mlp_ratio=mlp_ratio,
531
+ drop=drop,
532
+ drop_path=(
533
+ drop_path[i]
534
+ if isinstance(drop_path, list)
535
+ else drop_path
536
+ ),
537
+ local_conv_size=local_conv_size,
538
+ activation=activation,
539
+ )
540
+ for i in range(depth)
541
+ ]
542
+ )
543
+
544
+ # patch merging layer
545
+ if downsample is not None:
546
+ self.downsample = downsample(
547
+ input_resolution,
548
+ dim=dim,
549
+ out_dim=out_dim,
550
+ activation=activation,
551
+ )
552
+ else:
553
+ self.downsample = None
554
+
555
+ def forward(self, x):
556
+ for blk in self.blocks:
557
+ if self.use_checkpoint:
558
+ x = checkpoint.checkpoint(blk, x)
559
+ else:
560
+ x = blk(x)
561
+ if self.downsample is not None:
562
+ x = self.downsample(x)
563
+ return x
564
+
565
+ def extra_repr(self) -> str:
566
+ return f"dim={self.dim}, input_resolution={self.input_resolution}, depth={self.depth}"
567
+
568
+
569
+ class LayerNorm2d(nn.Module):
570
+ def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
571
+ super().__init__()
572
+ self.weight = nn.Parameter(torch.ones(num_channels))
573
+ self.bias = nn.Parameter(torch.zeros(num_channels))
574
+ self.eps = eps
575
+
576
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
577
+ u = x.mean(1, keepdim=True)
578
+ s = (x - u).pow(2).mean(1, keepdim=True)
579
+ x = (x - u) / torch.sqrt(s + self.eps)
580
+ x = self.weight[:, None, None] * x + self.bias[:, None, None]
581
+ return x
582
+
583
+
584
+ class TinyViT(nn.Module):
585
+ def __init__(
586
+ self,
587
+ img_size=224,
588
+ in_chans=3,
589
+ num_classes=1000,
590
+ embed_dims=[96, 192, 384, 768],
591
+ depths=[2, 2, 6, 2],
592
+ num_heads=[3, 6, 12, 24],
593
+ window_sizes=[7, 7, 14, 7],
594
+ mlp_ratio=4.0,
595
+ drop_rate=0.0,
596
+ drop_path_rate=0.1,
597
+ use_checkpoint=False,
598
+ mbconv_expand_ratio=4.0,
599
+ local_conv_size=3,
600
+ layer_lr_decay=1.0,
601
+ ):
602
+ super().__init__()
603
+ self.img_size = img_size
604
+ # import pdb;pdb.set_trace()
605
+ self.num_classes = num_classes
606
+ self.depths = depths
607
+ self.num_layers = len(depths)
608
+ self.mlp_ratio = mlp_ratio
609
+
610
+ activation = nn.GELU
611
+
612
+ self.patch_embed = PatchEmbed(
613
+ in_chans=in_chans,
614
+ embed_dim=embed_dims[0],
615
+ resolution=img_size,
616
+ activation=activation,
617
+ )
618
+
619
+ patches_resolution = self.patch_embed.patches_resolution
620
+ self.patches_resolution = patches_resolution
621
+
622
+ # stochastic depth
623
+ dpr = [
624
+ x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))
625
+ ] # stochastic depth decay rule
626
+
627
+ # build layers
628
+ self.layers = nn.ModuleList()
629
+ for i_layer in range(self.num_layers):
630
+ kwargs = dict(
631
+ dim=embed_dims[i_layer],
632
+ input_resolution=(
633
+ patches_resolution[0]
634
+ // (2 ** (i_layer - 1 if i_layer == 3 else i_layer)),
635
+ patches_resolution[1]
636
+ // (2 ** (i_layer - 1 if i_layer == 3 else i_layer)),
637
+ ),
638
+ # input_resolution=(patches_resolution[0] // (2 ** i_layer),
639
+ # patches_resolution[1] // (2 ** i_layer)),
640
+ depth=depths[i_layer],
641
+ drop_path=dpr[
642
+ sum(depths[:i_layer]) : sum(depths[: i_layer + 1])
643
+ ],
644
+ downsample=(
645
+ PatchMerging if (i_layer < self.num_layers - 1) else None
646
+ ),
647
+ use_checkpoint=use_checkpoint,
648
+ out_dim=embed_dims[min(i_layer + 1, len(embed_dims) - 1)],
649
+ activation=activation,
650
+ )
651
+ if i_layer == 0:
652
+ layer = ConvLayer(
653
+ conv_expand_ratio=mbconv_expand_ratio,
654
+ **kwargs,
655
+ )
656
+ else:
657
+ layer = BasicLayer(
658
+ num_heads=num_heads[i_layer],
659
+ window_size=window_sizes[i_layer],
660
+ mlp_ratio=self.mlp_ratio,
661
+ drop=drop_rate,
662
+ local_conv_size=local_conv_size,
663
+ **kwargs,
664
+ )
665
+ self.layers.append(layer)
666
+
667
+ # Classifier head
668
+ self.norm_head = nn.LayerNorm(embed_dims[-1])
669
+ self.head = (
670
+ nn.Linear(embed_dims[-1], num_classes)
671
+ if num_classes > 0
672
+ else torch.nn.Identity()
673
+ )
674
+
675
+ # init weights
676
+ self.apply(self._init_weights)
677
+ self.set_layer_lr_decay(layer_lr_decay)
678
+ self.neck = nn.Sequential(
679
+ nn.Conv2d(
680
+ embed_dims[-1],
681
+ 256,
682
+ kernel_size=1,
683
+ bias=False,
684
+ ),
685
+ LayerNorm2d(256),
686
+ nn.Conv2d(
687
+ 256,
688
+ 256,
689
+ kernel_size=3,
690
+ padding=1,
691
+ bias=False,
692
+ ),
693
+ LayerNorm2d(256),
694
+ )
695
+
696
+ def set_layer_lr_decay(self, layer_lr_decay):
697
+ decay_rate = layer_lr_decay
698
+
699
+ # layers -> blocks (depth)
700
+ depth = sum(self.depths)
701
+ lr_scales = [decay_rate ** (depth - i - 1) for i in range(depth)]
702
+ # print("LR SCALES:", lr_scales)
703
+
704
+ def _set_lr_scale(m, scale):
705
+ for p in m.parameters():
706
+ p.lr_scale = scale
707
+
708
+ self.patch_embed.apply(lambda x: _set_lr_scale(x, lr_scales[0]))
709
+ i = 0
710
+ for layer in self.layers:
711
+ for block in layer.blocks:
712
+ block.apply(lambda x: _set_lr_scale(x, lr_scales[i]))
713
+ i += 1
714
+ if layer.downsample is not None:
715
+ layer.downsample.apply(
716
+ lambda x: _set_lr_scale(x, lr_scales[i - 1])
717
+ )
718
+ assert i == depth
719
+ for m in [self.norm_head, self.head]:
720
+ m.apply(lambda x: _set_lr_scale(x, lr_scales[-1]))
721
+
722
+ for k, p in self.named_parameters():
723
+ p.param_name = k
724
+
725
+ def _check_lr_scale(m):
726
+ for p in m.parameters():
727
+ assert hasattr(p, "lr_scale"), p.param_name
728
+
729
+ self.apply(_check_lr_scale)
730
+
731
+ def _init_weights(self, m):
732
+ if isinstance(m, nn.Linear):
733
+ trunc_normal_(m.weight, std=0.02)
734
+ if isinstance(m, nn.Linear) and m.bias is not None:
735
+ nn.init.constant_(m.bias, 0)
736
+ elif isinstance(m, nn.LayerNorm):
737
+ nn.init.constant_(m.bias, 0)
738
+ nn.init.constant_(m.weight, 1.0)
739
+
740
+ @torch.jit.ignore
741
+ def no_weight_decay_keywords(self):
742
+ return {"attention_biases"}
743
+
744
+ def forward_features(self, x):
745
+ # x: (N, C, H, W)
746
+ x = self.patch_embed(x)
747
+
748
+ x = self.layers[0](x)
749
+ start_i = 1
750
+
751
+ for i in range(start_i, len(self.layers)):
752
+ layer = self.layers[i]
753
+ x = layer(x)
754
+ B, _, C = x.size()
755
+ x = x.view(B, 64, 64, C)
756
+ x = x.permute(0, 3, 1, 2)
757
+ x = self.neck(x)
758
+ return x
759
+
760
+ def forward(self, x):
761
+ x = self.forward_features(x)
762
+ # x = self.norm_head(x)
763
+ # x = self.head(x)
764
+ return x
765
+
766
+
767
+ _checkpoint_url_format = "https://github.com/wkcn/TinyViT-model-zoo/releases/download/checkpoints/{}.pth"
768
+ _provided_checkpoints = {
769
+ "tiny_vit_5m_224": "tiny_vit_5m_22kto1k_distill",
770
+ "tiny_vit_11m_224": "tiny_vit_11m_22kto1k_distill",
771
+ "tiny_vit_21m_224": "tiny_vit_21m_22kto1k_distill",
772
+ "tiny_vit_21m_384": "tiny_vit_21m_22kto1k_384_distill",
773
+ "tiny_vit_21m_512": "tiny_vit_21m_22kto1k_512_distill",
774
+ }
775
+
776
+
777
+ def register_tiny_vit_model(fn):
778
+ """Register a TinyViT model
779
+ It is a wrapper of `register_model` with loading the pretrained checkpoint.
780
+ """
781
+
782
+ def fn_wrapper(pretrained=False, **kwargs):
783
+ model = fn()
784
+ if pretrained:
785
+ model_name = fn.__name__
786
+ assert (
787
+ model_name in _provided_checkpoints
788
+ ), f"Sorry that the checkpoint `{model_name}` is not provided yet."
789
+ url = _checkpoint_url_format.format(
790
+ _provided_checkpoints[model_name]
791
+ )
792
+ checkpoint = torch.hub.load_state_dict_from_url(
793
+ url=url,
794
+ map_location="cpu",
795
+ check_hash=False,
796
+ )
797
+ model.load_state_dict(checkpoint["model"])
798
+
799
+ return model
800
+
801
+ # rename the name of fn_wrapper
802
+ fn_wrapper.__name__ = fn.__name__
803
+ if fn_wrapper.__name__ in list_models():
804
+ return
805
+ return register_model(fn_wrapper)
806
+
807
+
808
+ @register_tiny_vit_model
809
+ def tiny_vit_5m_224(pretrained=False, num_classes=1000, drop_path_rate=0.0):
810
+ return TinyViT(
811
+ num_classes=num_classes,
812
+ embed_dims=[64, 128, 160, 320],
813
+ depths=[2, 2, 6, 2],
814
+ num_heads=[2, 4, 5, 10],
815
+ window_sizes=[7, 7, 14, 7],
816
+ drop_path_rate=drop_path_rate,
817
+ )
818
+
819
+
820
+ @register_tiny_vit_model
821
+ def tiny_vit_11m_224(pretrained=False, num_classes=1000, drop_path_rate=0.1):
822
+ return TinyViT(
823
+ num_classes=num_classes,
824
+ embed_dims=[64, 128, 256, 448],
825
+ depths=[2, 2, 6, 2],
826
+ num_heads=[2, 4, 8, 14],
827
+ window_sizes=[7, 7, 14, 7],
828
+ drop_path_rate=drop_path_rate,
829
+ )
830
+
831
+
832
+ @register_tiny_vit_model
833
+ def tiny_vit_21m_224(pretrained=False, num_classes=1000, drop_path_rate=0.2):
834
+ return TinyViT(
835
+ num_classes=num_classes,
836
+ embed_dims=[96, 192, 384, 576],
837
+ depths=[2, 2, 6, 2],
838
+ num_heads=[3, 6, 12, 18],
839
+ window_sizes=[7, 7, 14, 7],
840
+ drop_path_rate=drop_path_rate,
841
+ )
842
+
843
+
844
+ @register_tiny_vit_model
845
+ def tiny_vit_21m_384(pretrained=False, num_classes=1000, drop_path_rate=0.1):
846
+ return TinyViT(
847
+ img_size=384,
848
+ num_classes=num_classes,
849
+ embed_dims=[96, 192, 384, 576],
850
+ depths=[2, 2, 6, 2],
851
+ num_heads=[3, 6, 12, 18],
852
+ window_sizes=[12, 12, 24, 12],
853
+ drop_path_rate=drop_path_rate,
854
+ )
855
+
856
+
857
+ @register_tiny_vit_model
858
+ def tiny_vit_21m_512(pretrained=False, num_classes=1000, drop_path_rate=0.1):
859
+ return TinyViT(
860
+ img_size=512,
861
+ num_classes=num_classes,
862
+ embed_dims=[96, 192, 384, 576],
863
+ depths=[2, 2, 6, 2],
864
+ num_heads=[3, 6, 12, 18],
865
+ window_sizes=[16, 16, 32, 16],
866
+ drop_path_rate=drop_path_rate,
867
+ )