o & ilK@sddlmZddlZddlmZddlmZddlmZddl Z ddl m Z ddl m Z ddlmZddlmZdd lmZmZmZdd lmZmZmZmZmZdd lmZgd Zd ddddddZGddde j Z!d*ddZ"Gddde!Z#Gd d!d!e!Z$Gd"d#d#e!Z%Gd$d%d%e!Z&Gd&d'd'e!Z'Gd(d)d)e!Z(e!Z)Z*e#Z+Z,Z-e$Z.Z/Z0e%Z1Z2Z3e&Z4Z5Z6e'Z7Z8Z9Z:e(Z;Z<Z=Z>dS)+) annotationsN) OrderedDict)Sequence)Any)load_state_dict_from_url) download_url) Convolution) SEBottleneckSEResNetBottleneckSEResNeXtBottleneck)ActConvDropoutNormPool)look_up_option)SENetSENet154 SEResNet50 SEResNet101 SEResNet152 SEResNeXt50 SEResNext101 SE_NET_MODELSzAhttp://data.lip6.fr/cadene/pretrainedmodels/senet154-c7b49a05.pthzDhttp://data.lip6.fr/cadene/pretrainedmodels/se_resnet50-ce0d4300.pthzEhttp://data.lip6.fr/cadene/pretrainedmodels/se_resnet101-7e38fcc6.pthzEhttp://data.lip6.fr/cadene/pretrainedmodels/se_resnet152-d17c99b7.pthzKhttp://data.lip6.fr/cadene/pretrainedmodels/se_resnext50_32x4d-a260b3a4.pthzLhttp://data.lip6.fr/cadene/pretrainedmodels/se_resnext101_32x4d-3b2fe3d8.pth)senet154 se_resnet50 se_resnet101 se_resnet152se_resnext50_32x4dse_resnext101_32x4dcs^eZdZdZ      d,d-fdd Z  d.d/d"d#Zd0d&d'Zd0d(d)Zd1d*d+ZZ S)2ra SENet based on `Squeeze-and-Excitation Networks `_. Adapted from `Cadene Hub 2D version `_. Args: spatial_dims: spatial dimension of the input data. in_channels: channel number of the input data. block: SEBlock class or str. for SENet154: SEBottleneck or 'se_bottleneck' for SE-ResNet models: SEResNetBottleneck or 'se_resnet_bottleneck' for SE-ResNeXt models: SEResNeXtBottleneck or 'se_resnetxt_bottleneck' layers: number of residual blocks for 4 layers of the network (layer1...layer4). groups: number of groups for the 3x3 convolution in each bottleneck block. for SENet154: 64 for SE-ResNet models: 1 for SE-ResNeXt models: 32 reduction: reduction ratio for Squeeze-and-Excitation modules. for all models: 16 dropout_prob: drop probability for the Dropout layer. if `None` the Dropout layer is not used. for SENet154: 0.2 for SE-ResNet models: None for SE-ResNeXt models: None dropout_dim: determine the dimensions of dropout. Defaults to 1. When dropout_dim = 1, randomly zeroes some of the elements for each channel. When dropout_dim = 2, Randomly zeroes out entire channels (a channel is a 2D feature map). When dropout_dim = 3, Randomly zeroes out entire channels (a channel is a 3D feature map). inplanes: number of input channels for layer1. for SENet154: 128 for SE-ResNet models: 64 for SE-ResNeXt models: 64 downsample_kernel_size: kernel size for downsampling convolutions in layer2, layer3 and layer4. for SENet154: 3 for SE-ResNet models: 1 for SE-ResNeXt models: 1 input_3x3: If `True`, use three 3x3 convolutions instead of a single 7x7 convolution in layer0. - For SENet154: True - For SE-ResNet models: False - For SE-ResNeXt models: False num_classes: number of outputs in `last_linear` layer. for all models: 1000 皙?T spatial_dimsint in_channelsblockCtype[SEBottleneck | SEResNetBottleneck | SEResNeXtBottleneck] | strlayers Sequence[int]groups reduction dropout_prob float | None dropout_diminplanesdownsample_kernel_size input_3x3bool num_classesreturnNonec stt|tr%|dkrt}n|dkrt}n |dkrt}ntd|ttj } t t j |f}t t j |f}ttj|f}ttj|f}t t j|f}| |_||_| rd||dddd d d fd |dd fd| ddfd|dddd d d d fd|dd fd| ddfd|d| dd d d d fd|| d fd| ddfg }nd||| dddd d fd || d fd| ddfg}|d|ddddftt||_|j|d|d||d d|_|j|d|d d||| d|_|j|d|dd||| d|_|j|d|dd||| d|_|d |_|dur||nd|_ t!d|j"| |_#|$D]E}t||r8tj%&t'(|j)q$t||rVtj%*t'(|j)d tj%*t'(|j+dq$t|tj!rhtj%*t'(|j+dq$dS) NZ se_bottleneckZse_resnet_bottleneckZse_resnetxt_bottleneckzUUnknown block '%s', use se_bottleneck, se_resnet_bottleneck or se_resnetxt_bottleneckconv1@r#r!F)r' out_channels kernel_sizestridepaddingbiasbn1) num_featuresrelu1T)inplaceconv2bn2relu2conv3bn3relu3pool)r<r= ceil_moder)planesblocksr,r-r2r")rMrNr=r,r-r2i),super__init__ isinstancestrr r r ValueErrorr RELUr CONVrMAXrBATCHrDROPOUT ADAPTIVEAVGr1r%appendnn Sequentialrlayer0 _make_layerlayer1layer2layer3layer4adaptive_avg_pooldropoutLinear expansion last_linearmodulesinitkaiming_normal_torch as_tensorweight constant_r?)selfr%r'r(r*r,r-r.r0r1r2r3r5 relu_type conv_type pool_type norm_type dropout_type avg_pool_typeZlayer0_modulesm __class__[/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/networks/nets/senet.pyrQ`s                   zSENet.__init__=type[SEBottleneck | SEResNetBottleneck | SEResNeXtBottleneck]rMrNr= nn.Sequentialc Csd}|dks|j||jkr t|j|j||j||dtjdd}g} | ||j|j|||||d||j|_td|D]} | ||j|j|||dq=tj | S)Nr!F)r%r'r;stridesr<actnormr?)r%r1rMr,r-r= downsample)r%r1rMr,r-) r1rgrr%rrXr[ranger\r]) rpr(rMrNr,r-r=r2rr*_numrzrzr{r_sH    zSENet._make_layerx torch.TensorcCs6||}||}||}||}||}|SN)r^r`rarbrcrprrzrzr{featuress     zSENet.featurescCs8||}|jdur||}t|d}||}|S)Nr!)rdrerlflattenrhrrzrzr{logitss     z SENet.logitscCs||}||}|Sr)rrrrzrzr{forwards  z SENet.forward)r r!r"r#Tr$)r%r&r'r&r(r)r*r+r,r&r-r&r.r/r0r&r1r&r2r&r3r4r5r&r6r7)r!r!)r(r|rMr&rNr&r,r&r-r&r=r&r2r&r6r})rr)rrr6r) __name__ __module__ __qualname____doc__rQr_rrr __classcell__rzrzrxr{r2s5} 1 rmodel nn.ModulearchrSprogressr4c st|td}|durtdtd}td}td}td}td}td} t|trFt|d |d d tj |d dd d nt ||dt  D]l} d} | | rct|d| } nP| | rpt|d| } nC| | r| | <t|d| } n.| | r| | <t|d| } n| | rt|d| } n | | rt| d| } | r| | <| =qR|fddD|dS)z: This function is used to load pretrained models. Nzonly 'senet154', 'se_resnet50', 'se_resnet101', 'se_resnet152', 'se_resnext50_32x4d', and se_resnext101_32x4d are supported to load pretrained weights.z%^(layer[1-4]\.\d\.(?:conv)\d\.)(\w*)$z%^(layer[1-4]\.\d\.)(?:bn)(\d\.)(\w*)$z+^(layer[1-4]\.\d\.)(?:se_module.fc1.)(\w*)$z+^(layer[1-4]\.\d\.)(?:se_module.fc2.)(\w*)$z*^(layer[1-4]\.\d\.)(?:downsample.0.)(\w*)$z*^(layer[1-4]\.\d\.)(?:downsample.1.)(\w*)$urlfilename)filepathT) map_location weights_only)rz \1conv.\2z\1conv\2adn.N.\3z\1se_layer.fc.0.\2z\1se_layer.fc.2.\2z\1project.conv.\2z\1project.adn.N.\2cs2i|]\}}|vr|j|jkr||qSrz)shape).0kv model_dict state_dictrzr{ Is,z$_load_state_dict..)rrrTrecompilerRdictrrlloadrlistkeysmatchsubsqueezeritemsupdateload_state_dict) rrr model_urlZ pattern_convZ pattern_bnZ pattern_seZ pattern_se2Zpattern_down_convZpattern_down_bnkeynew_keyrzrr{_load_state_dictsR                  rcs.eZdZdZ     ddfdd ZZS)rzlSENet154 based on `Squeeze-and-Excitation Networks` with optional pretrained support when spatial_dims is 2.r#$r#r9FTr*r+r,r&r- pretrainedr4rr6r7c s4tjdt|||d||rt|d|dSdS)N)r(r*r,r-rrz)rPrQr r)rpr*r,r-rrkwargsrxrzr{rQSs zSENet154.__init__)rr9rFT) r*r+r,r&r-r&rr4rr4r6r7rrrrrQrrzrzrxr{rPsrcs6eZdZdZ         ddfdd ZZS)rznSEResNet50 based on `Squeeze-and-Excitation Networks` with optional pretrained support when spatial_dims is 2.r#r#r!rNr9FTr*r+r,r&r-r.r/r1r2r3r4rrr6r7c s<tjdt|||||||d| |rt|d| dSdS)N)r(r*r,r-r.r1r2r3rrzrPrQr r rpr*r,r-r.r1r2r3rrrrxrzr{rQe   zSEResNet50.__init__) rr!rNr9r!FFTr*r+r,r&r-r&r.r/r1r&r2r&r3r4rr4rr4r6r7rrzrzrxr{rbsrc4eZdZdZ        ddfdd ZZS)rzy SEResNet101 based on `Squeeze-and-Excitation Networks` with optional pretrained support when spatial_dims is 2. r#rr#r!rr9FTr*r+r,r&r-r1r2r3r4rrr6r7c :tjdt||||||d| |rt|d|dSdS)Nr(r*r,r-r1r2r3rrzr rpr*r,r-r1r2r3rrrrxrzr{rQ  zSEResNet101.__init__)rr!rr9r!FFTr*r+r,r&r-r&r1r&r2r&r3r4rr4rr4r6r7rrzrzrxr{rrcr)rzy SEResNet152 based on `Squeeze-and-Excitation Networks` with optional pretrained support when spatial_dims is 2. rr!rr9FTr*r+r,r&r-r1r2r3r4rrr6r7c r)Nrrrzrrrxrzr{rQrzSEResNet152.__init__)rr!rr9r!FFTrrrzrzrxr{rrrc6eZdZdZ         ddfdd ZZS) SEResNext50zy SEResNext50 based on `Squeeze-and-Excitation Networks` with optional pretrained support when spatial_dims is 2. r rNr9r!FTr*r+r,r&r-r.r/r1r2r3r4rrr6r7c <tjdt|||||||d| |rt|d| dSdS)Nr(r*r,r.r-r1r2r3rrzrPrQr rrrxrzr{rQrzSEResNext50.__init__) rrrNr9r!FFTrrrzrzrxr{rrcr)rzz SEResNext101 based on `Squeeze-and-Excitation Networks` with optional pretrained support when spatial_dims is 2. rrrNr9r!FTr*r+r,r&r-r.r/r1r2r3r4rrr6r7c r)Nrrrzrrrxrzr{rQrzSEResNext101.__init__) rrrNr9r!FFTrrrzrzrxr{rrr)rrrrSrr4)? __future__rr collectionsrcollections.abcrtypingrrltorch.nnr\Z torch.hubrmonai.apps.utilsr"monai.networks.blocks.convolutionsrZ,monai.networks.blocks.squeeze_and_excitationr r r monai.networks.layers.factoriesr r rrrmonai.utils.moduler__all__rModulerrrrrrrrSEnetSenetSEnet154Senet154r SEresnet50 Seresnet50 seresnet50 SEresnet101 Seresnet101 seresnet101 SEresnet152 Seresnet152 seresnet152r SEresnext50 Seresnext50 seresnext50 SEResNeXt101 SEresnext101 Seresnext101 seresnext101rzrzrzr{sJ            l3   ""