Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
swin.cpython-38.pyc291 linesDownload Raw Back to __pycache__
1U

2E$�c��@s>ddlZddlZddlZddlmZddlmmZddl	m3mZddlm
Z
mZmZddlmZmZddlmZddlmZe�e�ZGdd�dej�Zd	d4�Zdd�ZGd
d�dej�ZGdd�dej�ZGdd�dej�Z Gdd�dej�Z!Gdd�dej�Z"Gdd�dej�Z#Gdd�de#e�Z$edd��Z%dS)�N)�DropPath�	to_2tuple�
trunc_normal_)�Backbone�	ShapeSpec)�PathManager�)�register_backbonecs4eZdZdZddejdf�fdd�	Zdd�Z�ZS)�MlpzMultilayer perceptron.N�csNt���|p|}|p|}t�||�|_|�|_t�||�|_t�|�|_dS�N)	�super�__init__�nn�Linear�fc1�act�fc2�Dropout�drop)�self�in_features�hidden_features�out_features�	act_layerr��	__class__��//data/arXiv/demo/Demo/xdecoder/backbone/swin.pyrs5zMlp.__init__cCs6|�|�}|�|�}|�|�}|�|�}|�|�}|Sr)rrrr)r�xrrr�forward(s678910zMlp.forward)	�__name__�11__module__�__qualname__�__doc__r�GELUrr �
__classcell__rrrrr12s�r13cCsR|j\}}}}|�||||||||�}|�dddddd����d|||�}|S)z�14    Args:15        x: (B, H, W, C)16        window_size (int): window size17    Returns:18        windows: (num_windows*B, window_size, window_size, C)19    rr���������)�shape�view�permute�20contiguous)r�window_size�B�H�W�C�windowsrrr�window_partition1s$r6cCsbt|jd||||�}|�|||||||d�}|�dddddd����|||d�}|S)z�21    Args:22        windows: (num_windows*B, window_size, window_size, C)23        window_size (int): Window size24        H (int): Height of image25        W (int): Width of image26    Returns:27        x: (B, H, W, C)28    rr+rr'r(r)r*)�intr,r-r.r/)r5r0r2r3r1rrrr�window_reverse?s29$r8cs,eZdZdZd	�fdd�	Zd30dd�Z�ZS)�WindowAttentiona�Window based multi-head self attention (W-MSA) module with relative position bias.31    It supports both of shifted and non-shifted window.32    Args:33        dim (int): Number of input channels.34        window_size (tuple[int]): The height and width of the window.35        num_heads (int): Number of attention heads.36        qkv_bias (bool, optional):  If True, add a learnable bias to query, key, value. Default: True37        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set38        attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.039        proj_drop (float, optional): Dropout ratio of output. Default: 0.040    TNrcs�t���||_||_||_||}|p.|d|_t�t�	d|ddd|dd|��|_41t�|jd�}	t�|jd�}42t�t�
|	|43g��}t�|d�}|dd�dd�df|dd�ddd�f}
|
�ddd���}
|
dd�dd�df|jdd7<|
dd�dd�df|jdd7<|
dd�dd�dfd|jdd9<|
�d�}|�d|�tj||d|d�|_t�|�|_t�||�|_t�|�|_t|j44d	d45�tjdd�|_dS)Ng�r(rrr+�relative_position_indexr'��bias�{�G�z�?��std)�dim)r
rr@r0�	num_heads�scaler�	Parameter�torch�zeros�relative_position_bias_table�arange�stack�meshgrid�flattenr.r/�sum�register_bufferr�qkvr�	attn_drop�proj�	proj_dropr�Softmax�softmax)rr@r0rA�qkv_bias�qk_scalerNrP�head_dimZcoords_hZcoords_w�coordsZcoords_flattenZrelative_coordsr:rrrr\s446&�,((,47zWindowAttention.__init__c
Csl|j\}}}|�|��||d|j||j��ddddd�}|d|d|d}}}	||j}||�dd�}48|j|j�	d��	|j49d|j50d|j51d|j52dd�}|�ddd���}|53|�d�}54|dk	�r&|jd}|55�	||||j||�|�d��d�}56|57�	d|j||�}58|�
|59�}60n61|�
|62�}63|�|64�}65|66|	�dd��|||�}|�|�}|�|�}|S)	z�Forward function.67        Args:68            x: input features with shape of (num_windows*B, N, C)69            mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None70        r'r(rrr)�����r+N)r,rM�reshaperAr.rB�	transposerFr:r-r0r/�	unsqueezerRrNrOrP)
rr�mask�B_�Nr4rM�q�k�v�attnZrelative_position_biasZnWrrrr �sT71���7273���7475(76777879zWindowAttention.forward)TNrr)N�r!r"r#r$rr r&rrrrr9Os�,r9c80sBeZdZdZddddddddejejf81�fdd	�	Zd82d�Z�Z	S)�SwinTransformerBlocka[Swin Transformer Block.83    Args:84        dim (int): Number of input channels.85        num_heads (int): Number of attention heads.86        window_size (int): Window size.87        shift_size (int): Shift size for SW-MSA.88        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.89        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True90        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.91        drop (float, optional): Dropout rate. Default: 0.092        attn_drop (float, optional): Attention dropout rate. Default: 0.093        drop_path (float, optional): Stochastic depth rate. Default: 0.094        act_layer (nn.Module, optional): Activation layer. Default: nn.GELU95        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm96    �r�@TNrc
	s�t���||_||_||_||_||_d|jkr@|jksJntd��||�|_t	|t97|j�||||	|d�|_|98dkr�t|99�nt
