o & iW@sddlmZddlZddlZddlmZddlmZddlm Z ddl m Z ddl m Z ddlmZmZmZed \ZZd d gZGd d d ejZGd d d ejZdS)) annotationsN)nn) functional) USE_COMPILED) grid_pull) meshgrid_ij)GridSampleModeGridSamplePadModeoptional_importzmonai._CWarpDVF2DDFcsFeZdZdZejjejjdffdd Z ddd dZ dddZ Z S)r zB Warp an image with given dense displacement field (DDF). Fcsttr2|ddtDvr.t|}|tjkrd}n|tjkr$d}n |tjkr,d}nd}||_n t dt|j |_trj|ddt Dvrft |}|t j krTd}n|t j kr\d}n |t jkrdd}nd}||_nt |j |_d |_||_d S) ac For pytorch native APIs, the possible values are: - mode: ``"nearest"``, ``"bilinear"``, ``"bicubic"``. - padding_mode: ``"zeros"``, ``"border"``, ``"reflection"`` See also: https://pytorch.org/docs/stable/generated/torch.nn.functional.grid_sample.html For MONAI C++/CUDA extensions, the possible values are: - mode: ``"nearest"``, ``"bilinear"``, ``"bicubic"``, 0, 1, ... - padding_mode: ``"zeros"``, ``"border"``, ``"reflection"``, 0, 1, ... See also: :py:class:`monai.networks.layers.grid_pull` - jitter: bool, default=False Define reference grid on non-integer values Reference: B. Likar and F. Pernus. A heirarchical approach to elastic registration based on mutual information. Image and Vision Computing, 19:33-44, 2001. cs|]}|jVqdSNvalue).0interr\/home/dell461/cl/sdc2/last_ska_mid/HISourceFinder-master-l/src/monai/networks/blocks/warp.py <z Warp.__init__..rz=monai.networks.blocks.Warp: Using PyTorch native grid_sample.csr rr)rpadrrrrMrN)super__init__rrBILINEARNEARESTBICUBIC _interp_modewarningswarnrr ZEROSBORDER REFLECTION _padding_moderef_gridjitter)selfmode padding_moder( __class__rrr#s8           z Warp.__init__rddf torch.Tensorr(boolseedintreturncCs|jdur"|jjd|jdkr"|jjdd|jddkr"|jSdd|jddD}tjt|dd}tj|g|jddd}|||_|rptjj|dtj||t |7}Wdn1skwYd|j_ |jS) NrrcSsg|]}td|qS)r)torcharange)rdimrrr esz+Warp.get_reference_grid..)r7)enabledF) r'shaper5stackrtorandomfork_rng manual_seed rand_like requires_grad)r)r.r(r1Z mesh_pointsgridrrrget_reference_grid^s   zWarp.get_reference_gridimagec CsBt|jd}|dvrtd|d|jd|ft|jdd}|j|kr;td|d|jd |d |jd |j||jd |}|dgtt dd|d g}t st |jd dD]\}}|d|fd|d d |d|f<qbtt |d dd}|d|f}t j |||j|jddSt|||jd|jdS)a+ Args: image: Tensor in shape (batch, num_channels, H, W[, D]) ddf: Tensor in the same spatial size as image, in shape (batch, ``spatial_dims``, H, W[, D]) Returns: warped_image in the same shape as image (batch, num_channels, H, W[, D]) r4)r4rzgot unsupported spatial_dims=z, currently support 2 or 3.rNz Given input z-d image shape z, the input DDF shape must be z, Got z instead.)r(r.T)r*r+ align_corners)bound extrapolate interpolation)lenr:NotImplementedErrortuple ValueErrorrCr(permutelistranger enumerateF grid_sampler r&r) r)rDr. spatial_dimsZ ddf_shaperBir7Zindex_orderingrrrforwardqs*   $& z Warp.forward)Fr)r.r/r(r0r1r2r3r/)rDr/r.r/) __name__ __module__ __qualname____doc__rrrr r$rrCrV __classcell__rrr,rr s  ;cs<eZdZdZdejjejjfd fdd Z d d d Z Z S)r z Layer calculates a dense displacement field (DDF) from a dense velocity field (DVF) with scaling and squaring. Adapted from: DeepReg (https://github.com/DeepRegNet/DeepReg) r num_stepsr2cs8t|dkrtd|||_t||d|_dS)Nrz"expecting positive num_steps, got )r*r+)rrrMr\r warp_layer)r)r\r*r+r,rrrs zDVF2DDF.__init__dvfr/r3cCs4|d|j}t|jD] }||j||d}q |S)z Args: dvf: dvf to be transformed, in shape (batch, ``spatial_dims``, H, W[,D]) Returns: a dense displacement field r4)rDr.)r\rPr])r)r^r._rrrrVszDVF2DDF.forward)r\r2)r^r/r3r/) rWrXrYrZrrrr r#rrVr[rrr,rr s   ) __future__rr!r5rtorch.nnrrRZmonai.config.deviceconfigrZ(monai.networks.layers.spatial_transformsrmonai.networks.utilsr monai.utilsrr r _Cr___all__Moduler r rrrrs       u