o

    h                     @   sd   d dl mZ d dlZd dlZd dlmZ ddlmZ d dlm	Z	 G dd deZ
G dd	 d	eZdS )
    N)partial   )
BaseModule)
digit_versionc                       s*   e Zd ZdZd	 fdd	Zdd Z  ZS )
ConditionalPositionEncodingar  The Conditional Position Encoding (CPE) module.

    The CPE is the implementation of 'Conditional Positional Encodings
    for Vision Transformers <https://arxiv.org/abs/2102.10882>'_.

    Args:
       in_channels (int): Number of input channels.
       embed_dims (int): The feature dimension. Default: 768.
       stride (int): Stride of conv layer. Default: 1.
       r   Nc              	      s6   t t| j|d tj||d|dd|d| _|| _d S )Ninit_cfg   r   T)kernel_sizestridepaddingbiasgroups)superr   __init__nnConv2dprojr   )selfin_channels
embed_dimsr   r	   	__class__ ;/root/Awesome-Backbones/configs/common/position_encoding.pyr      s   
z$ConditionalPositionEncoding.__init__c           
      C   sj   |j \}}}|\}}|}|dd||||}	| jdkr%| |	|	 }n| |	}|ddd}|S )Nr      )shape	transposeviewr   r   flatten)
r   xhw_shapeBNCHWZ
feat_tokenZcnn_featr   r   r   forward!   s   

z#ConditionalPositionEncoding.forward)r   r   N)__name__
__module____qualname____doc__r   r(   
__classcell__r   r   r   r   r   	   s    r   c                       s6   e Zd ZdZdddejdf fdd	Zdd	 Z  ZS )
PositionEncodingFouriera  The Position Encoding Fourier (PEF) module.

    The PEF is adopted from EdgeNeXt <https://arxiv.org/abs/2206.10589>'_.
    Args:
        in_channels (int): Number of input channels.
            Default: 32
        embed_dims (int): The feature dimension.
            Default: 768.
        temperature (int): Temperature.
            Default: 10000.
        dtype (torch.dtype): The data type.
            Default: torch.float32.
        init_cfg (dict): The config dict for initializing the module.
            Default: None.
        r   i'  Nc                    s   t t| j|d tj|d |dd| _dtj | _|| _	|| _
|| _tt
jtdk r0t
j}ntt
jdd}t
j|| jd}|d||d |  | _d S )	Nr   r   r   )r   z1.8.0floor)
rounding_modedtype)r   r.   r   r   r   r   mathpiscaler   r   r3   r   torch__version__floor_divider   divarangedim_t)r   r   r   temperaturer3   r	   	floor_divr<   r   r   r   r   ?   s   z PositionEncodingFourier.__init__c              	   C   s  |\}}}t ||| | jjj}| }d}|jd| jd}|jd| jd}	||d d dd d d f |  | j	 }|	|	d d d d dd f |  | j	 }	| j
|j}
|	d d d d d d d f |
 }|d d d d d d d f |
 }t j|d d d d d d dd df  |d d d d d d dd df 
 fddd	}t j|d d d d d d dd df  |d d d d d d dd df 
 fddd	}t j||fd	ddd	dd}
| |
}
|
S )
Ngư>r   r2   r   r      )dimr
   )r7   zerosbooltor   weightdevicecumsumr3   r6   r<   stacksincosr    catpermute)r   Z	bhw_shaper#   r&   r'   maskZnot_maskepsZy_embedZx_embedr<   Zpos_xZpos_yposr   r   r   r(   S   s4   
((  JJ
zPositionEncodingFourier.forward)	r)   r*   r+   r,   r7   float32r   r(   r-   r   r   r   r   r.   .   s    r.   )torch.nnr   r7   r4   	functoolsr   base_moduler   utils.version_utilsr   r   r.   r   r   r   r   <module>   s   %