(self, img_size=256, vae_stride=16, patch_size=1,
encoder_embed_dim=1024, encoder_depth=16, encoder_num_heads=16,
decoder_embed_dim=1024, decoder_depth=16, decoder_num_heads=16,
mlp_ratio=4., norm_layer=nn.LayerNorm,
vae_embed_dim=16,
mask_ratio_min=0.7,
label_drop_prob=0.1,
class_num=1000,
attn_dropout=0.1,
proj_dropout=0.1,
buffer_size=64,
diffloss_d=3,
diffloss_w=1024,
num_sampling_steps='100',
diffusion_batch_mul=4,
grad_checkpointing=False,
)
| 23 | """ Masked Autoencoder with VisionTransformer backbone |
| 24 | """ |
| 25 | def __init__(self, img_size=256, vae_stride=16, patch_size=1, |
| 26 | encoder_embed_dim=1024, encoder_depth=16, encoder_num_heads=16, |
| 27 | decoder_embed_dim=1024, decoder_depth=16, decoder_num_heads=16, |
| 28 | mlp_ratio=4., norm_layer=nn.LayerNorm, |
| 29 | vae_embed_dim=16, |
| 30 | mask_ratio_min=0.7, |
| 31 | label_drop_prob=0.1, |
| 32 | class_num=1000, |
| 33 | attn_dropout=0.1, |
| 34 | proj_dropout=0.1, |
| 35 | buffer_size=64, |
| 36 | diffloss_d=3, |
| 37 | diffloss_w=1024, |
| 38 | num_sampling_steps='100', |
| 39 | diffusion_batch_mul=4, |
| 40 | grad_checkpointing=False, |
| 41 | ): |
| 42 | super().__init__() |
| 43 | |
| 44 | # -------------------------------------------------------------------------- |
| 45 | # VAE and patchify specifics |
| 46 | self.vae_embed_dim = vae_embed_dim |
| 47 | |
| 48 | self.img_size = img_size |
| 49 | self.vae_stride = vae_stride |
| 50 | self.patch_size = patch_size |
| 51 | self.seq_h = self.seq_w = img_size // vae_stride // patch_size |
| 52 | self.seq_len = self.seq_h * self.seq_w |
| 53 | self.token_embed_dim = vae_embed_dim * patch_size**2 |
| 54 | self.grad_checkpointing = grad_checkpointing |
| 55 | |
| 56 | # -------------------------------------------------------------------------- |
| 57 | # Class Embedding |
| 58 | self.num_classes = class_num |
| 59 | self.class_emb = nn.Embedding(class_num, encoder_embed_dim) |
| 60 | self.label_drop_prob = label_drop_prob |
| 61 | # Fake class embedding for CFG's unconditional generation |
| 62 | self.fake_latent = nn.Parameter(torch.zeros(1, encoder_embed_dim)) |
| 63 | |
| 64 | # -------------------------------------------------------------------------- |
| 65 | # MAR variant masking ratio, a left-half truncated Gaussian centered at 100% masking ratio with std 0.25 |
| 66 | self.mask_ratio_generator = stats.truncnorm((mask_ratio_min - 1.0) / 0.25, 0, loc=1.0, scale=0.25) |
| 67 | |
| 68 | # -------------------------------------------------------------------------- |
| 69 | # MAR encoder specifics |
| 70 | self.z_proj = nn.Linear(self.token_embed_dim, encoder_embed_dim, bias=True) |
| 71 | self.z_proj_ln = nn.LayerNorm(encoder_embed_dim, eps=1e-6) |
| 72 | self.buffer_size = buffer_size |
| 73 | self.encoder_pos_embed_learned = nn.Parameter(torch.zeros(1, self.seq_len + self.buffer_size, encoder_embed_dim)) |
| 74 | |
| 75 | self.encoder_blocks = nn.ModuleList([ |
| 76 | Block(encoder_embed_dim, encoder_num_heads, mlp_ratio, qkv_bias=True, norm_layer=norm_layer, |
| 77 | proj_drop=proj_dropout, attn_drop=attn_dropout) for _ in range(encoder_depth)]) |
| 78 | self.encoder_norm = norm_layer(encoder_embed_dim) |
| 79 | |
| 80 | # -------------------------------------------------------------------------- |
| 81 | # MAR decoder specifics |
| 82 | self.decoder_embed = nn.Linear(encoder_embed_dim, decoder_embed_dim, bias=True) |
nothing calls this directly
no test coverage detected