MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / MaskedAutoencoderViT

Class MaskedAutoencoderViT

SwissArmyTransformer/examples/mae/models_mae.py:22–222  ·  view source on GitHub ↗

Masked Autoencoder with VisionTransformer backbone

Source from the content-addressed store, hash-verified

20
21
22class 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

Calls

no outgoing calls

Tested by

no test coverage detected