o , iH@sddlmZddlZddlmZddlmZddlmZm Z m Z m Z m Z ddl mZgdZGdd d ejZGd d d ejZGd d d eZGdddejZGdddeZdS)) annotationsN)nn) functional)RegistrationDownSampleBlockRegistrationExtractionBlockRegistrationResidualConvBlockget_conv_blockget_deconv_block) meshgrid_ij)RegUNet AffineHead GlobalNetLocalNetcseZdZdZ       d.d/fdd ZddZddZddZd0d d!Zd1d"d#Z d$d%Z d2d(d)Z d3d*d+Z d,d-Z ZS)4r u Class that implements an adapted UNet. This class also serve as the parent class of LocalNet and GlobalNet Reference: O. Ronneberger, P. Fischer, and T. Brox, “U-net: Convolutional networks for biomedical image segmentation,”, Lecture Notes in Computer Science, 2015, vol. 9351, pp. 234–241. https://arxiv.org/abs/1505.04597 Adapted from: DeepReg (https://github.com/DeepRegNet/DeepReg) kaiming_uniformNTF spatial_dimsint in_channelsnum_channel_initialdepthout_kernel_initializer str | Noneout_activation out_channelsextract_levelstuple[int] | Nonepoolingbool concat_skipencode_kernel_sizesint | list[int]c st|s |f}t||krt|_|_|_|_|_|_ |_ |_ | _ | _ t| tr=| gjd} t| jdkrHt| _fddtjdD_tj _dS)a, Args: spatial_dims: number of spatial dims in_channels: number of input channels num_channel_initial: number of initial channels depth: input is at level 0, bottom is at level depth. out_kernel_initializer: kernel initializer for the last layer out_activation: activation at the last layer out_channels: number of channels for the output extract_levels: list, which levels from net to extract. The maximum level must equal to ``depth`` pooling: for down-sampling, use non-parameterized pooling if true, otherwise use conv concat_skip: when up-sampling, concatenate skipped tensor if true, otherwise use addition encode_kernel_sizes: kernel size for down-sampling csg|] }jd|qS)r.0dself]/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/networks/nets/regunet.py `sz$RegUNet.__init__..N)super__init__maxAssertionErrorrrrrrrrrrr isinstancerlenrrange num_channelsminmin_extract_level build_layers) r(rrrrrrrrrrr __class__r'r*r-,s:     zRegUNet.__init__cCs||dS)N)build_encode_layersbuild_decode_layersr'r)r)r*r6os zRegUNet.build_layerscs`tfddtjD_tfddtjD_jjdjdd_dS)Ncs@g|]}j|dkr jnj|dj|j|dqS)rr!rr kernel_size)build_conv_blockrr3rr$r'r)r*r+vsz/RegUNet.build_encode_layers..csg|] }jj|dqS))channels)build_down_sampling_blockr3r$r'r)r*r+srr) r ModuleListr2r encode_convs encode_poolsbuild_bottom_blockr3 bottom_blockr'r)r'r*r9ss   zRegUNet.build_encode_layersc Cs(tt|j|||dt|j|||dSN)rrrr<)r Sequentialrrrr(rrr<r)r)r*r=szRegUNet.build_conv_blockr>cCst|j||jdS)N)rr>r)rrr)r(r>r)r)r*r?sz!RegUNet.build_down_sampling_blockc Cs4|j|j}tt|j|||dt|j|||dSrH)rrrrIrrrrJr)r)r*rFs zRegUNet.build_bottom_blockcsjtfddtjdjddD_tfddtjdjddD__dS)Ncs*g|]}jj|dj|dqS)r!rB)build_up_sampling_blockr3r$r'r)r*r+sz/RegUNet.build_decode_layers..r!rAcs<g|]}jjrdj|nj|j|ddqS)r#rr;)r=rr3r$r'r)r*r+s) rrCr2rr5decode_deconvs decode_convsbuild_output_block output_blockr'r)r'r*r:s   zRegUNet.build_decode_layersreturn nn.ModulecCst|j||dSNrrr)r rr(rrr)r)r*rKszRegUNet.build_up_sampling_blockcCs t|j|j|j|j|j|jdS)N)rrr3rkernel_initializer activation)rrrr3rrrr'r)r)r*rNszRegUNet.build_output_blockcCs|jdd}g}|}t|j|jD]\}}||}||}||q||}|g} tt|j|jD].\} \} } | |}|j rQt j ||| dgdd}n ||| d}| |}| |q5|j | |d} | S)z Args: x: Tensor in shape (batch, ``in_channels``, insize_1, insize_2, [insize_3]) Returns: Tensor in shape (batch, ``out_channels``, insize_1, insize_2, [insize_3]), with the same spatial size as ``x`` r#Nr!dim) image_size) shapeziprDrEappendrG enumeraterLrMrtorchcatrO)r(xrYskipsencodedZ encode_convZ encode_poolskipdecodedoutsiZ decode_deconvZ decode_convoutr)r)r*forwards$   zRegUNet.forward)rNrNTFr)rrrrrrrrrrrrrrrrrrrrrr )r>rrrrrrrrrrPrQ)rPrQ)__name__ __module__ __qualname____doc__r-r6r9r=r?rFr:rKrNrh __classcell__r)r)r7r*r s&C     r csDeZdZ ddfd d ZedddZdddZdddZZS)r FrrrY list[int] decode_sizer save_thetarc st||_|dkr#||d|d}d}tjgdtjd}n&|dkrB||d|d|d}d}tjgd tjd}ntd |tj||d |_ | ||_ |j j j |j jj |||_t|_d S) aR Args: spatial_dims: number of spatial dimensions image_size: output spatial size decode_size: input spatial size (two or three integers depending on ``spatial_dims``) in_channels: number of input channels save_theta: whether to save the theta matrix estimation r#rr!)r!rrrr!rdtyper ) r!rrrrr!rrrrr!rz/only support 2D/3D operation, got spatial_dims=) in_features out_featuresN)r,r-rr^tensorfloat ValueErrorrLinearfcget_reference_gridgridweightdatazero_biascopy_rrTensortheta) r(rrYrqrrrrwrxZout_initr7r)r*r-s"  zAffineHead.__init__tuple[int] | list[int]rP torch.TensorcCs.dd|D}tjt|dd}|jtjdS)NcSsg|]}td|qS)r)r^arange)r%rXr)r)r*r+z1AffineHead.get_reference_grid..rrWrt)r^stackr torz)rY mesh_pointsrr)r)r*r~szAffineHead.get_reference_gridrc Cs|t|jt|jddg}|jdkr#td||ddd}|S|jdkr6td||ddd}|Std|j) Nr!r#z qij,bpq->bpijrArzqijk,bpq->bpijkzdo not support spatial_dims=)r^r_r ones_likereinsumreshaper{)r(rZ grid_paddedZ grid_warpedr)r)r*affine_transforms   zAffineHead.affine_transformr`list[torch.Tensor]cCsV|d}|jj|jd|_|||jdd}|jr!||_| ||j}|S)Nr)devicerA) rrrr}rrZrrdetachrr)r(r`rYfrrgr)r)r*rh(s zAffineHead.forward)F) rrrYrprqrprrrrr)rYrrPr)rr)r`rrYrprPr) rkrlrmr- staticmethodr~rrhror)r)r7r*r s'   r cs8eZdZdZ      ddfdd ZddZZS)r z Build GlobalNet for image registration. Reference: Hu, Yipeng, et al. "Label-driven weakly-supervised learning for multimodal deformable image registration," https://arxiv.org/abs/1711.01666 rNTFrrYrprrrrrrrrrrrrr rrc s||D]} | ddkrtdddd|q||_fdd|D|_| |_tj|||||||| | d d S) a Args: image_size: output displacement field spatial size spatial_dims: number of spatial dims in_channels: number of input channels num_channel_initial: number of initial channels depth: input is at level 0, bottom is at level depth. out_kernel_initializer: kernel initializer for the last layer out_activation: activation at the last layer pooling: for down-sampling, use non-parameterized pooling if true, otherwise use conv concat_skip: when up-sampling, concatenate skipped tensor if true, otherwise use addition encode_kernel_sizes: kernel size for down-sampling save_theta: whether to save the theta matrix estimation r#rz given depth z3, all input spatial dimension must be divisible by z, got input of size csg|]}|dqSr"r)r%sizerr)r*r+arz&GlobalNet.__init__..) rrrrrrrrrrN)r{rYrqrrr,r-) r(rYrrrrrrrrrrrrr7rr*r-=s2 zGlobalNet.__init__cCs t|j|j|j|jd|jdS)NrA)rrYrqrrr)r rrYrqr3rrr'r)r)r*rNpszGlobalNet.build_output_block)rNTFrF)rYrprrrrrrrrrrrrrrrrrr rrr)rkrlrmrnr-rNror)r)r7r*r 2s3r cs.eZdZ  ddfd d ZdddZZS)AdditiveUpSampleBlocknearestNrrrrmodestr align_corners bool | Nonecs*tt|||d|_||_||_dSrR)r,r-r deconvrr)r(rrrrrr7r)r*r-|s  zAdditiveUpSampleBlock.__init__r`rrPcCspdd|jddD}||}tj|||j|jd}tjtj|j |jdddddddd}||}|S) NcSsg|]}|dqSr"r)rr)r)r*r+sz1AdditiveUpSampleBlock.forward..r#)rrr!) split_sizerXrArW) rZrF interpolaterrr^sumrsplit)r(r` output_sizeZdeconvedresizedrgr)r)r*rhs  ,zAdditiveUpSampleBlock.forward)rN) rrrrrrrrrr)r`rrPr)rkrlrmr-rhror)r)r7r*rzs  rcsHeZdZdZ        d"d#fdd Zd$ddZd%d d!ZZS)&ra Reimplementation of LocalNet, based on: `Weakly-supervised convolutional neural networks for multimodal image registration `_. `Label-driven weakly-supervised learning for multimodal deformable image registration `_. Adapted from: DeepReg (https://github.com/DeepRegNet/DeepReg) rNrTFrrrrrr tuple[int]rrrrrruse_additive_samplingrrrrrc sL| |_| |_| |_tj||||t|||||| dgdgt|d dS)a Args: spatial_dims: number of spatial dims in_channels: number of input channels num_channel_initial: number of initial channels out_kernel_initializer: kernel initializer for the last layer out_activation: activation at the last layer out_channels: number of channels for the output extract_levels: list, which levels from net to extract. The maximum level must equal to ``depth`` pooling: for down-sampling, use non-parameterized pooling if true, otherwise use conv3d use_additive_sampling: whether use additive up-sampling layer for decoding. concat_skip: when up-sampling, concatenate skipped tensor if true, otherwise use addition mode: mode for interpolation when use_additive_sampling, default is "nearest". align_corners: align_corners for interpolation when use_additive_sampling, default is None. r) rrrrrrrrrrrN)use_additive_upsamplingrrr,r-r.) r(rrrrrrrrrrrrr7r)r*r-s  zLocalNet.__init__cCs|j|j}t|j|||dSrH)rrrrrJr)r)r*rFs  zLocalNet.build_bottom_blockrPrQcCs.|jrt|j|||j|jdSt|j||dS)N)rrrrrrS)rrrrrr rTr)r)r*rKsz LocalNet.build_up_sampling_block)rNrTTFrN)rrrrrrrrrrrrrrrrrrrrrrrrrirj)rkrlrmrnr-rFrKror)r)r7r*rs /r) __future__rr^rtorch.nnrrZ#monai.networks.blocks.regunet_blockrrrrr monai.networks.utilsr __all__Moduler r r rrr)r)r)r*s    OFH