o , iBTã@s*ddlmZddlZddlmZddlmZddlZddlm Z ddl m m Z ddl mZddlmZmZmZmZgd¢ZGdd „d e jƒZGd d „d e jƒZGd d „d e jƒZGdd„de jƒZGdd„de jƒZGdd„de jƒZGdd„de jƒZGdd„de jƒZdd„Zdd„Z eZ!Z"dS)é)Ú annotationsN)ÚSequence)ÚUnion)ÚFCN)ÚActÚConvÚNormÚPool)ÚAHnetÚAhnetÚAHNetcs0eZdZdZ  dd‡fdd„ Zdd„Z‡ZS)ÚBottleneck3x3x1ééNÚ spatial_dimsÚintÚinplanesÚplanesÚstrideúSequence[int] | intÚ downsampleúnn.Sequential | NoneÚreturnÚNonec sðtƒ ¡ttj|f}ttj|f}ttj|f}ttj } |||ddd|_ ||ƒ|_ |||d| d…|d| d…dd|_ ||ƒ|_ |||dddd|_||dƒ|_| dd |_||_||_|d | d…d | d…d |_dS) NrF)Ú kernel_sizeÚbias©érr©rrr©rrÚpaddingrrT©Úinplace©rré©rr)ÚsuperÚ__init__rÚCONVrÚBATCHr ÚMAXrÚRELUÚconv1Úbn1Úconv2Úbn2Úconv3Úbn3ÚrelurrÚpool) ÚselfrrrrrÚ conv_typeÚ norm_typeÚ pool_typeÚ relu_type©Ú __class__©ú[/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/networks/nets/ahnet.pyr's,     ú  &zBottleneck3x3x1.__init__cCs˜|}| |¡}| |¡}| |¡}| |¡}| |¡}| |¡}| |¡}| |¡}|jdurA| |¡}| ¡| ¡krA|  |¡}||7}| |¡}|S©N) r,r-r2r.r/r0r1rÚsizer3)r4ÚxÚresidualÚoutr;r;r<Úforward@s             zBottleneck3x3x1.forward)rN) rrrrrrrrrrrr)Ú__name__Ú __module__Ú __qualname__Ú expansionr'rBÚ __classcell__r;r;r9r<r s ú!r cseZdZd‡fdd„ Z‡ZS)Ú ProjectionrrÚnum_input_featuresÚnum_output_featuresc sptƒ ¡ttj|f}ttj|f}ttj}| d||ƒ¡| d|dd¡| d|||dddd¡dS) NÚnormr2Tr!ÚconvrF©rrr) r&r'rr(rr)rr+Ú add_module)r4rrIrJr5r6r8r9r;r<r'[s  zProjection.__init__)rrrIrrJr©rCrDrEr'rGr;r;r9r<rHYórHcseZdZd ‡fd d „ Z‡ZS) Ú DenseBlockrrÚ num_layersrIÚbn_sizeÚ growth_rateÚ dropout_probÚfloatc sHtƒ ¡t|ƒD]}t|||||||ƒ}| d|d|¡q dS)Nz denselayer%dr)r&r'ÚrangeÚ Pseudo3DLayerrN) r4rrRrIrSrTrUÚiÚlayerr9r;r<r'is ÿüzDenseBlock.__init__) rrrRrrIrrSrrTrrUrVrOr;r;r9r<rQgrPrQcó"eZdZ d d ‡fdd „ Z‡ZS) Ú UpTransitionÚ transposerrrIrJÚ upsample_modeÚstrc sÌtƒ ¡ttj|f}ttj|f}ttj}| d||ƒ¡| d|dd¡| d|||dddd¡|d krPttj |f}| d |||d d dd¡dSd} |d vrXd} | d t j d || d ¡dS)NrKr2Tr!rLrFrMr]Úupr$©Ú trilinearÚbilinear©Ú scale_factorÚmodeÚ align_corners© r&r'rr(rr)rr+rNÚ CONVTRANSÚnnÚUpsample© r4rrIrJr^r5r6r8Úconv_trans_typergr9r;r<r'|s  ÿzUpTransition.__init__©r]©rrrIrrJrr^r_rOr;r;r9r<r\zóÿr\cr[) ÚFinalr]rrrIrJr^r_c sâtƒ ¡ttj|f}ttj|f}ttj}| d||ƒ¡| d|dd¡| d|||d| d…dd| d…d d ¡|d kr[ttj |f}| d |||d d d d¡dSd} |dvrcd} | d t j d || d¡dS)NrKr2Tr!rLrrrFrr]r`r$rMrardrhrlr9r;r<r'–s4    úþ ÿzFinal.__init__rnrorOr;r;r9r<rq”rprqcs&eZdZd ‡fdd „ Zd d „Z‡ZS) rXrrrIrTrSrUrVc stƒ ¡ttj|f}ttj|f}ttj}||ƒ|_|dd|_ ||||dddd|_ |||ƒ|_ |dd|_ ||||d| d…dd| d…dd|_ ||ƒ|_|dd|_|||d | d…dd | d…dd|_||ƒ|_|dd|_|||dddd|_||_dS) NTr!rFrMrrr)rrr)rrr)r&r'rr(rr)rr+r-Úrelu1r,r/Úrelu2r.r1Úrelu3r0Úbn4Úrelu4Úconv4rU) r4rrIrTrSrUr5r6r8r9r;r<r'ºs>       ú   ú  zPseudo3DLayer.__init__cCs¸|}| |¡}| |¡}| |¡}| |¡}| |¡}| |¡}| |¡}| |¡}| |¡}||}|  |¡}|  |¡}|  |¡}d|_ |j dkrTt j||j |jd}t ||gd¡S)Nç)ÚpÚtrainingr)r-rrr,r/rsr.r1rtr0rurvrwrUÚFÚdropoutrzÚtorchÚcat)r4r?ZinxZx3x3x1Zx1x1x3Z new_featuresr;r;r<rBás$             zPseudo3DLayer.forward) rrrIrrTrrSrrUrV©rCrDrEr'rBrGr;r;r9r<rX¸s'rXcs*eZdZdd‡fdd „ Zdd d„Z‡ZS)ÚPSPr]rrÚ psp_block_numÚin_chr^r_c sZtƒ ¡t ¡|_ttj|f}ttj|f}t ¡|_ t ¡|_ t |ƒD]5}d|dd|ddf| d…}|j   |||d¡|j   ||dd| d…dd| d…d¡q&||_ ||_||_|jdkr©ttj|f} t |ƒD]5}d|dd|ddf| d…}d|dd|dd f| d…} |j  | dd||| d¡qudSdS) Nr$rrr%)rrrr©rrr r]r)r&r'rjÚ ModuleListÚ up_modulesrr(r r*Ú pool_modulesÚproject_modulesrWÚappendrrr^ri) r4rrr‚r^r5r7rYr>rmÚpad_sizer9r;r<r'ýs.     $$ÿ  $$ûz PSP.__init__r?ú torch.Tensorrc Cs¸g}|jdkr$t|j|j|jƒD]\}}}||||ƒƒƒ}| |¡qn/t|j|jƒD]'\}}|jdd…}d}|jdvr?d}tj|||ƒƒ||j|d}| |¡q+t j |dd}|S)Nr]r$raT)r>rfrgr©Údim) r^Úzipr‡r†r…rˆÚshaper{Ú interpolater}r~) r4r?ÚoutputsZproject_moduleZ pool_moduleZ up_moduleÚoutputZinterpolate_sizergr;r;r<rBs(  þ  ü z PSP.forwardrn)rrrrr‚rr^r_)r?rŠrrŠrr;r;r9r<r€ûsr€csPeZdZdZ        d$d%‡fdd„ Zd&d'dd„Zd d!„Zd"d#„Z‡ZS)(r a4 AHNet based on `Anisotropic Hybrid Network `_. Adapted from `lsqshr's official code `_. Except from the original network that supports 3D inputs, this implementation also supports 2D inputs. According to the `tests for deconvolutions `_, using ``"transpose"`` rather than linear interpolations is faster. Therefore, this implementation sets ``"transpose"`` as the default upsampling method. To meet the requirements of the structure, the input size for each spatial dimension (except the last one) should be: divisible by 2 ** (psp_block_num + 3) and no less than 32 in ``transpose`` mode, and should be divisible by 32 and no less than 2 ** (psp_block_num + 3) in other upsample modes. In addition, the input size for the last spatial dimension should be divisible by 32, and at least one spatial size should be no less than 64. Args: layers: number of residual blocks for 4 layers of the network (layer1...layer4). Defaults to ``(3, 4, 6, 3)``. spatial_dims: spatial dimension of the input data. Defaults to 3. in_channels: number of input channels for the network. Default to 1. out_channels: number of output channels for the network. Defaults to 1. psp_block_num: the number of pyramid volumetric pooling modules used at the end of the network before the final output layer for extracting multiscale features. The number should be an integer that belongs to [0,4]. Defaults to 4. upsample_mode: [``"transpose"``, ``"bilinear"``, ``"trilinear"``, ``nearest``] The mode of upsampling manipulations. Using the last two modes cannot guarantee the model's reproducibility. Defaults to ``transpose``. - ``"transpose"``, uses transposed convolution layers. - ``"bilinear"``, uses bilinear interpolate. - ``"trilinear"``, uses trilinear interpolate. - ``"nearest"``, uses nearest interpolate. pretrained: whether to load pretrained weights from ResNet50 to initialize convolution layers, default to False. progress: If True, displays a progress bar of the download of pretrained weights to stderr. ©rrérrrrr]FTÚlayersÚtuplerrÚ in_channelsÚ out_channelsrr^r_Ú pretrainedÚboolÚprogressc sžd|_tƒ ¡ttj|f} ttj|f} ttj|f} ttj |f} t t j } ttjdf}ttjdf}||_ ||_ | |_| |_| |_| |_||_||_||dvrYtdƒ‚|dvratdƒ‚| |dd| d…d| d…d | d…d d |_| d | d…d | d…d |_| dƒ|_| dd|_|dvr¨| d| d…dd |_n | d| d…ddd|_|jtd|ddd|_|jtd|ddd|_|jtd|ddd|_|jtd|ddd|_d}d}d}d}d}d}d}d}t ||||ƒ|_!t"|||||dƒ|_#|||}t ||||ƒ|_$t"|||||dƒ|_%|||}t ||||ƒ|_&t"|||||dƒ|_'|||}t(|||ƒ|_)t"|||||dƒ|_*|||}t ||||ƒ|_+t"|||||dƒ|_,|||}t-||||ƒ|_.t/|||||ƒ|_0| 1¡D]:}t2|| | fƒr¨|j3d|j3d|j4}|j5j6 7dt8 9d |¡¡q‚t2|| ƒr»|j5j6 :d¡|j;j6 <¡q‚|rÍt=d|d!}| >|¡dSdS)"Né@r$)r$rz spatial_dims can only be 2 or 3.)rrr$rrz:psp_block_num should be an integer that belongs to [0, 4].)érœr)r$r$rrFrr#r%Tr!)r]Únearest)r$r$r$)rrrrrƒr)ré€éirériirxg@)r˜rš)?rr&r'rr(rirr)r r*rr+Ú conv2d_typeÚ norm2d_typer5r6r8r7rrÚAssertionErrorr,Úpool1Úbn0r2ÚmaxpoolÚ _make_layerr Úlayer1Úlayer2Úlayer3Úlayer4r\Úup0rQÚdense0Úup1Údense1Úup2Údense2rHÚtrans1Údense3Úup3Údense4r€ÚpsprqÚfinalÚmodulesÚ isinstancerr—ÚweightÚdataÚnormal_ÚmathÚsqrtÚfill_rÚzero_rÚ copy_from) r4r”rr–r—rr^r˜ršr5rmr6r7r8r¡r¢Z densegrowthZdensebnZ ndenselayerZnum_init_featuresZnoutres1Znoutres2Znoutres3Znoutres4Z noutdenseZ noutdense1Z noutdense2Z noutdense3Z noutdense4ÚmÚnZnet2dr9r;r<r'Rsš     ú"          € þzAHNet.__init__Úblockútype[Bottleneck3x3x1]rÚblocksrrú nn.Sequentialc Csòd}|dks|j||jkrDt |j|j||jd||dfd|j…dd|jdd|fd|j…dd|fd|j…d| ||j¡¡}g}| ||j|j|||dfd|j…|ƒ¡||j|_t d|ƒD] }| ||j|j|ƒ¡qftj|ŽS)NrFrMr%) rrFrjÚ Sequentialr5rr7r6rˆrW)r4rÄrrÆrrr”Ú_r;r;r<r§½s.û$ÿõ"ÿ  zAHNet._make_layercCs| |¡}| |¡}| |¡}| |¡}|}| |¡}|}| |¡}| |¡}| |¡}| |¡}|  |¡|}|  |¡} |  | ¡|} |  | ¡} |  | ¡|} | | ¡} | | ¡|}| |¡}| |¡|}| |¡}|jdkr| |¡}tj||fdd}n|}| |¡S)Nrrr‹)r,r¤r¥r2r¦r¨r©rªr«r¬r­r®r¯r°r±r²r³r´rµrr¶r}r~r·)r4r?Zconv_xZpool_xZfm1Zfm2Zfm3Zfm4Úsum0Úd0Zsum1Úd1Zsum2Úd2Zsum3Úd3Zsum4Úd4r¶r;r;r<rB×s4                 z AHNet.forwardc Cs<t|j ¡ƒt|j ¡ƒ}}|jjdd ddddd¡ ¡}| d|jddddg¡|_t |j |j ƒt ddƒD]b}dt |ƒ}g}g}t |ƒd | ¡D]} t| |j|jfƒr_| | ¡qOt |ƒd | ¡D]} t| |j|jfƒrz| | ¡qjt||ƒD]\} } t| |jƒrt| | ƒt| |jƒršt | | ƒq€q9dS) Nrr‹rr$rrérZÚ_modules)Únextr,Ú parametersr»Ú unsqueezeÚpermuteÚcloneÚrepeatrŽÚ copy_bn_paramr¥rWr_Úvarsr¸r¹r¢r¡rˆr6r5rÚcopy_conv_param) r4ÚnetÚp2dÚp3dÚweightsrYZ layer_numZlayer_2dZlayer_3dÚm1Úm2r;r;r<rÁús0   € €    €üôzAHNet.copy_from)r’rrrrr]FT)r”r•rrr–rr—rrrr^r_r˜r™ršr™)r) rÄrÅrrrÆrrrrrÇ) rCrDrEÚ__doc__r'r§rBrÁrGr;r;r9r<r /s$÷ k#r cCsDt| ¡| ¡ƒD]\}}|jjdd ¡dd…|jdd…<q dS)Nrr‹)rrÓr»rÔrÖ©Zmodule2dZmodule3drÜrÝr;r;r<rÚs&ÿrÚcCs8t| ¡| ¡ƒD]\}}|jdd…|jdd…<q dSr=)rrÓr»râr;r;r<rØsÿrØ)#Ú __future__rr½Úcollections.abcrÚtypingrr}Útorch.nnrjÚtorch.nn.functionalÚ functionalr{Zmonai.networks.blocks.fcnrÚmonai.networks.layers.factoriesrrrr Ú__all__ÚModuler rÈrHrQr\rqrXr€r rÚrØr r r;r;r;r<Ús,     =$C4k