MCPcopy Create free account
hub / github.com/OpenGVLab/UniFormerV2 / MViT

Class MViT

slowfast/models/video_model_builder.py:765–1013  ·  view source on GitHub ↗

Multiscale Vision Transformers Haoqi Fan, Bo Xiong, Karttikeya Mangalam, Yanghao Li, Zhicheng Yan, Jitendra Malik, Christoph Feichtenhofer https://arxiv.org/abs/2104.11227

Source from the content-addressed store, hash-verified

763
764@MODEL_REGISTRY.register()
765class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected