U PhD*@sddlmZddlZddlmZddlmZddlmZm Z m Z m Z m Z ddl mZdgZdddd d d ZGd d d ejZdddddddddZGdddejZGdddejZGdddejZGdddejZGdddejZdS)) annotationsN) Convolution)ActConvDropoutNorm split_args)deprecated_argVNettuple[str, dict] | strint)actnchancCs2|dkrdd|if}t|\}}t|}|f|S)Nprelunum_parameters)rr)r ract_nameact_argsact_typerM/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/networks/nets/vnet.pyget_acti_layers   rcs2eZdZd dddddfdd Zdd ZZS) LUConvFr r bool) spatial_dimsrr biasc s4tt|||_t|||ddtj|d|_dS)Nr in_channels out_channels kernel_sizer normr)super__init__r act_functionrrBATCH conv_block)selfrrr r __class__rrr""s  zLUConv.__init__cCs||}||}|SN)r%r#r&xoutrrrforward0s  zLUConv.forward)F__name__ __module__ __qualname__r"r- __classcell__rrr'rr srFr)rrdepthr rcCs0g}t|D]}|t||||q tj|Sr))rangeappendrnn Sequential)rrr3r rlayers_rrr _make_nconv6s r:cs4eZdZd ddddddfdd Zdd ZZS) InputTransitionFr r rrrrr rc sht||dkr,td|d|d||_||_||_t|||_t|||ddt j |d|_ dS)NrzAout channels should be divisible by in_channels. Got in_channels=z, out_channels=.rr) r!r" ValueErrorrrrrr#rrr$r%)r&rrrr rr'rrr"?s$   zInputTransition.__init__cCsN||}|j|j}|d|dddgd|jd}|t||}|S)N)r%rrrepeatrr#torchadd)r&r+r,Z repeat_numx16rrrr-Ws   "zInputTransition.forward)Fr.rrr'rr;=sr;c s8eZdZd ddddddddfd d Zd d ZZS)DownTransitionNFr r float | Noner)rrnconvsr dropout_prob dropout_dimrc stttj|f}ttj|f} ttj|f} d|} ||| dd|d|_| | |_ t || |_ t || |_ t || ||||_|dk r| |nd|_dS)Nr@)rstrider)r!r"rCONVrr$rDROPOUT down_convbn1r act_function1 act_function2r:opsdropout) r&rrrHr rIrJr conv_type norm_type dropout_typerr'rrr"as    zDownTransition.__init__cCsP||||}|jdk r,||}n|}||}|t||}|Sr))rPrOrNrSrRrQrBrC)r&r+downr,rrrr-ys   zDownTransition.forward)NrFFr.rrr'rrE_s  rEc s8eZdZd ddddddddfdd Zd d ZZS) UpTransitionN?rFr r tuple[float | None, float])rrrrHr rIrJc stttj|f}ttj|f} ttj|f} |||dddd|_| |d|_ |ddk rp| |dnd|_ | |d|_ t ||d|_ t |||_t|||||_dS)Nr@)rrKrr?)r!r"r CONVTRANSrr$rrMup_convrOrSdropout2rrPrQr:rR) r&rrrrHr rIrJconv_trans_typerUrVr'rrr"s  zUpTransition.__init__cCsj|jdk r||}n|}||}||||}t||fd}||}|t ||}|S)Nr?) rSr^rPrOr]rBcatrRrQrC)r&r+Zskipxr,ZskipxdoZxcatrrrr-s    zUpTransition.forward)rYrFr.rrr'rrXs  rXcs4eZdZd ddddddfdd Zdd ZZS) OutputTransitionFr r rr<c sRtttj|f}t|||_t|||ddtj|d|_ |||dd|_ dS)Nrrr?)r) r!r"rrLrrPrrr$r%conv2)r&rrrr rrTr'rrr"s   zOutputTransition.__init__cCs"||}||}||}|Sr))r%rPrbr*rrrr-s   zOutputTransition.forward)Fr.rrr'rrasrac szeZdZdZedddddedddddd d d d d d ifdddd df dddddddddd fdd ZddZZS)r a V-Net based on `Fully Convolutional Neural Networks for Volumetric Medical Image Segmentation `_. Adapted from `the official Caffe implementation `_. and `another pytorch implementation `_. The model supports 2D or 3D inputs. Args: spatial_dims: spatial dimension of the input data. Defaults to 3. in_channels: number of input channels for the network. Defaults to 1. The value should meet the condition that ``16 % in_channels == 0``. out_channels: number of output channels for the network. Defaults to 1. act: activation type in the network. Defaults to ``("elu", {"inplace": True})``. dropout_prob_down: dropout ratio for DownTransition blocks. Defaults to 0.5. dropout_prob_up: dropout ratio for UpTransition blocks. Defaults to (0.5, 0.5). dropout_dim: determine the dimensions of dropout. Defaults to (0.5, 0.5). - ``dropout_dim = 1``, randomly zeroes some of the elements for each channel. - ``dropout_dim = 2``, Randomly zeroes out entire channels (a channel is a 2D feature map). - ``dropout_dim = 3``, Randomly zeroes out entire channels (a channel is a 3D feature map). bias: whether to have a bias term in convolution blocks. Defaults to False. According to `Performance Tuning Guide `_, if a conv layer is directly followed by a batch norm layer, bias should be False. .. deprecated:: 1.2 ``dropout_prob`` is deprecated in favor of ``dropout_prob_down`` and ``dropout_prob_up``. rIz1.2dropout_prob_downz'please use `dropout_prob_down` instead.)namesincenew_name msg_suffixdropout_prob_upz%please use `dropout_prob_up` instead.rFr?eluinplaceTrZ)rZrZFr r rGr[r) rrrr rIrcrhrJrc st|dkrtdt||d|| d|_t|dd|| d|_t|dd|| d|_t|dd ||| d |_t|d d||| d |_ t |d d d||d |_ t |d d d||d |_ t |d dd||_ t |ddd||_t|d||| d|_dS)N)r@rFz spatial_dims can only be 2 or 3.)rr? r@@rF)rIr)rI)r!r"AssertionErrorr;in_trrE down_tr32 down_tr64 down_tr128 down_tr256rXup_tr256up_tr128up_tr64up_tr32raout_tr) r&rrrr rIrcrhrJrr'rrr"s z VNet.__init__cCsp||}||}||}||}||}|||}|||}|||}|||}| |}|Sr)) rqrrrsrtrurvrwrxryrz)r&r+Zout16Zout32Zout64Zout128Zout256rrrr- s          z VNet.forward)r/r0r1__doc__r r"r-r2rrr'rr s0 ()r)F) __future__rrBtorch.nnr6"monai.networks.blocks.convolutionsrmonai.networks.layers.factoriesrrrrr monai.utilsr __all__rModulerr:r;rErXrar rrrr s    "%'