MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / build_swin_transformer

Function build_swin_transformer

models/aios/backbones/swin_transformer.py:760–796  ·  view source on GitHub ↗
(modelname, pretrain_img_size, **kw)

Source from the content-addressed store, hash-verified

758
759
760def build_swin_transformer(modelname, pretrain_img_size, **kw):
761 assert modelname in [
762 'swin_T_224_1k', 'swin_B_224_22k', 'swin_B_384_22k', 'swin_L_224_22k',
763 'swin_L_384_22k'
764 ]
765
766 model_para_dict = {
767 'swin_T_224_1k':
768 dict(embed_dim=96,
769 depths=[2, 2, 6, 2],
770 num_heads=[3, 6, 12, 24],
771 window_size=7),
772 'swin_B_224_22k':
773 dict(embed_dim=128,
774 depths=[2, 2, 18, 2],
775 num_heads=[4, 8, 16, 32],
776 window_size=7),
777 'swin_B_384_22k':
778 dict(embed_dim=128,
779 depths=[2, 2, 18, 2],
780 num_heads=[4, 8, 16, 32],
781 window_size=12),
782 'swin_L_224_22k':
783 dict(embed_dim=192,
784 depths=[2, 2, 18, 2],
785 num_heads=[6, 12, 24, 48],
786 window_size=7),
787 'swin_L_384_22k':
788 dict(embed_dim=192,
789 depths=[2, 2, 18, 2],
790 num_heads=[6, 12, 24, 48],
791 window_size=12),
792 }
793 kw_cgf = model_para_dict[modelname]
794 kw_cgf.update(kw)
795 model = SwinTransformer(pretrain_img_size=pretrain_img_size, **kw_cgf)
796 return model
797
798
799if __name__ == '__main__':

Callers 2

build_backboneFunction · 0.85

Calls 2

SwinTransformerClass · 0.85
updateMethod · 0.45

Tested by

no test coverage detected