MCPcopy Create free account
hub / github.com/LTH14/mar / __init__

Method __init__

models/mar.py:25–105  ·  view source on GitHub ↗
(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,
                 )

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 2

initialize_weightsMethod · 0.95
DiffLossClass · 0.90

Tested by

no test coverage detected