(modelname, pretrain_img_size, **kw)
| 758 | |
| 759 | |
| 760 | def 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 | |
| 799 | if __name__ == '__main__': |
no test coverage detected