U PhBT@s,ddlmZddlZddlmZddlmZddlZddlm Z ddl m m Z ddl mZddlmZmZmZmZddd gZGd d d e jZGd d d e jZGddde jZGddde jZGddde jZGddde jZGddde jZGdd d e jZddZddZ eZ!Z"dS)) annotationsN)Sequence)Union)FCN)ActConvNormPoolAHnetAhnetAHNetcs:eZdZdZd dddddddfd d Zd d ZZS)Bottleneck3x3x1NintzSequence[int] | intznn.Sequential | NoneNone) spatial_dimsinplanesplanesstride downsamplereturnc stttj|f}ttj|f}ttj|f}ttj } |||ddd|_ |||_ |||d| d|d| ddd|_ |||_ |||dddd|_||d|_| dd |_||_||_|d | dd | dd |_dS) NrF) kernel_sizebiasrrrrrrrpaddingrrTinplacerrrr)super__init__rCONVrBATCHr MAXrRELUconv1bn1conv2bn2conv3bn3relurrpool) selfrrrrr conv_type norm_type pool_type relu_type __class__N/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/networks/nets/ahnet.pyr%s,       zBottleneck3x3x1.__init__cCs|}||}||}||}||}||}||}||}||}|jdk r||}||kr| |}||7}||}|SN) r*r+r0r,r-r.r/rsizer1)r2xresidualoutr9r9r:forward@s             zBottleneck3x3x1.forward)rN)__name__ __module__ __qualname__ expansionr%r@ __classcell__r9r9r7r:r s !r cs&eZdZddddfdd ZZS) Projectionr)rnum_input_featuresnum_output_featuresc sptttj|f}ttj|f}ttj}|d|||d|dd|d|||dddddS) Nnormr0TrconvrFrrr) r$r%rr&rr'rr) add_module)r2rrGrHr3r4r6r7r9r:r%[s  zProjection.__init__rArBrCr%rEr9r9r7r:rFYsrFcs,eZdZdddddddfdd ZZS) DenseBlockrfloat)r num_layersrGbn_size growth_rate dropout_probc sHtt|D]0}t|||||||}|d|d|qdS)Nz denselayer%dr)r$r%range Pseudo3DLayerrL) r2rrPrGrQrRrSilayerr7r9r:r%is   zDenseBlock.__init__rMr9r9r7r:rNgsrNcs*eZdZddddddfdd ZZS) UpTransition transposerstrrrGrH upsample_modec stttj|f}ttj|f}ttj}|d|||d|dd|d|||dddd|d krttj |f}|d |||d d ddn(d} |d krd} |d t j d || d dS)NrIr0TrrJrFrKrYupr" trilinearbilinear scale_factormode align_corners r$r%rr&rr'rr)rL CONVTRANSnnUpsample r2rrGrHr\r3r4r6conv_trans_typerdr7r9r:r%|s"  zUpTransition.__init__)rYrMr9r9r7r:rXzsrXcs*eZdZddddddfdd ZZS)FinalrYrrZr[c stttj|f}ttj|f}ttj}|d|||d|dd|d|||d| ddd| dd d |d krttj |f}|d |||d d d dn(d} |dkrd} |d t j d || ddS)NrIr0TrrJrrrFrrYr]r"rKr^rarerir7r9r:r%s6     zFinal.__init__)rYrMr9r9r7r:rksrkcs2eZdZddddddfdd ZddZZS)rUrrO)rrGrRrQrSc stttj|f}ttj|f}ttj}|||_|dd|_ ||||dddd|_ ||||_ |dd|_ ||||d| ddd| ddd|_ |||_|dd|_|||d | ddd | ddd|_|||_|dd|_|||dddd|_||_dS) NTrrFrKrrr)rrr)rrr)r$r%rr&rr'rr)r+relu1r*r-relu2r,r/relu3r.bn4relu4conv4rS) r2rrGrRrQrSr3r4r6r7r9r:r%s>             zPseudo3DLayer.__init__cCs|}||}||}||}||}||}||}||}||}||}||}| |}| |}| |}d|_ |j dkrt j||j |jd}t||gdS)N)ptrainingr)r+rlr*r-rmr,r/rnr.rorprqrSFdropoutrttorchcat)r2r=ZinxZx3x3x1Zx1x1x3Z new_featuresr9r9r:r@s$             zPseudo3DLayer.forwardrArBrCr%r@rEr9r9r7r:rUs'rUcs:eZdZd dddddfdd Zdddd d ZZS) PSPrYrrZ)r psp_block_numin_chr\c sXtt|_ttj|f}ttj|f}t|_ t|_ t |D]j}d|dd|ddf| d}|j |||d|j ||dd| ddd| ddqL||_ ||_||_|jdkrTttj|f} t |D]f}d|dd|ddf| d}d|dd|dd f| d} |j | dd||| dqdS) Nr"rrr#)rrrrrrrrYr)r$r%rg ModuleList up_modulesrr&r r( pool_modulesproject_modulesrTappendrr{r\rf) r2rr{r|r\r3r5rVr<rjpad_sizer7r9r:r%s*     $$  $$z PSP.__init__z torch.Tensor)r=rc Csg}|jdkrHt|j|j|jD]$\}}}||||}||q n^t|j|jD]N\}}|jdd}d}|jdkr~d}tj|||||j|d}||qVt j |dd}|S)NrYr"r^T)r<rcrdrdim) r\ziprrrrshaperu interpolaterwrx) r2r=outputsZproject_moduleZ pool_moduleZ up_moduleoutputZinterpolate_sizerdr9r9r:r@s&    z PSP.forward)rYryr9r9r7r:rzsrzc s^eZdZdZdd d d d d d d d d fdd Zddd d d ddddZddZddZZS)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. rrrrrrrYFTtuplerrZbool)layersr in_channels out_channelsr{r\ pretrainedprogressc sd|_tttj|f} ttj|f} ttj|f} ttj |f} t t j } ttjdf}ttjdf}||_ ||_ | |_| |_| |_| |_||_||_||dkrtd|dkrtd| |dd| dd| dd | dd d |_| d | dd | dd |_| d|_| dd|_|dkrR| d| ddd |_n| d| dddd|_|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|1D]r}t2|| | frP|j3d|j3d|j4}|j5j67dt89d |n&t2|| r|j5j6:d|j;j6<q|rt=d|d!}|>|dS)"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].)rr)r"r"rrFrr!r#Tr)rYnearest)r"r"r")rrrrr}r)rirriirrg@)rr)?rr$r%rr&rfrr'r r(rr) conv2d_type norm2d_typer3r4r6r5rr{AssertionErrorr*pool1bn0r0maxpool _make_layerr layer1layer2layer3layer4rXup0rNdense0up1dense1up2dense2rFtrans1dense3up3dense4rzpsprkfinalmodules isinstancerrweightdatanormal_mathsqrtfill_rzero_r copy_from) r2rrrrr{r\rrr3rjr4r5r6rrZ densegrowthZdensebnZ ndenselayerZnum_init_featuresZnoutres1Znoutres2Znoutres3Znoutres4Z noutdenseZ noutdense1Z noutdense2Z noutdense3Z noutdense4mnZnet2dr7r9r:r%Rs      "           zAHNet.__init__ztype[Bottleneck3x3x1]z nn.Sequential)blockrblocksrrc Csd}|dks|j||jkrt|j|j||jd||dfd|jdd|jdd|fd|jdd|fd|jd|||j}g}|||j|j|||dfd|j|||j|_t d|D]}|||j|j|qtj|S)NrFrKr#) rrDrg Sequentialr3rr5r4rrT)r2rrrrrr_r9r9r:rs0" zAHNet._make_layercCs||}||}||}||}|}||}|}||}||}||}||}| ||}| |} | | |} | | } | | |} || } || |}||}|||}||}|jdkr||}tj||fdd}n|}||S)Nrrr)r*rrr0rrrrrrrrrrrrrrrr{rrwrxr)r2r=Zconv_xZpool_xZfm1Zfm2Zfm3Zfm4sum0d0Zsum1d1Zsum2d2Zsum3d3Zsum4d4rr9r9r:r@s4                z AHNet.forwardc CsBt|jt|j}}|jjddddddd}|d|jddddg|_t |j |j t ddD]}dt |}g}g}t |d |D] } t| |j|jfr|| qt |d |D] } t| |j|jfr|| qt||D]:\} } t| |jr t| | t| |jrt | | qqrdS) Nrrrr"rrrW_modules)nextr* parametersr unsqueezepermuteclonerepeatr copy_bn_paramrrTrZvarsrrrrrr4r3rcopy_conv_param) r2netp2dp3dweightsrVZ layer_numZlayer_2dZlayer_3dm1m2r9r9r:rs&     zAHNet.copy_from)rrrrrrYFT)r) rArBrC__doc__r%rr@rrEr9r9r7r:r /s$"k#cCsDt||D],\}}|jjdddd|jdd<qdS)Nrr)rrrrrZmodule2dZmodule3drrr9r9r:rsrcCs8t||D] \}}|jdd|jdd<qdSr;)rrrrr9r9r:rsr)# __future__rrcollections.abcrtypingrrwtorch.nnrgtorch.nn.functional functionalruZmonai.networks.blocks.fcnrmonai.networks.layers.factoriesrrrr __all__Moduler rrFrNrXrkrUrzr rrr r r9r9r9r: s*      =$C4k