(self, img_size=256, patch_size=16, in_chans=3,
embed_dim=1024, depth=24, num_heads=16,
decoder_embed_dim=512, decoder_depth=8, decoder_num_heads=16,
mlp_ratio=4., norm_layer=nn.LayerNorm, norm_pix_loss=False,
mask_ratio_min=0.5, mask_ratio_max=1.0, mask_ratio_mu=0.55, mask_ratio_std=0.25,
vqgan_ckpt_path='vqgan_jax_strongaug.ckpt', use_rep=True, rep_dim=256,
rep_drop_prob=0.0,
use_class_label=False,
pretrained_enc_arch='mocov3_vit_base',
pretrained_enc_path='pretrained_enc_ckpts/mocov3/vitb.pth.tar',
pretrained_enc_proj_dim=256,
pretrained_enc_withproj=False,
pretrained_rdm_ckpt=None,
pretrained_rdm_cfg=None)
| 162 | """ Masked Autoencoder with VisionTransformer backbone |
| 163 | """ |
| 164 | def __init__(self, img_size=256, patch_size=16, in_chans=3, |
| 165 | embed_dim=1024, depth=24, num_heads=16, |
| 166 | decoder_embed_dim=512, decoder_depth=8, decoder_num_heads=16, |
| 167 | mlp_ratio=4., norm_layer=nn.LayerNorm, norm_pix_loss=False, |
| 168 | mask_ratio_min=0.5, mask_ratio_max=1.0, mask_ratio_mu=0.55, mask_ratio_std=0.25, |
| 169 | vqgan_ckpt_path='vqgan_jax_strongaug.ckpt', use_rep=True, rep_dim=256, |
| 170 | rep_drop_prob=0.0, |
| 171 | use_class_label=False, |
| 172 | pretrained_enc_arch='mocov3_vit_base', |
| 173 | pretrained_enc_path='pretrained_enc_ckpts/mocov3/vitb.pth.tar', |
| 174 | pretrained_enc_proj_dim=256, |
| 175 | pretrained_enc_withproj=False, |
| 176 | pretrained_rdm_ckpt=None, |
| 177 | pretrained_rdm_cfg=None): |
| 178 | super().__init__() |
| 179 | assert not (use_rep and use_class_label) |
| 180 | |
| 181 | # -------------------------------------------------------------------------- |
| 182 | # VQGAN specifics |
| 183 | vqgan_config = OmegaConf.load('config/mage/vqgan.yaml').model |
| 184 | self.vqgan_cfg = vqgan_config |
| 185 | |
| 186 | self.codebook_size = vqgan_config.params.n_embed |
| 187 | vocab_size = self.codebook_size + 1000 + 1 # 1024 codebook size, 1000 classes, 1 for mask token. |
| 188 | self.fake_class_label = self.codebook_size + 1100 - 1024 |
| 189 | self.mask_token_label = vocab_size - 1 |
| 190 | self.token_emb = BertEmbeddings(vocab_size=vocab_size, |
| 191 | hidden_size=embed_dim, |
| 192 | max_position_embeddings=256+1, |
| 193 | dropout=0.1) |
| 194 | self.use_rep = use_rep |
| 195 | self.use_class_label = use_class_label |
| 196 | if self.use_rep: |
| 197 | print("Use representation as condition!") |
| 198 | self.latent_prior_proj = nn.Linear(rep_dim, embed_dim, bias=True) |
| 199 | if self.use_class_label: |
| 200 | print("Use class label as condition!") |
| 201 | self.class_emb = nn.Embedding(1000, embed_dim) |
| 202 | |
| 203 | # CFG config |
| 204 | self.rep_drop_prob = rep_drop_prob |
| 205 | self.fake_latent = nn.Parameter(torch.zeros(1, rep_dim)) |
| 206 | torch.nn.init.normal_(self.fake_latent, std=.02) |
| 207 | |
| 208 | # MAGE variant masking ratio |
| 209 | self.mask_ratio_min = mask_ratio_min |
| 210 | self.mask_ratio_generator = stats.truncnorm((mask_ratio_min - mask_ratio_mu) / mask_ratio_std, |
| 211 | (mask_ratio_max - mask_ratio_mu) / mask_ratio_std, |
| 212 | loc=mask_ratio_mu, scale=mask_ratio_std) |
| 213 | |
| 214 | # -------------------------------------------------------------------------- |
| 215 | # MAGE encoder specifics |
| 216 | dropout_rate = 0.1 |
| 217 | self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) |
| 218 | num_patches = self.patch_embed.num_patches |
| 219 | |
| 220 | self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) |
| 221 | self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim), requires_grad=False) # fixed sin-cos embedding |
nothing calls this directly
no test coverage detected