| 283 | |
| 284 | |
| 285 | class 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} |
no outgoing calls
no test coverage detected