Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

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

2E$�c�j�@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__��0/data/arXiv/demo/Demo/xdecoder/backbone/focal.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 �BZnH�nW�C�q�ctx�gatesZctx_all�lZ44ctx_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?cs�t���||_||_||_||_|	|_||_||�|_t	||j|j||61|d�|_62|dkrbt|�nt�
�|_||�|_t||�}t||||d�|_d|_d|_d|_d|_|jr�tj|
t�|�dd�|_tj|
t�|�dd�|_dS)N)r2r1r:r4r5r)rrrr��?T)�
requires_grad)rrr0�	mlp_ratior2r1rB�use_layerscale�norm1r(�63modulationrr�Identity�	drop_path�norm2�intr�mlp�H�W�gamma_1�gamma_2�	ParameterrI�ones)rr0rZrr_r�64norm_layerr1r2rBr4r5r[Zlayerscale_value�mlp_hidden_dimrrrr�s66566�67zFocalModulationBlock.__init__c	Cs�|j\}}}|j|j}}|||ks.td��|}|jsB|�|�}|�||||�}|�|��||||�}|jrz|�|�}||�|j	|�}|jr�||�|j68|�|�|���}n ||�|j69|�|�|���}|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 size)
rFrcrd�AssertionErrorrBr\�viewr]r_rerfr`rb)rr rL�LrNrcrd�shortcutrrrr!�s7576" zFocalModulationBlock.forward)77r"r#r$r%rr&r=rr!r'rrrrrUvs
�"rUc
sFeZdZdZdddejdddddddddf
�fdd	�	Zd78d�Z�ZS)�79BasicLayeraj A basic focal modulation layer for one stage.80 81    Args:82        dim (int): Number of feature channels83        depth (int): Depths of this stage.84        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.85        drop (float, optional): Dropout rate. Default: 0.086        drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.087        norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm88        downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None89        focal_level (int): Number of focal levels90        focal_window (int): Focal window size at focal level 191        use_conv_embed (bool): Use overlapped convolution for patch embedding or now. Default: False92        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False93    rVrNrWr)Fc
svt���||_||_t�����������	�94fdd�t|�D��|_|dk	rl|d�d�|95�dd�|_nd|_dS)Ncs<g|]4}t���t�t�r"�|n����	�96���d��qS))r0rZrr_r2r1rBr4r5r[ri)rU�97isinstance�list��.0�i�r0rr_r1r2rZrir5r[rBr4rr�98<listcomp>�s
��z'BasicLayer.__init__.<locals>.<listcomp>r)F)�99patch_size�in_chans�	embed_dim�use_conv_embedri�is_stem)	rr�depth�use_checkpointrr;r?�blocks�100downsample)rr0r}rZrr_rir�r2r1r{rBr4r5r[r~rrvrr�s 101"
�102�103	zBasicLayer.__init__c	Cs�|jD].}|||_|_|jr,t�||�}q||�}q|jdk	r�|�dd��|jd|jd||�}|�|�}|�	d��dd�}|dd|dd}}||||||fS||||||fSdS)rkNr	r)r�����)104rrcrdr~�105checkpointr��	transposermrF�flatten)	rr rcrd�blkZ106x_reshaped�x_down�Wh�Wwrrrr!s107108109$110zBasicLayer.forward)	r"r#r$r%rr=rr!r'rrrrrp�s �2rpcs*eZdZdZd�fdd�	Zd	d111�Z�ZS)�112PatchEmbeda� Image to Patch Embedding113 114    Args:115        patch_size (int): Patch token size. Default: 4.116        in_chans (int): Number of input image channels. Default: 3.117        embed_dim (int): Number of linear projection output channels. Default: 96.118        norm_layer (nn.Module, optional): Normalization layer. Default: None119        use_conv_embed (bool): Whether use overlapped convolution for patch embedding. Default: False120        is_stem (bool): Is the stem block or not. 121    �rD�`NFc122s�t���t|�}||_||_||_|r^|r:d}d}d}	nd}d}d}	tj||||	|d�|_ntj||||d�|_|dk	r�||�|_	nd|_	dS)Nr*r)r�rDr	)r,r-r.)r,r-)123rrrrxryrzrr7r9�norm)124rrxryrzrir{r|r,r.r-rrrr+s$125zPatchEmbed.__init__c126Cs�|��\}}}}||jddkrFt�|d|jd||jdf�}||jddkr�t�|ddd|jd||jdf�}|�|�}|jdk	r�|�d�|�d�}}|�d��dd�}|�|�}|�dd��d|j	||�}|S)�Forward function.r	rNr)rDr�)127�sizerx�F�padr9r�r�r�rmrz)rr �_rcrdr�r�rrrr!Bs$(128129130zPatchEmbed.forward)r�rDr�NFFrTrrrrr�sr�cs�eZdZdZddddddddgdd	d131ejddd
ddgdddddgddddgddddddf�fdd�	Zdd�Zddd�Zdgdfdd�Z	dd�Z132d�fdd�	Z�ZS) �FocalNetaS FocalNet backbone.133 134    Args:135        pretrain_img_size (int): Input image size for training the pretrained model,136            used in absolute postion embedding. Default 224.137        patch_size (int | tuple(int)): Patch size. Default: 4.138        in_chans (int): Number of input image channels. Default: 3.139        embed_dim (int): Number of linear projection output channels. Default: 96.140        depths (tuple[int]): Depths of each Swin Transformer stage.141        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.142        drop_rate (float): Dropout rate.143        drop_path_rate (float): Stochastic depth rate. Default: 0.2.144        norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.145        patch_norm (bool): If True, add normalization after patch embedding. Default: True.146        out_indices (Sequence[int]): Output from which stages.147        frozen_stages (int): Stages to be frozen (stop grad and set eval mode).148            -1 means not freezing any parameters.149        focal_levels (Sequence[int]): Number of focal levels at four stages150        focal_windows (Sequence[int]): Focal window sizes at first focal level at four stages151        use_conv_embed (bool): Whether use overlapped convolution for patch embedding152        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.153    i@r�rDr�r)�rVrg�������?Trr	r�rWFcsnt���||_t|�|_�|_|154|_||_||_t	||�|jrD|	nd|dd�|_155tj|d�|_
dd�t�d|t|��D�}t��|_t|j�D]�}tt�d|�|||||t|d|��t|d|d���|	||jdkr�t	nd|||
|||||||d	�}|j�|�q��fd156d�t|j�D�}||_|D](}|	||�}d|��}|�||��q8|��dS)NT)rxryrzrir{r|)�pcSsg|]}|���qSr)�item)rtr rrrrw�sz%FocalNet.__init__.<locals>.<listcomp>rr)r	)r0r}rZrr_rir�r2r1r{rBr4r5r[r~csg|]}t�d|��qS)r))rars�rzrrrw�sr�)rr�pretrain_img_size�len�157num_layersrz�158patch_norm�out_indices�
frozen_stagesr��patch_embedrr�pos_droprI�linspace�sumr;�layersr?rprar@�num_features�159add_module�_freeze_stages)rr�rxryrz�depthsrZ�	drop_rate�drop_path_raterir�r�r��focal_levels�
focal_windowsr{rBr4r5r[r~�dpr�i_layer�layerr��160layer_namerr�rrlsX161162�163&�164zFocalNet.__init__cCs~|jdkr*|j��|j��D]165}d|_q|jdkrz|j��td|jd�D]*}|j|}|��|��D]166}d|_qlqNdS)NrFr)r	)r�r��eval�167parametersrYr�r?r�)r�paramru�mrrrr��s168169170171172zFocalNet._freeze_stagesNcCsTdd�}t|t�r4|�|�t�}t||d|d�n|dkrH|�|�ntd��dS)z�Initialize the weights in backbone.173 174        Args:175            pretrained (str, optional): Path to pre-trained weights.176                Defaults to None.177        cSsrt|tj�rBt|jdd�t|tj�rn|jdk	rntj�|jd�n,t|tj�rntj�|jd�tj�|jd�dS)Ng{�G�z�?)�stdrrX)	rqrrr�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)rq�str�applyZget_root_logger�load_checkpoint�	TypeError)r�178pretrainedr�r�rrr�init_weights�s	179180zFocalNet.init_weightsc	s4|����fdd����D�}t�d|����fdd����D�}t�d|����fdd����D��i}���D�]�\}}|�d�d	|ks�|d	d181ko�d|ko�d|k}	|	rvd
|ks�d|k�r�|���|��k�r�|}182�|}|183jd}|jd}
||
k�rZt�	|j�}|184|dd�dd�|
|d|
|d�|
|d|
|d�f<|}nR||
k�r�|185dd�dd�||
d||
d�||
d||
d�f}|}d|k�s�d|k�r|}186�|}|187j|jk�rt188|189j�dk�r�|190jd}|jd|k�st�|191jd	}|jd	}||k�r�t�	|j�}|192dd|�|dd|�<|193d|d<|194d|d�|d|d||d|d�<|}n||k�rt�nxt195|196j�dk�r|197jd	}|198jd	}|jd	}||k�r199t�	|j�}|200d|�|d|�<|201d|d<|}n||k�rt�|||<qv|j
|dd�dS)Ncsg|]}|�kr|�qSrr�rtrC)�pretrained_dictrrrw�sz)FocalNet.load_weights.<locals>.<listcomp>z=> Missed keys csg|]}|�kr|�qSrrr���202model_dictrrrw�sz=> Unexpected keys cs"i|]\}}|���kr||�qSr)�keys)rtrC�vr�rr�203<dictcomp>�s�z)FocalNet.load_weights.<locals>.<dictcomp>�.r�*�relative_position_index�	attn_maskZpool_layersr<r)zmodulation.fZpre_convr	r�F)r�)�204state_dictr�r��info�itemsrJr�rFrI�zerosr�rl�NotImplementedError�load_state_dict)rr��pretrained_layers�verboseZmissed_dictZunexpected_dict�need_init_state_dictrCr��	need_initZtable_pretrainedZ
table_currentZfsize1Zfsize2Ztable_pretrained_resizedr0�L1�L2r)r�r�r�load_weights�sx205�206���	(207208209D210D2112122132140215216217218219220221zFocalNet.load_weightscCst��}|�|�}|�d�|�d�}}|�d��dd�}|�|�}i}t|j�D]�}|j|}||||�\}}	}222}}}||j	krRt223|d|���}||�}|�d|	|224|j|��
dddd���}||d�|d�<qRt|j	�dk�r|�d|	|225|j|��
dddd���|d<t��}
|S)	r�r)rDr	r�r�rzres{}�res5)�timer�r�r�r�r�r?r�r�r��getattrrmr�rGrH�formatr�)rr �ticr�r��outsrur�rSrcrdri�outZtocrrrr!6s$226227228229&*zFocalNet.forwardcstt|��|�|��dS)z?Convert the model into training mode while keep layers freezed.N)rr��trainr�)r�moderrrr�PszFocalNet.train)N)T)
r"r#r$r%rr=rr�r�r�r!r�r'rrrrr�Ts6230231232233�J234Xr�cs<eZdZ�fdd�Z�fdd�Zdd�Zedd��Z�ZS)	�235D2FocalNetcs�|ddd}|ddd}d}|ddd}|ddd}|ddd}|ddd	}	|ddd236}237tj}|ddd}|ddd}
|ddd
}|dd�dd�}t�j|||||||	|238||||ddd|ddd|ddd|ddd|ddd||ddd|
d�|ddd|_ddddd�|_|jd|jd|jd|jdd�|_dS) N�BACKBONE�FOCAL�PRETRAIN_IMG_SIZE�239PATCH_SIZErD�	EMBED_DIM�DEPTHS�	MLP_RATIO�	DROP_RATE�DROP_PATH_RATE�240PATCH_NORM�USE_CHECKPOINT�OUT_INDICESZSCALING_MODULATORFZFOCAL_LEVELSZ
FOCAL_WINDOWSZUSE_CONV_EMBEDZ241USE_POSTLNZUSE_POSTLN_IN_MODULATIONZUSE_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�rxryrzr�rZr�r�rir�r~r�r5rrrrWsZ���zD2FocalNet.__init__csV|��dkstd|j�d���i}t��|�}|��D]}||jkr6||||<q6|S)z�242        Args:243            x: Tensor of shape (N,C,H,W). H, W must be a multiple of ``self.size_divisibility``.244        Returns:245            dict[str->Tensor]: names and the corresponding features246        r�z:SwinTransformer takes an input of shape (N, C, H, W). Got z	 instead!)r0rlrFrr!r�r�)rr �outputs�yrCrrrr!�s247��248zD2FocalNet.forwardcs�fdd��jD�S)Ncs&i|]}|t�j|�j|d��qS))�channelsr-)rr�r�)rt�name�rrrr��s��z+D2FocalNet.output_shape.<locals>.<dictcomp>)r�r�rr�r�output_shape�s249�zD2FocalNet.output_shapecCsdS)Nr�rr�rrr�size_divisibility�szD2FocalNet.size_divisibility)	r"r#r$rr!r��propertyrr'rrrrr�Vs2505r�c	Cs�t|dd�}|ddddkr�|ddd}t�d|���t�|d��}t�|�d	}W5QRX|�||ddd251�ddg�|d
�|S)N�MODEL��r��LOAD_PRETRAINEDT�252PRETRAINEDz
=> init from �rb�modelr��PRETRAINED_LAYERSr��VERBOSE)	r�r�r�r�openrI�loadr�r�)r��focal�filenamer6�ckptrrr�get_focal_backbone�s(r)&�mathr��numpy�np�loggingrI�torch.nnrZtorch.nn.functional�253functionalr��torch.utils.checkpoint�utilsr��timm.models.layersrrr�detectron2.utils.file_ior�detectron2.modelingrrr�registryr254�	getLoggerr"r��Modulerr(rUrpr�r�r�rrrrr�<module>s.255JOZ5S