��|_||�|_t||�}
t||
||d�|_d|_d|_dS)Nrz shift_size must in 0-window_size)r0rArSrTrNrPr)rrrr)r
rr@rAr0�100shift_size�	mlp_ratio�AssertionError�norm1r9rrarr�Identity�	drop_path�norm2r7r101�mlpr2r3)rr@rAr0rfrgrSrTrrNrkr�102norm_layerZmlp_hidden_dimrrrr�s8103"104�105106�zSwinTransformerBlock.__init__c	Cs�|j\}}}|j|j}}|||ks.td��|}|�|�}|�||||�}d}	}107|j||j|j}|j||j|j}t�|dd|	||108|f�}|j\}
}}}
|j	dkr�t109j||j	|j	fdd�}|}n|}d}t||j�}|�d|j|j|�}|j
||d�}|�d|j|j|�}t||j||�}|j	dk�rTt110j||j	|j	fdd�}n|}|dk�sl|dk�r�|dd�d|�d|�dd�f��}|�||||�}||�|�}||�|�|�|���}|S)z�Forward function.111        Args:112            x: Input feature, tensor size (B, H*W, C).113            H, W: Spatial resolution of the input feature.114            mask_matrix: Attention mask for cyclic shift.115        �input feature has wrong sizer)rr()�shifts�dimsNr+)r[)r,r2r3rhrir-r0�F�padrfrD�rollr6rar8r/rkrmrl)rrZmask_matrixr1�Lr4r2r3�shortcutZpad_lZpad_tZpad_rZpad_b�_�Hp�WpZ	shifted_x�	attn_maskZ	x_windowsZattn_windowsrrrr �sJ116117�118�$zSwinTransformerBlock.forward)119r!r"r#r$rr%�	LayerNormrr r&rrrrrc�s�,rccs.eZdZdZejf�fdd�	Zdd�Z�ZS)�PatchMergingz�Patch Merging Layer120    Args:121        dim (int): Number of input channels.122        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm123    cs<t���||_tjd|d|dd�|_|d|�|_dS)Nr)r(Fr;)r
rr@rr�	reduction�norm)rr@rnrrrr<s124zPatchMerging.__init__c125Cs:|j\}}}|||ks td��|�||||�}|ddkpF|ddk}|rlt�|ddd|dd|df�}|dd�ddd�ddd�dd�f}|dd�ddd�ddd�dd�f}	|dd�ddd�ddd�dd�f}126|dd�ddd�ddd�dd�f}t�||	|127|gd�}|�|dd|�}|�|�}|�|�}|S)��Forward function.128        Args:129            x: Input feature, tensor size (B, H*W, C).130            H, W: Spatial resolution of the input feature.131        ror(rrNr+r))	r,rhr-rrrsrD�catr~r})rrr2r3r1rur4Z	pad_input�x0�x1�x2�x3rrrr Bs $$$$132133zPatchMerging.forward�	r!r"r#r$rr{rr r&rrrrr|5sr|c134s@eZdZdZdddddddejddf135�fdd	�	Zd136d�Z�ZS)�137BasicLayeraA basic Swin Transformer layer for one stage.138    Args:139        dim (int): Number of feature channels140        depth (int): Depths of this stage.141        num_heads (int): Number of attention head.142        window_size (int): Local window size. Default: 7.143        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.144        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True145        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.146        drop (float, optional): Dropout rate. Default: 0.0147        attn_drop (float, optional): Attention dropout rate. Default: 0.0148        drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0149        norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm150        downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None151        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.152    rdreTNrFcsxt����	|_�	d|_||_|
|_t�����������	f153dd�t|�D��|_	|dk	rn|��d�|_154nd|_155dS)Nr(csPg|]H}t���	|ddkr dn�	d�����t�t�rB�|n��d��qS)r(r)r@rAr0rfrgrSrTrrNrkrn)rc�156isinstance�list��.0�i�157rNr@rrkrgrnrArTrSr0rr�158<listcomp>�s��z'BasicLayer.__init__.<locals>.<listcomp>)r@rn)r
rr0rf�depth�use_checkpointr�159ModuleList�range�blocks�160downsample)rr@r�rAr0rgrSrTrrNrkrnr�r�rr�rrqs161162��zBasicLayer.__init__c	Cs�tt�||j��|j}tt�||j��|j}tjd||df|jd�}td|j�t|j|j�t|jd�f}td|j�t|j|j�t|jd�f}d}	|D].}163|D]$}|	|dd�|164|dd�f<|	d7}	q�q�t	||j�}|�165d|j|j�}|�d�|�d�}
|
�|
dkt
d���|
dkt
d���|j�}
|jD]6}|||_|_|j�rlt�|||
�}n166|||
�}�qB|jdk	�r�|�|||�}|dd|dd}}||||||fS||||||fSdS)	rr)�devicerNr+r(gY�r)r7�np�ceilr0rDrEr��slicerfr6r-rZ�masked_fill�float�type�dtyper�r2r3r��167checkpointr�)rrr2r3rxryZimg_maskZh_slicesZw_slices�cnt�h�wZmask_windowsrz�blkZx_down�Wh�Wwrrrr �sL�����168zBasicLayer.forwardr�rrrrr�_s�0r�cs*eZdZdZd169�fdd�	Zdd	�Z�ZS)�170PatchEmbedaCImage to Patch Embedding171    Args:172        patch_size (int): Patch token size. Default: 4.173        in_chans (int): Number of input image channels. Default: 3.174        embed_dim (int): Number of linear projection output channels. Default: 96.175        norm_layer (nn.Module, optional): Normalization layer. Default: None176    r)r'�`NcsVt���t|�}||_||_||_tj||||d�|_|dk	rL||�|_	nd|_	dS)N)�kernel_size�stride)177r
rr�178patch_size�in_chans�	embed_dimr�Conv2drOr~)rr�r�r�rnrrrr�s179zPatchEmbed.__init__c180Cs�|��\}}}}||jddkrFt�|d|jd||jdf�}||jddkr�t�|ddd|jd||jdf�}|�|�}|jdk	r�|�d�|�d�}}|�d��dd�}|�|�}|�dd��d|j	||�}|S)�Forward function.rrNr(r'r+)181�sizer�rrrsrOr~rJrYr-r�)rrrwr2r3r�r�rrrr �s$(182183184zPatchEmbed.forward)r)r'r�Nrbrrrrr��sr�cs�eZdZdZddddddddgdddd	gd185ddd
dddejdddddf�fdd�	Zdd�Zddd�Zd
gdfdd�Z	dd�Z186d �fdd�	Z�ZS)!�SwinTransformera�Swin Transformer backbone.187        A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows`  -188          https://arxiv.org/pdf/2103.14030189    Args:190        pretrain_img_size (int): Input image size for training the pretrained model,191            used in absolute postion embedding. Default 224.192        patch_size (int | tuple(int)): Patch size. Default: 4.193        in_chans (int): Number of input image channels. Default: 3.194        embed_dim (int): Number of linear projection output channels. Default: 96.195        depths (tuple[int]): Depths of each Swin Transformer stage.196        num_heads (tuple[int]): Number of attention head of each stage.197        window_size (int): Window size. Default: 7.198        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.199        qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True200        qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.201        drop_rate (float): Dropout rate.202        attn_drop_rate (float): Attention dropout rate. Default: 0.203        drop_path_rate (float): Stochastic depth rate. Default: 0.2.204        norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.205        ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.206        patch_norm (bool): If True, add normalization after patch embedding. Default: True.207        out_indices (Sequence[int]): Output from which stages.208        frozen_stages (int): Stages to be frozen (stop grad and set eval mode).209            -1 means not freezing any parameters.210        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.211    ��r)r'r�r(���rdreTNrg�������?F)rrr(r'r+cs�t���||_t|�|_�|_||_||_||_||_	t212||�|jrJ|ndd�|_|jr�t|�}t|�}|d|d|d|dg}t
�t�d�|d|d��|_t|jdd�t
j|d�|_dd�t�d|
t|��D�}t
��|_t|j�D]~}tt�d	|�|||||||	|213|||t|d|��t|d|d���|||jdk�r^tnd|d214�
}|j�|�q��fdd�t|j�D�}||_|D](}|||�}d|��}|�||��q�|� �dS)
N)r�r�r�rnrrr=r>)�pcSsg|]}|���qSr)�item)r�rrrrr�Rsz,SwinTransformer.__init__.<locals>.<listcomp>r()
r@r�rAr0rgrSrTrrNrkrnr�r�csg|]}t�d|��qS)r()r7r��r�rrr�jsr~)!r
r�pretrain_img_size�len�215num_layersr��ape�216patch_norm�out_indices�
frozen_stagesr��patch_embedrrrCrDrE�absolute_pos_embedrr�pos_drop�linspacerKr��layersr�r�r7r|�append�num_features�217add_module�_freeze_stages)rr�r�r�r��depthsrAr0rgrSrT�	drop_rate�attn_drop_rate�drop_path_raternr�r�r�r�r�Zpatches_resolutionZdprZi_layer�layerr�Z218layer_namerr�rrsj219220����221&�222zSwinTransformer.__init__cCs�|jdkr*|j��|j��D]223}d|_q|jdkrB|jrBd|j_|jdkr�|j��td|jd�D]*}|j	|}|��|��D]224}d|_q�qfdS)NrFrr()225r�r��eval�226parameters�
requires_gradr�r�r�r�r�)r�paramr��mrrrr�us227228229230231zSwinTransformer._freeze_stagescCsdd�}dS)z�Initialize the weights in backbone.232        Args:233            pretrained (str, optional): Path to pre-trained weights.234                Defaults to None.235        cSsrt|tj�rBt|jdd�t|tj�rn|jdk	rntj�|jd�n,t|tj�rntj�|jd�tj�|jd�dS)Nr=r>rg�?)	r�rrr�weightr<�init�	constant_r{)r�rrr�
_init_weights�sz3SwinTransformer.init_weights.<locals>._init_weightsNr)r�236pretrainedr�rrr�init_weights�szSwinTransformer.init_weightsc	sT|����fdd�|��D�}i}|��D�]\}}|�d�d|ksR|ddko`d|ko`d|k}|r*d|k�rB|���|��k�rB|}�|}	|��\}237}|	��\}}
||
kr�t�d	|�d238��n||239|k�rBt�d�|240|f||
f��t|241d�}t|d�}tj	j242j|�d
d��
d
|||�||fdd�}|�
|
|��d
d�}d|k�r8|���|��k�r8|}�|}|��\}}243}|��\}}}||k�r�t�d	|�d244��n�|245|k�r8t�d�d
|246|fd
||f��t|247d�}t|d�}|�d|||�}|�ddd
d�}tj	j248j|||fdd�}|�dddd
��d
d�}|||<q*|j|dd�dS)Ncs"i|]\}}|���kr||�qSr)�keys)r�r_r`�Z249model_dictrr�250<dictcomp>�s�z0SwinTransformer.load_weights.<locals>.<dictcomp>�.r�*r:rzrFzError in loading z	, passingz-=> load_pretrained: resized variant: {} to {}g�?r�bicubic�r��moder�r+r'r(F)�strict)�251state_dict�items�splitr��logger�info�formatr7rDr�252functional�interpolater.r-rXrJ�load_state_dict)rZpretrained_dictZpretrained_layers�verboseZneed_init_state_dictr_r`Z	need_initZ'relative_position_bias_table_pretrainedZ$relative_position_bias_table_current�L1ZnH1�L2ZnH2�S1ZS2Z/relative_position_bias_table_pretrained_resizedZabsolute_pos_embed_pretrainedZabsolute_pos_embed_currentrwZC1ZC2Z%absolute_pos_embed_pretrained_resizedrr�r�load_weights�s|253�254���	 255��� 256257���258zSwinTransformer.load_weightsc
Cs>|�|�}|�d�|�d�}}|jrTtj|j||fdd�}||�d��dd�}n|�d��dd�}|�|�}i}t	|j259�D]�}|j|}||||�\}}	}260}}}||jkr~t
|d|���}||�}|�d|	|261|j|��dddd���}||d	�|d�<q~t|j�dk�r:|�d|	|262|j|��dddd���|d263<|S)r�r(r'r�r�rr~r+rzres{}�res5)r�r�r�rrr�r�rJrYr�r�r�r�r��getattrr-r�r.r/r�r�)
rrr�r�r��outsr�r�Zx_outr2r3rn�outrrrr �s.264�265266267&*zSwinTransformer.forwardcstt|��|�|��dS)z?Convert the model into training mode while keep layers freezed.N)r
r��trainr�)rr�rrrr��szSwinTransformer.train)N)T)
r!r"r#r$rr{rr�r�r�r r�r&rrrrr��s4268269�\270C!r�cs<eZdZ�fdd�Z�fdd�Zdd�Zedd��Z�ZS)	�D2SwinTransformercsvt�j||||||||	|271|||
||||||d�|d|_ddddd�|_|jd|jd	|jd272|jdd�|_dS)N�r��OUT_FEATURESr)��� )�res2�res3�res4r�rrr(r')r
r�
_out_features�_out_feature_stridesr��_out_feature_channels)r�cfgr�r�r�r�r�rAr0rgrSrTr�r�r�rnr�r�r�r�rrrrs>�273��zD2SwinTransformer.__init__csV|��dkstd|j�d���i}t��|�}|��D]}||jkr6||||<q6|S)z�274        Args:275            x: Tensor of shape (N,C,H,W). H, W must be a multiple of ``self.size_divisibility``.276        Returns:277            dict[str->Tensor]: names and the corresponding features278        r)z:SwinTransformer takes an input of shape (N, C, H, W). Got z	 instead!)r@rhr,r
r r�r�)rr�outputs�yr_rrrr *s279��280zD2SwinTransformer.forwardcs.tt�j���t�j�@�}�fdd�|D�S)Ncs&i|]}|t�j|�j|d��qS))�channelsr�)rr�r�)r��name�rrrr�=s��z2D2SwinTransformer.output_shape.<locals>.<dictcomp>)r��setr�r�r�)rZ
feature_namesrrr�output_shape;s281�zD2SwinTransformer.output_shapecCsdS)Nr�rrrrr�size_divisibilityDsz#D2SwinTransformer.size_divisibility)	r!r"r#rr r�propertyrr&rrrrr�s282(	r�cCsH|ddd}|d}|d}d}|d}|d}|d	}|d283}|d}	|d}284|d
}|d}|d}
|d}tj}|d}|d}|d}|�dddddg�}t|||||||||	|285|||
||||||d�}|ddddk�rD|ddd}t�|d��}tj||dd�d}W5QRX|�||�d d!g�|d"�|S)#N�MODEL�BACKBONEZSWINZPRETRAIN_IMG_SIZEZ286PATCH_SIZEr'Z	EMBED_DIMZDEPTHSZ	NUM_HEADSZWINDOW_SIZEZ	MLP_RATIOZQKV_BIASZQK_SCALEZ	DROP_RATEZATTN_DROP_RATEZDROP_PATH_RATEZAPEZ287PATCH_NORMZUSE_CHECKPOINTZOUT_INDICESrrr(r��LOAD_PRETRAINEDT�288PRETRAINED�rbr�)�map_location�modelZPRETRAINED_LAYERSr��VERBOSE)	rr{�getr�r�openrD�loadr�)r�Zswin_cfgr�r�r�r�r�rAr0rgrSrTr�r�r�rnr�r�r�r��swin�filename�f�ckptrrr�get_swin_backboneIs\� r)&�logging�numpyr�rD�torch.nnrZtorch.nn.functionalr�rr�torch.utils.checkpoint�utilsr��timm.models.layersrrr�detectron2.modelingrr�detectron2.utils.file_ior�registryr	�	getLoggerr!r��Moduler289r6r8r9rcr|r�r�r�r�rrrrr�<module>290s2291e*t*H