a
    d7                     @   s   d Z ddlZddlmZmZmZmZ ddlZddlmZ ddl	m
Z
 ddejdfeeeejeej d	d
dZddejdfeeeejeej ejdddZddddejdfee eeeeejeej ejdddZddddddddejdf
ee eej eeeeeeeee  ejeej eej dddZG dd dejZdd Zejddd Zeej dd!d"Zejdd#d$Zd%d& Zdddddddejdf	ee eej eeeeeeee  ejeej d'
d(d)ZG d*d+ d+ejZG d,d- d-ejZdS ).zv Sin-cos, fourier, rotary position embedding modules and functions

Hacked together by / Copyright 2022 Ross Wightman
    N)ListTupleOptionalUnion)nn   )_assertg      l@T)	num_bandsmax_freqlinear_bandsdtypedevicec                 C   sL   |rt jd|d | ||d}n$dt jdt|dd | ||d }|t j S )N      ?   r   r   r   r   )torchlinspacemathlogpi)r	   r
   r   r   r   bands r   e/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/layers/pos_embed_sincos.pypixel_freq_bands   s    $r   g     @r   )r	   temperaturestepr   r   returnc              	   C   s$   d|t jd| |||d|    }|S )Nr   r   r   r   Zarange)r	   r   r   r   r   r   r   r   r   
freq_bands   s     r   @   F)
feat_shapedimr   reverse_coordinterleave_sin_cosr   r   r   c                    s   |d dksJ d|d }t ||d d}|r@| ddd } tt fdd	| D ddd}	|	d|d }
|rd
nd}tjt|
t|
g|dd}|S )a  

    Args:
        feat_shape:
        dim:
        temperature:
        reverse_coord: stack grid order W, H instead of H, W
        interleave_sin_cos: sin, cos, sin, cos stack instead of sin, sin, cos, cos
        dtype:
        device:

    Returns:

       r   zHEmbed dimension must be divisible by 4 for sin-cos 2D position embeddingr   r   r   r   r   Nc                    s   g | ]}t j| d qS r   r   r   .0sr(   r   r   
<listcomp>E       z,build_sincos2d_pos_embed.<locals>.<listcomp>r   r!   )	r   r   stackmeshgridflatten	transpose	unsqueezesincos)r    r!   r   r"   r#   r   r   Zpos_dimr   gridpos2Z	stack_dimZpos_embr   r(   r   build_sincos2d_pos_embed'   s     $r8      )r    r   r	   max_resr   r   include_grid	in_pixelsref_feat_shaper   r   r   c                    s   |du r8|r$t |t|| d}qTt||d d}n du rF|j du rT|j|rn fdd| D }n fdd| D }|durdd t|| |D }tjt|d	d
}|	d	}|| }|
 |  }}|r|||gn||g}|S )a[  

    Args:
        feat_shape: Feature shape for embedding.
        bands: Pre-calculated frequency bands.
        num_bands: Number of frequency bands (determines output dim).
        max_res: Maximum resolution for pixel based freq.
        temperature: Temperature for non-pixel freq.
        linear_bands: Linear band spacing for pixel based freq.
        include_grid: Include the spatial grid in output.
        in_pixels: Output in pixel freq.
        ref_feat_shape: Reference feature shape for resize / fine-tune.
        dtype: Output dtype.
        device: Output device.

    Returns:

    N)r   r   r   r   r%   c              	      s    g | ]}t jd d| dqS )g      r   )Zstepsr   r   )r   r   r)   r(   r   r   r,      r-   z+build_fourier_pos_embed.<locals>.<listcomp>c                    s   g | ]}t j| d qS r'   r   r)   r(   r   r   r,      r-   c                 S   s   g | ]\}}}|| | qS r   r   )r*   xfrr   r   r   r,      r-   r&   r.   )r   floatr   r   r   zipr   r/   r0   r3   r4   r5   )r    r   r	   r:   r   r   r;   r<   r=   r   r   tr6   posZpos_sinZpos_cosoutr   r(   r   build_fourier_pos_embedN   s>    
rF   c                       s.   e Zd Zd
eed fddZdd	 Z  ZS )FourierEmbedr9   r   TF)r:   r	   c                    s<   t    || _|| _|| _|| _| jdt||dd d S )Nr   F
persistent)super__init__r:   r	   concat_gridkeep_spatialregister_bufferr   )selfr:   r	   rL   rM   	__class__r   r   rK      s    
zFourierEmbed.__init__c                 C   s   |j d d \}}|j dd  }t|| j| j|j|jd}tj|dd}|dd	t
|}|fd|jd   }| jrtj||d|dd	ddgdd}n<tj|ddd	d|d|gdd}||| d}|S )
Nr   )r;   r   r   r&   r.   )r&   r   r      )shaperF   r   rL   r   r   r   catr2   r1   lenndimrM   r3   expandZpermutereshapeZnumel)rO   r>   BCr    embZbatch_expandr   r   r   forward   s"    ,*zFourierEmbed.forward)r9   r   TF)__name__
__module____qualname__intrK   r]   __classcell__r   r   rP   r   rG      s       rG   c                 C   s6   t | ddd df  | dd d df gd| jS )N.r   r   r&   )r   r/   rY   rT   r>   r   r   r   rot   s    rd   rc   c                 C   sJ   |j dkr6| |d|  t| |d|   S | | t| |  S )NrS   r   )rW   r3   	expand_asrd   r>   sin_embcos_embr   r   r   apply_rot_embed   s    
,ri   c                    s&   t | tjr| g}  fdd| D S )Nc                    s    g | ]}|  t |  qS r   )rd   )r*   rC   rh   rg   r   r   r,      r-   z(apply_rot_embed_list.<locals>.<listcomp>)
isinstancer   Tensorrf   r   rj   r   apply_rot_embed_list   s    rm   c                 C   sZ   | dd\}}|jdkrF| |d|  t| |d|   S | | t| |  S )Nr   r&   rS   r   )Ztensor_splitrW   r3   re   rd   )r>   r\   rg   rh   r   r   r   apply_rot_embed_cat   s    
,rn   c              	   C   s@   | d| jd dd}|d| ddd|jd }|S )Nr   r&   r   )r3   rX   rT   Zgather)r>   	pos_embedZkeep_indicesr   r   r   apply_keep_indices_nlc   s    "rp   )
r    r   r!   r:   r   r   r<   r=   r   r   c
                 C   sj   t | ||d ||||||	|d
