Spaces:
Sleeping
Sleeping
| # Copyright (c) Meta Platforms, Inc. and affiliates. | |
| # All rights reserved. | |
| # This source code is licensed under the license found in the | |
| # LICENSE file in the root directory of this source tree. | |
| # -------------------------------------------------------- | |
| # References: | |
| # timm: https://github.com/rwightman/pytorch-image-models/tree/master/timm | |
| # DeiT: https://github.com/facebookresearch/deit | |
| # -------------------------------------------------------- | |
| import sys | |
| sys.path.append(".") | |
| sys.path.append("..") | |
| from functools import partial | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| from torch.utils.checkpoint import checkpoint | |
| from utils import count_parameters | |
| from timm.models.vision_transformer import PatchEmbed, Block | |
| def load_part_mae_model(ckpt_path, model): | |
| checkpoint = torch.load(ckpt_path, map_location='cpu') | |
| load_state_dict = model.state_dict() | |
| for k in load_state_dict.keys(): | |
| load_state_dict[k] = checkpoint['model'][k] | |
| model.load_state_dict(load_state_dict, strict=False) | |
| return model | |
| def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False): | |
| """ | |
| grid_size: int of the grid height and width | |
| return: | |
| pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token) | |
| """ | |
| 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) # here w goes first | |
| grid = np.stack(grid, axis=0) | |
| grid = grid.reshape([2, 1, grid_size, grid_size]) | |
| pos_embed = 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(embed_dim, grid): | |
| assert embed_dim % 2 == 0 | |
| # use half of dimensions to encode grid_h | |
| emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2) | |
| emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2) | |
| emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D) | |
| return emb | |
| def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): | |
| """ | |
| embed_dim: output dimension for each position | |
| pos: a list of positions to be encoded: size (M,) | |
| out: (M, D) | |
| """ | |
| assert embed_dim % 2 == 0 | |
| omega = np.arange(embed_dim // 2, dtype=np.float32) | |
| omega /= embed_dim / 2. | |
| omega = 1. / 10000**omega # (D/2,) | |
| pos = pos.reshape(-1) # (M,) | |
| out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product | |
| emb_sin = np.sin(out) # (M, D/2) | |
| emb_cos = np.cos(out) # (M, D/2) | |
| emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) | |
| return emb | |
| class MaskedAutoencoderViT(nn.Module): | |
| """ Masked Autoencoder with VisionTransformer backbone | |
| """ | |
| def __init__(self, device, img_size=224, patch_size=16, in_chans=3, | |
| embed_dim=1024, depth=24, num_heads=16, | |
| mlp_ratio=4., norm_layer=nn.LayerNorm, | |
| grad_checkpointing=False, | |
| **kwargs): | |
| super().__init__() | |
| self.grad_checkpointing = grad_checkpointing | |
| # -------------------------------------------------------------------------- | |
| # MAE encoder specifics | |
| self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) | |
| num_patches = self.patch_embed.num_patches | |
| self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) | |
| self.pos_embed = torch.zeros(1, num_patches + 1, embed_dim) | |
| pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], int(self.patch_embed.num_patches**.5), cls_token=True) | |
| self.pos_embed = torch.from_numpy(pos_embed).float().unsqueeze(0).to(device) | |
| self.blocks = nn.ModuleList([ | |
| Block(embed_dim, num_heads, mlp_ratio, qkv_bias=True, norm_layer=norm_layer) | |
| for i in range(depth)]) | |
| self.norm = norm_layer(embed_dim) | |
| self.imagenet_mean = torch.tensor([0.485, 0.456, 0.406])[..., None, None].to(device) | |
| self.imagenet_std = torch.tensor([0.229, 0.224, 0.225])[..., None, None].to(device) | |
| def patchify(self, imgs): | |
| """ | |
| imgs: (N, 3, H, W) | |
| x: (N, L, patch_size**2 *3) | |
| """ | |
| p = self.patch_embed.patch_size[0] | |
| assert imgs.shape[2] == imgs.shape[3] and imgs.shape[2] % p == 0 | |
| h = w = imgs.shape[2] // p | |
| x = imgs.reshape(shape=(imgs.shape[0], 3, h, p, w, p)) | |
| x = torch.einsum('nchpwq->nhwpqc', x) | |
| x = x.reshape(shape=(imgs.shape[0], h * w, p**2 * 3)) | |
| return x | |
| def unpatchify(self, x): | |
| """ | |
| x: (N, L, patch_size**2 *3) | |
| imgs: (N, 3, H, W) | |
| """ | |
| p = self.patch_embed.patch_size[0] | |
| h = w = int(x.shape[1]**.5) | |
| assert h * w == x.shape[1] | |
| x = x.reshape(shape=(x.shape[0], h, w, p, p, 3)) | |
| x = torch.einsum('nhwpqc->nchpwq', x) | |
| imgs = x.reshape(shape=(x.shape[0], 3, h * p, h * p)) | |
| return imgs | |
| def random_masking(self, x, mask_ratio): | |
| """ | |
| Perform per-sample random masking by per-sample shuffling. | |
| Per-sample shuffling is done by argsort random noise. | |
| x: [N, L, D], sequence | |
| """ | |
| N, L, D = x.shape # batch, length, dim | |
| len_keep = int(L * (1 - mask_ratio)) | |
| noise = torch.rand(N, L, device=x.device) # noise in [0, 1] | |
| # sort noise for each sample | |
| ids_shuffle = torch.argsort(noise, dim=1) # ascend: small is keep, large is remove | |
| ids_restore = torch.argsort(ids_shuffle, dim=1) | |
| # keep the first subset | |
| ids_keep = ids_shuffle[:, :len_keep] | |
| x_masked = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D)) | |
| # generate the binary mask: 0 is keep, 1 is remove | |
| mask = torch.ones([N, L], device=x.device) | |
| mask[:, :len_keep] = 0 | |
| # unshuffle to get the binary mask | |
| mask = torch.gather(mask, dim=1, index=ids_restore) | |
| return x_masked, mask, ids_restore | |
| def forward_encoder(self, x, mask_ratio): | |
| x = (x - self.imagenet_mean) / self.imagenet_std | |
| # embed patches | |
| x = self.patch_embed(x) | |
| # add pos embed w/o cls token | |
| x = x + self.pos_embed[:, 1:, :] | |
| # masking: length -> length * mask_ratio | |
| x, mask, ids_restore = self.random_masking(x, mask_ratio) | |
| # append cls token | |
| cls_token = self.cls_token + self.pos_embed[:, :1, :] | |
| cls_tokens = cls_token.expand(x.shape[0], -1, -1) | |
| x = torch.cat((cls_tokens, x), dim=1) | |
| # apply Transformer blocks | |
| for blk in self.blocks: | |
| if self.grad_checkpointing and self.training: | |
| x = checkpoint(blk, x, use_reentrant=False) | |
| else: | |
| x = blk(x) | |
| x = self.norm(x) | |
| return x, mask, ids_restore | |
| def forward(self, x): | |
| x = (x - self.imagenet_mean) / self.imagenet_std | |
| # embed patches | |
| x = self.patch_embed(x) | |
| # add pos embed w/o cls token | |
| x = x + self.pos_embed[:, 1:, :] | |
| # append cls token | |
| cls_token = self.cls_token + self.pos_embed[:, :1, :] | |
| cls_tokens = cls_token.expand(x.shape[0], -1, -1) | |
| x = torch.cat((cls_tokens, x), dim=1) | |
| # apply Transformer blocks | |
| for blk in self.blocks: | |
| if self.grad_checkpointing and self.training: | |
| x = checkpoint(blk, x, use_reentrant=False) | |
| else: | |
| x = blk(x) | |
| x = self.norm(x) | |
| return x | |
| def mae_vit_base_patch16_dec512d8b(**kwargs): | |
| model = MaskedAutoencoderViT( | |
| patch_size=16, embed_dim=768, depth=12, num_heads=12, | |
| decoder_embed_dim=512, decoder_depth=8, decoder_num_heads=16, | |
| mlp_ratio=4, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs) | |
| return model | |
| def mae_vit_large_patch16_dec512d8b(**kwargs): | |
| model = MaskedAutoencoderViT( | |
| patch_size=16, embed_dim=1024, depth=24, num_heads=16, | |
| decoder_embed_dim=512, decoder_depth=8, decoder_num_heads=16, | |
| mlp_ratio=4, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs) | |
| return model | |
| def mae_vit_huge_patch14_dec512d8b(**kwargs): | |
| model = MaskedAutoencoderViT( | |
| patch_size=14, embed_dim=1280, depth=32, num_heads=16, | |
| decoder_embed_dim=512, decoder_depth=8, decoder_num_heads=16, | |
| mlp_ratio=4, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs) | |
| return model | |
| # set recommended archs | |
| mae_vit_base_patch16 = mae_vit_base_patch16_dec512d8b # decoder: 512 dim, 8 blocks | |
| mae_vit_large_patch16 = mae_vit_large_patch16_dec512d8b # decoder: 512 dim, 8 blocks | |
| mae_vit_huge_patch14 = mae_vit_huge_patch14_dec512d8b # decoder: 512 dim, 8 blocks | |
| def get_mae_encoder(img_size, model_type, device, pretrained=True, **kwargs): | |
| if model_type == "base": | |
| model = mae_vit_base_patch16(img_size=img_size, device=device, **kwargs) | |
| # model_path = 'pretrained/mae/mae_pretrain_vit_base.pth' | |
| model_path = 'pretrained/mae/mae_visualize_vit_base.pth' | |
| elif model_type == "large": | |
| model = mae_vit_large_patch16(img_size=img_size, device=device, **kwargs) | |
| model_path = 'pretrained/mae/mae_visualize_vit_large_ganloss.pth' | |
| else: | |
| raise NotImplementedError(f"Model type {model_type} not implemented.") | |
| if pretrained: | |
| model = load_part_mae_model(model_path, model) | |
| print("#" * 50) | |
| print("[HYX INFO] Loaded pretrained MAE encoder from:", model_path) | |
| print("#" * 50) | |
| else: | |
| print("#" * 50) | |
| print("[HYX INFO] DO NOT LOAD PRETRAINED WEIGHTS") | |
| print("#" * 50) | |
| model = model.to(device) | |
| return model | |
| if __name__ == "__main__": | |
| model = get_mae_encoder(img_size=224, model_type="base", device="cpu") | |
| count_parameters(model, detailed=True) | |