Spaces:
Sleeping
Sleeping
| import torch | |
| import torch.nn as nn | |
| import numpy as np | |
| from .transformer_block import TransformerBlock | |
| class Decoder(nn.Module): | |
| """ | |
| MAE Decoder based on Vision Transformer. | |
| """ | |
| def __init__( | |
| self, | |
| num_patches: int, | |
| patch_size: int = 16, | |
| in_chans: int = 3, | |
| embed_dim: int = 768, | |
| decoder_embed_dim: int = 512, | |
| decoder_depth: int = 8, | |
| decoder_num_heads: int = 16, | |
| mlp_ratio: float = 4.0, | |
| norm_layer: nn.Module = nn.LayerNorm | |
| ): | |
| super().__init__() | |
| self.decoder_embed = nn.Linear(embed_dim, decoder_embed_dim, bias=True) | |
| self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_embed_dim)) | |
| self.decoder_pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, decoder_embed_dim), requires_grad=False) | |
| self.decoder_blocks = nn.ModuleList([ | |
| TransformerBlock(decoder_embed_dim, decoder_num_heads, mlp_ratio, qkv_bias=True, norm_layer=norm_layer) | |
| for _ in range(decoder_depth) | |
| ]) | |
| self.decoder_norm = norm_layer(decoder_embed_dim) | |
| self.decoder_pred = nn.Linear(decoder_embed_dim, patch_size**2 * in_chans, bias=True) | |
| self.initialize_weights() | |
| def initialize_weights(self): | |
| # Initialize pos_embed with sin-cos embedding | |
| grid_size = int(self.decoder_pos_embed.shape[1] - 1)**.5 | |
| pos_embed = self.get_2d_sincos_pos_embed(self.decoder_pos_embed.shape[-1], int(grid_size), cls_token=True) | |
| self.decoder_pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0)) | |
| # Initialize mask_token | |
| torch.nn.init.normal_(self.mask_token, std=.02) | |
| # Initialize transformer blocks | |
| self.apply(self._init_weights) | |
| def _init_weights(self, m): | |
| if isinstance(m, nn.Linear): | |
| torch.nn.init.xavier_uniform_(m.weight) | |
| if isinstance(m, nn.Linear) and m.bias is not None: | |
| nn.init.constant_(m.bias, 0) | |
| elif isinstance(m, nn.LayerNorm): | |
| nn.init.constant_(m.bias, 0) | |
| nn.init.constant_(m.weight, 1.0) | |
| def get_2d_sincos_pos_embed(self, embed_dim, grid_size, cls_token=False): | |
| # Re-using the logic from Encoder (could be moved to utils for DRY) | |
| grid_h = np.arange(grid_size, dtype=np.float32) | |
| grid_w = np.arange(grid_size, dtype=np.float32) | |
| grid = np.meshgrid(grid_w, grid_h) | |
| grid = np.stack(grid, axis=0) | |
| grid = grid.reshape([2, 1, grid_size, grid_size]) | |
| pos_embed = self.get_2d_sincos_pos_embed_from_grid(embed_dim, grid) | |
| if cls_token: | |
| pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0) | |
| return pos_embed | |
| def get_2d_sincos_pos_embed_from_grid(self, embed_dim, grid): | |
| assert embed_dim % 2 == 0 | |
| emb_h = self.get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) | |
| emb_w = self.get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) | |
| emb = np.concatenate([emb_h, emb_w], axis=1) | |
| return emb | |
| def get_1d_sincos_pos_embed_from_grid(self, embed_dim, pos): | |
| assert embed_dim % 2 == 0 | |
| omega = np.arange(embed_dim // 2, dtype=np.float32) | |
| omega /= embed_dim / 2. | |
| omega = 1. / 10000**omega | |
| pos = pos.reshape(-1) | |
| out = np.einsum('m,d->md', pos, omega) | |
| emb_sin = np.sin(out) | |
| emb_cos = np.cos(out) | |
| emb = np.concatenate([emb_sin, emb_cos], axis=1) | |
| return emb | |
| def forward(self, x: torch.Tensor, ids_restore: torch.Tensor) -> torch.Tensor: | |
| # Project to decoder dimension | |
| x = self.decoder_embed(x) | |
| # Append mask tokens to sequence | |
| mask_tokens = self.mask_token.repeat(x.shape[0], ids_restore.shape[1] + 1 - x.shape[1], 1) | |
| x_ = torch.cat([x[:, 1:, :], mask_tokens], dim=1) # no cls token | |
| x_ = torch.gather(x_, dim=1, index=ids_restore.unsqueeze(-1).repeat(1, 1, x.shape[2])) # unshuffle | |
| x = torch.cat([x[:, :1, :], x_], dim=1) # append cls token | |
| # Add pos embed | |
| x = x + self.decoder_pos_embed | |
| # Apply transformer blocks | |
| for block in self.decoder_blocks: | |
| x = block(x) | |
| x = self.decoder_norm(x) | |
| # Predict pixels | |
| x = self.decoder_pred(x) | |
| # Remove cls token | |
| x = x[:, 1:, :] | |
| return x | |