MCPcopy Create free account
hub / github.com/microsoft/Cream / EfficientViT

Class EfficientViT

EfficientViT/classification/model/efficientvit.py:285–356  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

283
284
285class EfficientViT(torch.nn.Module):
286 def __init__(self, img_size=224,
287 patch_size=16,
288 in_chans=3,
289 num_classes=1000,
290 stages=['s', 's', 's'],
291 embed_dim=[64, 128, 192],
292 key_dim=[16, 16, 16],
293 depth=[1, 2, 3],
294 num_heads=[4, 4, 4],
295 window_size=[7, 7, 7],
296 kernels=[5, 5, 5, 5],
297 down_ops=[['subsample', 2], ['subsample', 2], ['']],
298 distillation=False,):
299 super().__init__()
300
301 resolution = img_size
302 # Patch embedding
303 self.patch_embed = torch.nn.Sequential(Conv2d_BN(in_chans, embed_dim[0] // 8, 3, 2, 1, resolution=resolution), torch.nn.ReLU(),
304 Conv2d_BN(embed_dim[0] // 8, embed_dim[0] // 4, 3, 2, 1, resolution=resolution // 2), torch.nn.ReLU(),
305 Conv2d_BN(embed_dim[0] // 4, embed_dim[0] // 2, 3, 2, 1, resolution=resolution // 4), torch.nn.ReLU(),
306 Conv2d_BN(embed_dim[0] // 2, embed_dim[0], 3, 2, 1, resolution=resolution // 8))
307
308 resolution = img_size // patch_size
309 attn_ratio = [embed_dim[i] / (key_dim[i] * num_heads[i]) for i in range(len(embed_dim))]
310 self.blocks1 = []
311 self.blocks2 = []
312 self.blocks3 = []
313
314 # Build EfficientViT blocks
315 for i, (stg, ed, kd, dpth, nh, ar, wd, do) in enumerate(
316 zip(stages, embed_dim, key_dim, depth, num_heads, attn_ratio, window_size, down_ops)):
317 for d in range(dpth):
318 eval('self.blocks' + str(i+1)).append(EfficientViTBlock(stg, ed, kd, nh, ar, resolution, wd, kernels))
319 if do[0] == 'subsample':
320 # Build EfficientViT downsample block
321 #('Subsample' stride)
322 blk = eval('self.blocks' + str(i+2))
323 resolution_ = (resolution - 1) // do[1] + 1
324 blk.append(torch.nn.Sequential(Residual(Conv2d_BN(embed_dim[i], embed_dim[i], 3, 1, 1, groups=embed_dim[i], resolution=resolution)),
325 Residual(FFN(embed_dim[i], int(embed_dim[i] * 2), resolution)),))
326 blk.append(PatchMerging(*embed_dim[i:i + 2], resolution))
327 resolution = resolution_
328 blk.append(torch.nn.Sequential(Residual(Conv2d_BN(embed_dim[i + 1], embed_dim[i + 1], 3, 1, 1, groups=embed_dim[i + 1], resolution=resolution)),
329 Residual(FFN(embed_dim[i + 1], int(embed_dim[i + 1] * 2), resolution)),))
330 self.blocks1 = torch.nn.Sequential(*self.blocks1)
331 self.blocks2 = torch.nn.Sequential(*self.blocks2)
332 self.blocks3 = torch.nn.Sequential(*self.blocks3)
333
334 # Classification head
335 self.head = BN_Linear(embed_dim[-1], num_classes) if num_classes > 0 else torch.nn.Identity()
336 self.distillation = distillation
337 if distillation:
338 self.head_dist = BN_Linear(embed_dim[-1], num_classes) if num_classes > 0 else torch.nn.Identity()
339
340 @torch.jit.ignore
341 def no_weight_decay(self):
342 return {x for x in self.state_dict().keys() if 'attention_biases' in x}

Callers 6

EfficientViT_M0Function · 0.90
EfficientViT_M1Function · 0.90
EfficientViT_M2Function · 0.90
EfficientViT_M3Function · 0.90
EfficientViT_M4Function · 0.90
EfficientViT_M5Function · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected