U PhW@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._CWarpDVF2DDFcsVeZdZdZejjejjdffdd Z dddddd d d Z ddd d dZ Z S)r zB Warp an image with given dense displacement field (DDF). Fcsttrd|ddtDkr\t|}|tjkr8d}n$|tjkrHd}n|tjkrXd}nd}||_nt dt|j |_tr|ddt Dkrt |}|t j krd}n$|t j krd}n|t jkrd}nd}||_n t |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. css|] }|jVqdSNvalue).0interrO/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/networks/blocks/warp.py <sz Warp.__init__..rz=monai.networks.blocks.Warp: Using PyTorch native grid_sample.css|] }|jVqdSr r)rpadrrrrMsN)super__init__rrBILINEARNEARESTBICUBIC _interp_modewarningswarnrr ZEROSBORDER REFLECTION _padding_moderef_gridjitter)selfmode padding_moder& __class__rrr#s8          z Warp.__init__r torch.Tensorboolint)ddfr&seedreturnc Cs|jdk rD|jjd|jdkrD|jjdd|jddkrD|jSdd|jddD}tjt|dd}tj|g|jddd}|||_|rtjj|d tj||t |7}W5QRXd|j_ |jS) NrrcSsg|]}td|qS)r)torcharange)rdimrrr esz+Warp.get_reference_grid..)r5)enabledF) r%shaper3stackrtorandomfork_rng manual_seed rand_like requires_grad)r'r/r&r0Z mesh_pointsgridrrrget_reference_grid^s"  zWarp.get_reference_gridimager/c CsDt|jd}|dkr&td|d|jd|ft|jdd}|j|krvtd|d|jd |d |jd |j||jd |}|dgtt dd|d g}t s.t |jd dD],\}}|d|fd|d d |d|f<qtt |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]) r2)r2rzgot 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)lenr8NotImplementedErrortuple ValueErrorrAr&permutelistranger enumerateF grid_samplerr$r) r'rCr/ spatial_dimsZ ddf_shaper@ir5Zindex_orderingrrrforwardqs.    $& z Warp.forward)Fr) __name__ __module__ __qualname____doc__rrrr r"rrArU __classcell__rrr*rr s;csFeZdZdZdejjejjfddfdd Z dddd 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) rr.) num_stepscs8t|dkr td|||_t||d|_dS)Nrz"expecting positive num_steps, got )r(r))rrrLr[r warp_layer)r'r[r(r)r*rrrs  zDVF2DDF.__init__r,)dvfr1cCs4|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 r2rB)r[rOr\)r'r]r/_rrrrUszDVF2DDF.forward) rVrWrXrYrrrr r!rrUrZrrr*rr s   ) __future__rrr3rtorch.nnrrQZmonai.config.deviceconfigrZ(monai.networks.layers.spatial_transformsrmonai.networks.utilsr monai.utilsrr r _Cr^__all__Moduler r rrrr s       u