U Ph $@sddlmZddlmZddlZddlZddlmZddl mm Z ddlm Z ddl mZddlmZmZddlmZmZmZddlmZed d d \ZZd d hZdddhZGdddejZGdddejZdS)) annotations)SequenceN) LayerNorm)build_sincos_position_embedding)Conv trunc_normal_)deprecated_argensure_tuple_repoptional_import)look_up_optionzeinops.layers.torch Rearrange)nameconv perceptronnone learnablesincoscs^eZdZdZedddddddd d d d d ddddd dd fdd ZddZddZZS)PatchEmbeddingBlocka A patch embedding block, based on: "Dosovitskiy et al., An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale " Example:: >>> from monai.networks.blocks import PatchEmbeddingBlock >>> PatchEmbeddingBlock(in_channels=4, img_size=32, patch_size=8, hidden_size=32, num_heads=4, >>> proj_type="conv", pos_embed_type="sincos") pos_embedz1.2z1.4 proj_typezplease use `proj_type` instead.)r sinceremovednew_name msg_suffixrrintSequence[int] | intstrfloatNone) in_channelsimg_size patch_size hidden_size num_headsrrpos_embed_type dropout_rate spatial_dimsreturnc svtd| krdks0ntd| d||dkrRtd|d|dt|t|_t|t|_t|| }t|| }t ||D]6\} } | | krtd|jd kr| | dkrtd qt d d t ||D|_ t |t ||_||jd krttj| f||||d|_n|jd krdd| } dddd| D}dddd | Ddddd | Dd}ddt|D}tt|d|f|t|j||_ttd|j ||_t| |_|jdkrnx|jdkrt|jdd d!d"d#nV|jd$krTg}t ||D]\}}|||q*t ||| |_ntd%|jd&|!|j"dS)'aX Args: in_channels: dimension of input channels. img_size: dimension of input image. patch_size: dimension of patch size. hidden_size: dimension of hidden layer. num_heads: number of attention heads. proj_type: patch embedding layer type. pos_embed_type: position embedding layer type. dropout_rate: fraction of the input units to drop. spatial_dims: number of spatial dimensions. .. deprecated:: 1.4 ``pos_embed`` is deprecated in favor of ``proj_type``. rz dropout_rate z should be between 0 and 1.z hidden size z" should be divisible by num_heads .z+patch_size should be smaller than img_size.rz:patch_size should be divisible by img_size for perceptron.cSsg|]\}}||qSr,).0Zim_dp_dr,r,Y/home/dell461/cl/sdc2/HISourceFinder-master-l/src/monai/networks/blocks/patchembedding.py ^sz0PatchEmbeddingBlock.__init__..rr! out_channels kernel_sizestride))hp1)wp2)dp3Nzb c  css$|]\}}d|d|dVqdS)(r;)Nr,)r-kvr,r,r/ isz/PatchEmbeddingBlock.__init__..zb (cSsg|] }|dqS)rr,r-cr,r,r/r0jsz) (cSsg|] }|dqS)r*r,rAr,r,r/r0jsz c)cSs i|]\}}d|d|qS)pr*r,)r-irCr,r,r/ ks z0PatchEmbeddingBlock.__init__..z -> rrr{Gz?@meanstdabrzpos_embed_type z not supported.)#super__init__ ValueErrorr SUPPORTED_PATCH_EMBEDDING_TYPESrSUPPORTED_POS_EMBEDDING_TYPESr&r zipnpprodZ n_patchesrZ patch_dimrCONVpatch_embeddingsjoin enumeratenn Sequentialr Linear Parametertorchzerosposition_embeddingsDropoutdropoutrappendrapply _init_weights)selfr!r"r#r$r%rrr&r'r(mrCcharsZ from_charsZto_charsZaxes_len grid_sizein_sizeZpa_size __class__r,r/rO-s\            2     zPatchEmbeddingBlock.__init__cCsxt|tjrHt|jdddddt|tjrt|jdk rttj|jdn,t|tjrttj|jdtj|jddS)NrrFrGrHrIrg?) isinstancerZr\rweightbiasinit constant_r)rfrgr,r,r/res  z!PatchEmbeddingBlock._init_weightscCs>||}|jdkr&|ddd}||j}||}|S)Nr)rWrflatten transposer`rb)rfx embeddingsr,r,r/forwards     zPatchEmbeddingBlock.forward)rrrrr) __name__ __module__ __qualname____doc__rrOrery __classcell__r,r,rkr/r s   *Q rcsFeZdZdZdddejdfdddddd d fd d Zd dZZS) PatchEmbeda0 Patch embedding block based on: "Liu et al., Swin Transformer: Hierarchical Vision Transformer using Shifted Windows " https://github.com/microsoft/Swin-Transformer Unlike ViT patch embedding block: (1) input is padded to satisfy window size requirements (2) normalized if specified (3) position embedding is not used. Example:: >>> from monai.networks.blocks import PatchEmbed >>> PatchEmbed(patch_size=2, in_chans=1, embed_dim=48, norm_layer=nn.LayerNorm, spatial_dims=3) rrr*0rrrztype[LayerNorm]r )r#in_chans embed_dim norm_layerr(r)csjt|dkrtdt||}||_||_ttj|f||||d|_|dk r`|||_ nd|_ dS)a Args: patch_size: dimension of patch size. in_chans: dimension of input channels. embed_dim: number of linear projection output channels. norm_layer: normalization layer. spatial_dims: spatial dimension. )rrrz#spatial dimension should be 2 or 3.r1N) rNrOrPr r#rrrVprojnorm)rfr#rrrr(rkr,r/rOs    zPatchEmbed.__init__c Cs |}t|dkr|\}}}}}||jddkrXt|d|jd||jdf}||jddkrt|ddd|jd||jdf}||jddkrt|ddddd|jd||jdf}nt|dkr`|\}}}}||jddkr$t|d|jd||jdf}||jddkr`t|ddd|jd||jdf}||}|jdk r|}|ddd}||}t|dkr|d|d|d}}}|dd d|j |||}n:t|dkr|d|d}}|dd d|j ||}|S)Nrrrr*rrs) sizelenr#Fpadrrrurvviewr) rfrwx_shape_r9r5r7whwwr,r,r/rys6 $(. $(   zPatchEmbed.forward) rzr{r|r}rZrrOryr~r,r,rkr/rs!r) __future__rcollections.abcrnumpyrTr^torch.nnrZtorch.nn.functional functionalrrZ%monai.networks.blocks.pos_embed_utilsrmonai.networks.layersrr monai.utilsrr r monai.utils.moduler r rrQrRModulerrr,r,r,r/ s       s