Masked Autoencoder with VisionTransformer backbone
| 20 | |
| 21 | |
| 22 | class MaskedAutoencoderViT(nn.Module): |
| 23 | """ Masked Autoencoder with VisionTransformer backbone |
| 24 | """ |
| 25 | def __init__(self, img_size=224, patch_size=16, in_chans=3, |
| 26 | embed_dim=1024, depth=24, num_heads=16, |
| 27 | decoder_embed_dim=512, decoder_depth=8, decoder_num_heads=16, |
| 28 | mlp_ratio=4., norm_layer=nn.LayerNorm, norm_pix_loss=False): |
| 29 | super().__init__() |
| 30 | |
| 31 | # -------------------------------------------------------------------------- |
| 32 | # MAE encoder specifics |
| 33 | self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) |
| 34 | num_patches = self.patch_embed.num_patches |
| 35 | |
| 36 | self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) |
| 37 | self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim), requires_grad=False) # fixed sin-cos embedding |
| 38 | |
| 39 | self.blocks = nn.ModuleList([ |
| 40 | Block(embed_dim, num_heads, mlp_ratio, qkv_bias=True, norm_layer=norm_layer) |
| 41 | for i in range(depth)]) |
| 42 | self.norm = norm_layer(embed_dim) |
| 43 | # -------------------------------------------------------------------------- |
| 44 | |
| 45 | # -------------------------------------------------------------------------- |
| 46 | # MAE decoder specifics |
| 47 | self.decoder_embed = nn.Linear(embed_dim, decoder_embed_dim, bias=True) |
| 48 | |
| 49 | self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_embed_dim)) |
| 50 | |
| 51 | self.decoder_pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, decoder_embed_dim), requires_grad=False) # fixed sin-cos embedding |
| 52 | |
| 53 | self.decoder_blocks = nn.ModuleList([ |
| 54 | Block(decoder_embed_dim, decoder_num_heads, mlp_ratio, qkv_bias=True, norm_layer=norm_layer) |
| 55 | for i in range(decoder_depth)]) |
| 56 | |
| 57 | self.decoder_norm = norm_layer(decoder_embed_dim) |
| 58 | self.decoder_pred = nn.Linear(decoder_embed_dim, patch_size**2 * in_chans, bias=True) # decoder to patch |
| 59 | # -------------------------------------------------------------------------- |
| 60 | |
| 61 | self.norm_pix_loss = norm_pix_loss |
| 62 | |
| 63 | self.initialize_weights() |
| 64 | |
| 65 | def initialize_weights(self): |
| 66 | # initialization |
| 67 | # initialize (and freeze) pos_embed by sin-cos embedding |
| 68 | grid_size = int(self.patch_embed.num_patches**.5) |
| 69 | pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (grid_size, grid_size), 1, 0) |
| 70 | self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0)) |
| 71 | |
| 72 | decoder_grid_size = int(self.patch_embed.num_patches**.5) |
| 73 | decoder_pos_embed = get_2d_sincos_pos_embed(self.decoder_pos_embed.shape[-1], (decoder_grid_size, decoder_grid_size), 1, 0) |
| 74 | self.decoder_pos_embed.data.copy_(torch.from_numpy(decoder_pos_embed).float().unsqueeze(0)) |
| 75 | |
| 76 | # initialize patch_embed like nn.Linear (instead of nn.Conv2d) |
| 77 | w = self.patch_embed.proj.weight.data |
| 78 | torch.nn.init.xavier_uniform_(w.view([w.shape[0], -1])) |
| 79 |
no outgoing calls
no test coverage detected