Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

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

2E$�c�{�@s0ddlZddlZddlZddlZddlZddlmZddlmm	Z3ddlmm
Z
ddlmZmZmZddlmZddlmZmZmZddlmZe�e�ZGdd�dej�ZGd	d4�d5ej�ZGdd�dej�Z Gd
d�dej�Z!Gdd�dej�Z"Gdd�dej�Z#Gdd�de#e�Z$edd��Z%dS)�N)�DropPath�	to_2tuple�
trunc_normal_)�PathManager)�BACKBONE_REGISTRY�Backbone�	ShapeSpec�)�register_backbonecs4eZdZdZddejdf�fdd�	Zdd�Z�ZS)�Mlpz Multilayer 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__��3/data/arXiv/demo/Demo/xdecoder/backbone/focal_dw.pyrs6zMlp.__init__cCs6|�|�}|�|�}|�|�}|�|�}|�|�}|Sr
)rrrr)r�xrrr�forward$s7891011zMlp.forward)	�__name__�12__module__�__qualname__�__doc__r�GELUrr!�
__classcell__rrrrrs	rcs*eZdZdZd13�fdd�	Zdd	�Z�ZS)�FocalModulationa� Focal Modulation14 15    Args:16        dim (int): Number of input channels.17        proj_drop (float, optional): Dropout ratio of output. Default: 0.018        focal_level (int): Number of focal levels19        focal_window (int): Focal window size at focal level 120        focal_factor (int, default=2): Step to increase the focal window21        use_postln (bool, default=False): Whether use post-modulation layernorm22    r��Fc	s�t���||_||_||_||_||_||_tj	|d||jddd�|_23tj||dddddd�|_t�
�|_t�	||�|_t�|�|_t��|_|jr�t�|�|_t|j�D]D}	|j|	|j}24|j�t�tj|||25d||26ddd�t�
���q�dS)	Nr)r	T)�biasr)�kernel_size�stride�padding�groupsr+F)r,r-r/r.r+)rr�dim�focal_level�focal_window�focal_factor�use_postln_in_modulation�scaling_modulatorrr�f�Conv2d�hr&r�projr�	proj_drop�27ModuleList�focal_layers�	LayerNorm�ln�range�append�28Sequential)rr0r:r1r2r3�29use_postlnr4r5�kr,rrrr8s430 3132���zFocalModulation.__init__c
Cs*|j\}}}}|�|�}|�dddd���}t�||||jdfd�\}}}d}	t|j�D]2}33|j|34|�}|	||dd�|35|36d�f}	qZ|�	|j37ddd�j38ddd��}|	||dd�|jd�f}	|jr�|	|jd}	||�|	�}|�dddd���}|j
�r|�|�}|�|�}|�|�}|S)zc Forward function.39 40        Args:41            x: input features with shape of (B, H, W, C)42        r�r	r)NT)�keepdim)�shaper6�permute�43contiguous�torch�splitr1r?r<r�meanr5r8r4r>r9r:)
rr �B�nH�nW�C�q�ctx�gates�ctx_all�l�44ctx_global�x_outrrrr!Ys&45 "464748zFocalModulation.forward)rr)r*r)FFF�r"r#r$r%rr!r'rrrrr(,s!r(csFeZdZdZdddejejdddddddf�fdd	�	Zd49d�Z�Z	S)�FocalModulationBlocka+ Focal Modulation Block.50 51    Args:52        dim (int): Number of input channels.53        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.54        drop (float, optional): Dropout rate. Default: 0.055        drop_path (float, optional): Stochastic depth rate. Default: 0.056        act_layer (nn.Module, optional): Activation layer. Default: nn.GELU57        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm58        focal_level (int): number of focal levels59        focal_window (int): focal kernel size at level 160    �@rr)�	Fg-C��6?cst���||_||_||_||_|	|_||_tj	||ddd|d�|_61||�|_t||j|j||62|d�|_
tj	||ddd|d�|_|dkr�t|�nt��|_||�|_t||�}t||||d�|_d|_d|_d|_d|_|j�rtj|
t�|�dd	�|_tj|
t�|�dd	�|_dS)63NrDr	)r,r-r.r/)r2r1r:r4r5r)rrrr��?T)�
requires_grad)rrr0�	mlp_ratior2r1rB�use_layerscalerr7�dw1�norm1r(�64modulation�dw2r�Identity�	drop_path�norm2�intr�mlp�H�W�gamma_1�gamma_2�	ParameterrI�ones)rr0r]rrdr�65norm_layerr1r2rBr4r5r^�layerscale_value�mlp_hidden_dimrrrr�s:6667�68zFocalModulationBlock.__init__c	Csx|j\}}}|j|j}}|||ks.td��|�||||��dddd���}||�|�}|�dddd����|||�}|}|js�|�	|�}|�||||�}|�69|��||||�}||�|j|�}|jr�|�	|�}|�||||��dddd���}||�
|�}|�dddd����|||�}|j�sP||�|j|�|�|���}n$||�|j|�|��}|�|�}|S)�� Forward function.70 71        Args:72            x: Input feature, tensor size (B, H*W, C).73            H, W: Spatial resolution of the input feature.74        zinput feature has wrong sizerrDr	r))rFrhri�AssertionError�viewrGrHr_rBr`rardrjrbrkrgre)rr rL�LrOrhri�shortcutrrrr!�s, 7576 "77zFocalModulationBlock.forward)78r"r#r$r%rr&r=rr!r'rrrrrXvs
�$rXcsHeZdZdZdddejddddddddddf�fdd	�	Zd79d�Z�ZS)�80BasicLayeraj A basic focal modulation layer for one stage.81 82    Args:83        dim (int): Number of feature channels84        depth (int): Depths of this stage.85        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.86        drop (float, optional): Dropout rate. Default: 0.087        drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.088        norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm89        downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None90        focal_level (int): Number of focal levels91        focal_window (int): Focal window size at focal level 192        use_conv_embed (bool): Use overlapped convolution for patch embedding or now. Default: False93        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False94    rYrNrZr)Fc
sxt���||_||_t�����������	�95fdd�t|�D��|_|dk	rn|d�d�|96�d|d�|_nd|_dS)Ncs<g|]4}t���t�t�r"�|n����	�97���d��qS))r0r]rrdr2r1rBr4r5r^rn)rX�98isinstance�list��.0�i�r0rrdr1r2r]rnr5r^rBr4rr�99<listcomp>�s
��z'BasicLayer.__init__.<locals>.<listcomp>r)F)�100patch_size�in_chans�	embed_dim�use_conv_embedrn�is_stem�use_pre_norm)	rr�depth�use_checkpointrr;r?�blocks�101downsample)rr0r�r]rrdrnr�r2r1r�rBr4r5r^r�r�rr|rr�s"102"
�103�104 105zBasicLayer.__init__c	Cs�|jD].}|||_|_|jr,t�||�}q||�}q|jdk	r�|�dd��|jd|jd||�}|�|�}|�	d��dd�}|dd|dd}}||||||fS||||||fSdS)rqNr	r)r�����)106r�rhrir��107checkpointr��	transposersrF�flatten)	rr rhri�blk�108x_reshaped�x_down�Wh�Wwrrrr!s109110111$112zBasicLayer.forward)	r"r#r$r%rr=rr!r'rrrrrv�s"�4rvcs*eZdZdZd�fdd�	Zd	d113�Z�ZS)�114PatchEmbeda� Image to Patch Embedding115 116    Args:117        patch_size (int): Patch token size. Default: 4.118        in_chans (int): Number of input image channels. Default: 3.119        embed_dim (int): Number of linear projection output channels. Default: 96.120        norm_layer (nn.Module, optional): Normalization layer. Default: None121        use_conv_embed (bool): Whether use overlapped convolution for patch embedding. Default: False122        is_stem (bool): Is the stem block or not. 123    �rD�`NFcs�t���t|�}||_||_||_||_|rd|r@d}d}	d}124nd}d}	d}125tj||||126|	d�|_	ntj||||d�|_	|jr�|dk	r�||�|_127q�d|_128n|dk	r�||�|_129nd|_130dS)Nr*rDr�r	r))r,r-r.)r,r-)rrrr~rr�r�rr7r9�norm)rr~rr�rnr�r�r�r,r.r-rrrr|s.131zPatchEmbed.__init__c132Cs2|��\}}}}||jddkrFt�|d|jd||jdf�}||jddkr�t�|ddd|jd||jdf�}|jr�|jdk	r�|�d��dd�}|�|��dd��||||�}|�	|�}nb|�	|�}|jdk	�r.|�d�|�d�}}|�d��dd�}|�|�}|�dd��d|j133||�}|S)�Forward function.r	rNr)rDr�)�sizer~�F�padr�r�r�r�rsr9r�)rr rLrOrhrir�r�rrrr!�s"$(134135136zPatchEmbed.forward)r�rDr�NFFFrWrrrrr�psr�cs�eZdZdZddddddddgdd	d137ejddd
ddgdddddgddddgddddgddddddf�fdd�	Zdd�Zddd�Zdgdfdd�Z	dd�Z138d�fdd�	Z�ZS) �FocalNetaS FocalNet backbone.139 140    Args:141        pretrain_img_size (int): Input image size for training the pretrained model,142            used in absolute postion embedding. Default 224.143        patch_size (int | tuple(int)): Patch size. Default: 4.144        in_chans (int): Number of input image channels. Default: 3.145        embed_dim (int): Number of linear projection output channels. Default: 96.146        depths (tuple[int]): Depths of each Swin Transformer stage.147        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.148        drop_rate (float): Dropout rate.149        drop_path_rate (float): Stochastic depth rate. Default: 0.2.150        norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.151        patch_norm (bool): If True, add normalization after patch embedding. Default: True.152        out_indices (Sequence[int]): Output from which stages.153        frozen_stages (int): Stages to be frozen (stop grad and set eval mode).154            -1 means not freezing any parameters.155        focal_levels (Sequence[int]): Number of focal levels at four stages156        focal_windows (Sequence[int]): Focal window sizes at first focal level at four stages157        use_conv_embed (bool): Whether use overlapped convolution for patch embedding158        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.159    i@r�rDr�r)�rYrg�������?Trr	r�rZFcsxt���||_t|�|_�|_|160|_||_||_t	||�|jrD|	nd|ddd�|_161tj|d�|_
dd�t�d|t|��D�}t��|_t|j�D]�}tt�d|�|||||t|d|��t|d|d	���|	||jd	kr�t	nd|||
|||||||||d162�}|j�|�q��fdd�t|j�D�}||_|jD](}|	||�}d|��}|�||��qB|��dS)
NTF)r~rr�rnr�r�r�)�pcSsg|]}|���qSr)�item)rzr rrrr}�sz%FocalNet.__init__.<locals>.<listcomp>rr)r	)r0r�r]rrdrnr�r2r1r�r�rBr4r5r^r�csg|]}t�d|��qS)r))rfry�r�rrr}sr�)rr�pretrain_img_size�len�163num_layersr��164patch_norm�out_indices�
frozen_stagesr��patch_embedrr�pos_droprI�linspace�sumr;�layersr?rvrfr@�num_features�165add_module�_freeze_stages)rr�r~rr��depthsr]�	drop_rate�drop_path_raternr�r�r��focal_levels�
focal_windowsZ
use_pre_normsr�rBr4r5r^r��dpr�i_layer�layerr��166layer_namerr�rr�s\167168�169&�170171zFocalNet.__init__cCs~|jdkr*|j��|j��D]172}d|_q|jdkrz|j��td|jd�D]*}|j|}|��|��D]173}d|_qlqNdS)NrFr)r	)r�r��eval�174parametersr\r�r?r�)r�paramr{�mrrrr�s175176177178179zFocalNet._freeze_stagesNcCsTdd�}t|t�r4|�|�t�}t||d|d�n|dkrH|�|�ntd��dS)z�Initialize the weights in backbone.180 181        Args:182            pretrained (str, optional): Path to pre-trained weights.183                Defaults to None.184        cSsrt|tj�rBt|jdd�t|tj�rn|jdk	rntj�|jd�n,t|tj�rntj�|jd�tj�|jd�dS)Ng{�G�z�?)�stdrr[)	rwrrr�weightr+�init�	constant_r=)r�rrr�
_init_weights-sz,FocalNet.init_weights.<locals>._init_weightsF)�strict�loggerNz pretrained must be a str or None)rw�str�apply�get_root_logger�load_checkpoint�	TypeError)r�185pretrainedr�r�rrr�init_weights%s	186187zFocalNet.init_weightsc	s4|����fdd����D�}t�d|����fdd����D�}t�d|����fdd����D��i}���D�]�\}}|�d�d	|ks�|d	d188ko�d|ko�d|k}	|	rvd
|ks�d|k�r�|���|��k�r�|}189�|}|190jd}|jd}
||
k�rZt�	|j�}|191|dd�dd�|
|d|
|d�|
|d|
|d�f<|}nR||
k�r�|192dd�dd�||
d||
d�||
d||
d�f}|}d|k�s�d|k�r|}193�|}|194j|jk�rt195|196j�dk�r�|197jd}|jd|k�st�|198jd	}|jd	}||k�r�t�	|j�}|199dd|�|dd|�<|200d|d<|201d|d�|d|d||d|d�<|}n||k�rt�nxt202|203j�dk�r|204jd	}|205jd	}|jd	}||k�r206t�	|j�}|207d|�|d|�<|208d|d<|}n||k�rt�|||<qv|j
|dd�dS)Ncsg|]}|�kr|�qSrr�rzrC)�pretrained_dictrrr}Bsz)FocalNet.load_weights.<locals>.<listcomp>z=> Missed keys csg|]}|�kr|�qSrrr���209model_dictrrr}Dsz=> Unexpected keys cs"i|]\}}|���kr||�qSr)�keys)rzrC�vr�rr�210<dictcomp>Gs�z)FocalNet.load_weights.<locals>.<dictcomp>�.r�*�relative_position_index�	attn_mask�pool_layersr<r)zmodulation.f�pre_convr	r�F)r�)�211state_dictr�r��info�itemsrJr�rFrI�zerosr�rr�NotImplementedError�load_state_dict)rr��pretrained_layers�verbose�missed_dict�unexpected_dict�need_init_state_dictrCr��	need_init�table_pretrained�
table_current�fsize1�fsize2�table_pretrained_resizedr0�L1�L2r)r�r�r�load_weights?sx212�213���	(214215216D217D2182192202210222223224225226227228zFocalNet.load_weightscCst��}|�|�}|�d�|�d�}}|�d��dd�}|�|�}i}t|j�D]�}|j|}||||�\}}	}229}}}||j	krRt230|d|���}||�}|�d|	|231|j|��
dddd���}||d�|d�<qRt|j	�dk�r|�d|	|232|j|��
dddd���|d<t��}
|S)	r�r)rDr	r�r�rzres{}�res5)�timer�r�r�r�r�r?r�r�r��getattrrsr�rGrH�formatr�)rr �ticr�r��outsr{r�rVrhrirn�out�tocrrrr!�s$233234235236&*zFocalNet.forwardcstt|��|�|��dS)z?Convert the model into training mode while keep layers freezed.N)rr��trainr�)r�moderrrr��szFocalNet.train)N)T)
r"r#r$r%rr=rr�r�r�r!r�r'rrrrr��s8237238239240241�M242Xr�cs<eZdZ�fdd�Z�fdd�Zdd�Zedd��Z�ZS)	�243D2FocalNetcs�|ddd}|ddd}d}|ddd}|ddd}|ddd}|ddd	}	|ddd244}245tj}|ddd}|ddd}
|ddd
}|dd�dd�}t�j|||||||	|246||||ddd|ddd|ddd|ddd|ddd||ddd|
d�|ddd|_ddddd�|_|jd|jd|jd|jdd�|_dS) N�BACKBONE�FOCAL�PRETRAIN_IMG_SIZE�247PATCH_SIZErD�	EMBED_DIM�DEPTHS�	MLP_RATIO�	DROP_RATE�DROP_PATH_RATE�248PATCH_NORM�USE_CHECKPOINT�OUT_INDICES�SCALING_MODULATORF�FOCAL_LEVELS�
FOCAL_WINDOWS�USE_CONV_EMBED�249USE_POSTLN�USE_POSTLN_IN_MODULATION�USE_LAYERSCALE)r�r�r�rBr4r5r^r��OUT_FEATURESr���� )�res2�res3�res4r�rr	r))	rr=�getrr�
_out_features�_out_feature_stridesr��_out_feature_channels)r�cfg�input_shaper�r~rr�r�r]r�r�rnr�r�r�r5rrrr�sZ���zD2FocalNet.__init__csV|��dkstd|j�d���i}t��|�}|��D]}||jkr6||||<q6|S)z�250        Args:251            x: Tensor of shape (N,C,H,W). H, W must be a multiple of ``self.size_divisibility``.252        Returns:253            dict[str->Tensor]: names and the corresponding features254        r�z:SwinTransformer takes an input of shape (N, C, H, W). Got z	 instead!)r0rrrFrr!r�r)rr �outputs�yrCrrrr!�s255��256zD2FocalNet.forwardcs�fdd��jD�S)Ncs&i|]}|t�j|�j|d��qS))�channelsr-)rrr)rz�name�rrrr��s��z+D2FocalNet.output_shape.<locals>.<dictcomp>)rrrrr�output_shape�s257�zD2FocalNet.output_shapecCsdS)Nr	rrrrr�size_divisibilityszD2FocalNet.size_divisibility)	r"r#r$rr!r�propertyrr'rrrrr��s2585r�c	Cs�t|dd�}|ddddkr�|ddd}t�d|���t�|d��}t�|�d	}W5QRX|�||ddd259�ddg�|d
�|S)N�MODEL��r��LOAD_PRETRAINEDT�260PRETRAINEDz
=> init from �rb�modelr��PRETRAINED_LAYERSr��VERBOSE)	r�r�r�r�openrI�loadr�r
)r�focal�filenamer6�ckptrrr�get_focal_backbone261s(r()&�mathr��numpy�np�loggingrI�torch.nnrZtorch.nn.functional�262functionalr��torch.utils.checkpoint�utilsr��timm.models.layersrrr�detectron2.utils.file_ior�detectron2.modelingrrr�registryr263�	getLoggerr"r��Modulerr(rXrvr�r�r�r(rrrr�<module>s0264JX#BS