\}
}d}| D ]}||9 }q,|
|ddd}
||ddd}|
|fS )a  

    Args:
        feat_shape: Spatial shape of the target tensor for embedding.
        bands: Optional pre-generated frequency bands
        dim: Output dimension of embedding tensor.
        max_res: Maximum resolution for pixel mode.
        temperature: Temperature (inv freq) for non-pixel mode
        linear_bands: Linearly (instead of log) spaced bands for pixel mode
        in_pixels: Pixel vs language (inv freq) mode.
        dtype: Output dtype.
        device: Output device.

    Returns:

    r$   )	r   r	   r:   r   r   r<   r=   r   r   r   r&   r   )rF   rY   Zrepeat_interleave)r    r   r!   r:   r   r   r<   r=   r   r   rg   rh   Znum_spatial_dimr>   r   r   r   build_rotary_pos_embed   s$    

rq   c                       s\   e Zd ZdZdeeee  eee  d fdd	Zdeee  d
ddZ	dd Z
  ZS )RotaryEmbeddinga   Rotary position embedding

    NOTE: This is my initial attempt at impl rotary embedding for spatial use, it has not
    been well tested, and will likely change. It will be moved to its own file.

    The following impl/resources were referenced for this impl:
    * https://github.com/lucidrains/vit-pytorch/blob/6f3a5fcf0bca1c5ec33a35ef48d97213709df4ba/vit_pytorch/rvt.py
    * https://blog.eleuther.ai/rotary-embeddings/
    r9   '  TFNr   r    r=   c                    s   t    || _|| _|| _|| _|| _|| _|d u r|rRt|d t	||d}nt
|d |dd}t| | jd|dd d | _d | _n@t|||||| jd\}	}
d | _| jd	|	dd | jd
|
dd d S )Nr$   r   r   r   r   r   FrH   r    r!   r:   r   r<   r=   pos_embed_sinpos_embed_cos)rJ   rK   r!   r:   r   r<   r    r=   r   rA   r   printrN   rx   ry   rq   r   )rO   r!   r:   r   r<   r   r    r=   r   Zemb_sinZemb_cosrP   r   r   rK     s\    


zRotaryEmbedding.__init__rT   c                 C   s8   | j d ur(|d usJ t|| j | jdS | j| jfS d S )Nr<   )r   rq   r<   rx   ry   )rO   rT   r   r   r   	get_embedY  s    
zRotaryEmbedding.get_embedc                 C   s$   |  |jdd  \}}t|||S Nr   )r}   rT   ri   )rO   r>   rg   rh   r   r   r   r]   e  s    zRotaryEmbedding.forward)r9   rs   TFNN)Nr^   r_   r`   __doc__boolr   r   ra   rK   r}   r]   rb   r   r   rP   r   rr     s         

>rr   c                       s\   e Zd ZdZdeeee  eee  d fdd	Zdeee  d
ddZ	dd Z
  ZS )RotaryEmbeddingCata   Rotary position embedding w/ concatenatd sin & cos

    The following impl/resources were referenced for this impl:
    * https://github.com/lucidrains/vit-pytorch/blob/6f3a5fcf0bca1c5ec33a35ef48d97213709df4ba/vit_pytorch/rvt.py
    * https://blog.eleuther.ai/rotary-embeddings/
    r9   rs   TFNrt   c           
         s   t    || _|| _|| _|| _|| _|| _|d u r|rRt|d t	||d}nt
|d |dd}t| | jd|dd d | _n4t|||||| jd}	d | _| jd	t|	d
dd d S )Nr$   ru   r   rv   r   FrH   rw   ro   r&   )rJ   rK   r!   r:   r   r<   r    r=   r   rA   r   rz   rN   embedrq   r   r   rU   )
rO   r!   r:   r   r<   r   r    r=   r   embedsrP   r   r   rK   s  sP    


zRotaryEmbeddingCat.__init__r{   c                 C   s@   | j d ur6t|d ud t|| j | jd}t|dS | jS d S )Nzvalid shape neededr|   r&   )r   r   rq   r<   r   rU   ro   )rO   rT   r   r   r   r   r}     s    
zRotaryEmbeddingCat.get_embedc                 C   s   |  |jdd  }t||S r~   )r}   rT   rn   )rO   r>   ro   r   r   r   r]     s    zRotaryEmbeddingCat.forward)r9   rs   TFNN)Nr   r   r   rP   r   r   k  s   
      

8r   ) r   r   typingr   r   r   r   r   r   Ztrace_utilsr   float32ra   rA   r   r   r   r   rl   r   r8   rF   ModulerG   rd   ri   rm   rn   rp   rq   rr   r   r   r   r   r   <module>   s   )
H,
1[