MCPcopy Create free account
hub / github.com/coperception/star / MultiAgentMaskedAutoencoderViT

Class MultiAgentMaskedAutoencoderViT

star/models/mae_base.py:257–750  ·  view source on GitHub ↗

Masked Autoencoder with VisionTransformer backbone

Source from the content-addressed store, hash-verified

255
256
257class 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":

Calls

no outgoing calls

Tested by

no test coverage detected