Masked Autoencoder with VisionTransformer backbone
| 255 | |
| 256 | |
| 257 | class MultiAgentMaskedAutoencoderViT(nn.Module): |
| 258 | """ Masked Autoencoder with VisionTransformer backbone |
| 259 | """ |
| 260 | def __init__(self, img_size=224, patch_size=16, in_chans=3, |
| 261 | embed_dim=1024, depth=24, num_heads=16, |
| 262 | decoder_embed_dim=512, decoder_depth=8, decoder_num_heads=16, |
| 263 | mlp_ratio=4., norm_layer=nn.LayerNorm, decoder_head="mlp", norm_pix_loss=False, time_stamp=1, mask_method="random", |
| 264 | inter_emb_dim=32): |
| 265 | super().__init__() |
| 266 | |
| 267 | # -------------------------------------------------------------------------- |
| 268 | # MAE encoder specifics |
| 269 | self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) |
| 270 | num_patches = self.patch_embed.num_patches |
| 271 | |
| 272 | self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) |
| 273 | self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim), requires_grad=False) # fixed sin-cos embedding |
| 274 | self.time_stamp = time_stamp |
| 275 | # temporal embeddings |
| 276 | self.temp_embed = nn.Parameter(torch.zeros(1, self.time_stamp, embed_dim)) # learnable temporal embeddings for 2 time stamp |
| 277 | self.decoder_temp_embed = nn.Parameter(torch.zeros(1, self.time_stamp, 1, decoder_embed_dim)) # learnable temporal embeddings for 2 time stamp |
| 278 | |
| 279 | self.blocks = nn.ModuleList([ |
| 280 | Block(embed_dim, num_heads, mlp_ratio, qkv_bias=True, qk_scale=None, norm_layer=norm_layer) |
| 281 | for i in range(depth)]) |
| 282 | self.norm = norm_layer(embed_dim) |
| 283 | |
| 284 | self.compressor = nn.Sequential( |
| 285 | nn.Linear(embed_dim, inter_emb_dim), |
| 286 | nn.ReLU(), |
| 287 | ) |
| 288 | # -------------------------------------------------------------------------- |
| 289 | |
| 290 | # -------------------------------------------------------------------------- |
| 291 | # MAE decoder specifics |
| 292 | self.decompressor = nn.Linear(inter_emb_dim, embed_dim) |
| 293 | self.decoder_embed = nn.Linear(embed_dim, decoder_embed_dim, bias=True) |
| 294 | |
| 295 | # try option1: |
| 296 | # self.decoder_embed = nn.Linear(inter_emb_dim, decoder_embed_dim, bias=True) |
| 297 | # try option2: |
| 298 | # self.decoder_embed = nn.Sequential( |
| 299 | # nn.Linear(inter_emb_dim, embed_dim), |
| 300 | # nn.ReLU(), |
| 301 | # nn.Linear(embed_dim, decoder_embed_dim, bias=True) |
| 302 | # ) |
| 303 | |
| 304 | self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_embed_dim)) |
| 305 | |
| 306 | self.decoder_pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, decoder_embed_dim), requires_grad=False) # fixed sin-cos embedding |
| 307 | |
| 308 | self.decoder_blocks = nn.ModuleList([ |
| 309 | Block(decoder_embed_dim, decoder_num_heads, mlp_ratio, qkv_bias=True, qk_scale=None, norm_layer=norm_layer) |
| 310 | for i in range(decoder_depth)]) |
| 311 | |
| 312 | self.decoder_norm = norm_layer(decoder_embed_dim) |
| 313 | |
| 314 | if decoder_head == "mlp": |
no outgoing calls
no test coverage detected