Multiscale Vision Transformers Haoqi Fan, Bo Xiong, Karttikeya Mangalam, Yanghao Li, Zhicheng Yan, Jitendra Malik, Christoph Feichtenhofer https://arxiv.org/abs/2104.11227
| 763 | |
| 764 | @MODEL_REGISTRY.register() |
| 765 | class MViT(nn.Module): |
| 766 | """ |
| 767 | Multiscale Vision Transformers |
| 768 | Haoqi Fan, Bo Xiong, Karttikeya Mangalam, Yanghao Li, Zhicheng Yan, Jitendra Malik, Christoph Feichtenhofer |
| 769 | https://arxiv.org/abs/2104.11227 |
| 770 | """ |
| 771 | |
| 772 | def __init__(self, cfg): |
| 773 | super().__init__() |
| 774 | # Get parameters. |
| 775 | # assert cfg.DATA.TRAIN_CROP_SIZE == cfg.DATA.TEST_CROP_SIZE |
| 776 | self.cfg = cfg |
| 777 | # Prepare input. |
| 778 | spatial_size = cfg.DATA.TRAIN_CROP_SIZE |
| 779 | temporal_size = cfg.DATA.NUM_FRAMES |
| 780 | in_chans = cfg.DATA.INPUT_CHANNEL_NUM[0] |
| 781 | use_2d_patch = cfg.MVIT.PATCH_2D |
| 782 | self.patch_stride = cfg.MVIT.PATCH_STRIDE |
| 783 | if use_2d_patch: |
| 784 | self.patch_stride = [1] + self.patch_stride |
| 785 | # Prepare output. |
| 786 | num_classes = cfg.MODEL.NUM_CLASSES |
| 787 | embed_dim = cfg.MVIT.EMBED_DIM |
| 788 | # Prepare backbone |
| 789 | num_heads = cfg.MVIT.NUM_HEADS |
| 790 | mlp_ratio = cfg.MVIT.MLP_RATIO |
| 791 | qkv_bias = cfg.MVIT.QKV_BIAS |
| 792 | self.drop_rate = cfg.MVIT.DROPOUT_RATE |
| 793 | depth = cfg.MVIT.DEPTH |
| 794 | drop_path_rate = cfg.MVIT.DROPPATH_RATE |
| 795 | mode = cfg.MVIT.MODE |
| 796 | self.cls_embed_on = cfg.MVIT.CLS_EMBED_ON |
| 797 | self.sep_pos_embed = cfg.MVIT.SEP_POS_EMBED |
| 798 | if cfg.MVIT.NORM == "layernorm": |
| 799 | norm_layer = partial(nn.LayerNorm, eps=1e-6) |
| 800 | else: |
| 801 | raise NotImplementedError("Only supports layernorm.") |
| 802 | self.num_classes = num_classes |
| 803 | self.patch_embed = stem_helper.PatchEmbed( |
| 804 | dim_in=in_chans, |
| 805 | dim_out=embed_dim, |
| 806 | kernel=cfg.MVIT.PATCH_KERNEL, |
| 807 | stride=cfg.MVIT.PATCH_STRIDE, |
| 808 | padding=cfg.MVIT.PATCH_PADDING, |
| 809 | conv_2d=use_2d_patch, |
| 810 | ) |
| 811 | self.input_dims = [temporal_size, spatial_size, spatial_size] |
| 812 | assert self.input_dims[1] == self.input_dims[2] |
| 813 | self.patch_dims = [ |
| 814 | self.input_dims[i] // self.patch_stride[i] |
| 815 | for i in range(len(self.input_dims)) |
| 816 | ] |
| 817 | num_patches = math.prod(self.patch_dims) |
| 818 | |
| 819 | dpr = [ |
| 820 | x.item() for x in torch.linspace(0, drop_path_rate, depth) |
| 821 | ] # stochastic depth decay rule |
| 822 |
nothing calls this directly
no outgoing calls
no test coverage detected