
    ^j                        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lmZ ddlmZ 	 	 	 dXd
edededeej$                     fdZ	 	 	 dYd
edededeej$                     dej(                  f
dZdddddej,                  fdee   dededededeej$                     dej.                  dej(                  fdZdee   dee   fdZddddddd	ddd dej,                  fdee   d!eej(                     d
ed"ededed#ed$ed%eee      d&ed'edeej$                     dej.                  deej(                     fd(Z G d) d*ej8                        Zd+ Zd,ej(                  dej(                  fd-Z	 dZd,ej(                  d.ej(                  d/ej(                  d0edej(                  f
d1Z 	 dZd,eej(                     d.ej(                  d/ej(                  d0edeej(                     f
d2Z!	 dZd,ej(                  d3ej(                  d0edej(                  fd4Z"	 dZd,ej(                  d5ej(                  d6ej(                  d7edej(                  f
d8Z#dddddd	ddd dej,                  fdee   d!eej(                     ded"ededed$ed%eee      d&ed'edeej$                     dej.                  fd9Z$ G d: d;ej8                        Z% G d< d=ej8                        Z&	 	 d[dej,                  d>d?ed@edAededBedej(                  fdCZ'ejP                  jR                  e
d dej,                  fdDee   d'edeej$                     dej.                  deej(                  ej(                  f   f
dE              Z*dFej(                  dGej(                  dHej(                  dej(                  fdIZ+ G dJ dKej8                        Z, G dL dMej8                        Z-ejP                  jR                  e
dNd ddOej,                  fdPedQedRed'ed&edej$                  dej.                  dej(                  fdS              Z. G dT dUej8                        Z/	 	 	 d\dVededAedej8                  fdWZ0y)]zv Sin-cos, fourier, rotary position embedding modules and functions

Hacked together by / Copyright 2022 Ross Wightman
    N)ListTupleOptionalUnion)nn   )register_notrace_function)ndgrid)_assertT	num_bandsmax_freqlinear_bandsdevicec                    |r,t        j                  d|dz  | t         j                  |      }nBdt        j                  dt        j                  |d      dz
  | t         j                  |      z  }|t         j
                  z  S )N      ?   dtyper   r   r   )torchlinspacefloat32mathlogpi)r   r   r   r   bandss        g/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/layers/pos_embed_sincos.pypixel_freq_bandsr      sh     sHqL)5==Y_`U^^Atxx!'<q'@)SXS`S`iopp588         @temperaturestepreturnc                     t        j                  d| |t         j                  |      j                  t         j                        | z  }d||z  z  }|S )Nr   r   r   )r   arangeint64tor   )r   r    r!   r   expr   s         r   
freq_bandsr(      sJ     ,,q)TV
L
O
OPUP]P]
^aj
jC+$%ELr   @   F
feat_shapedimreverse_coordinterleave_sin_cosr   c                    |dz  dk(  sJ d       |dz  }t        ||d|      }|r| ddd   } t        j                  t        | D 	cg c]E  }	t        j                  |	|t        j
                        j                  t        j                        G c}	            j                  d      j                  dd      }
|
j                  d      |j                  d      z  }|rd	nd}t        j                  t        j                  |      t        j                  |      g|
      j                  d      }|j                  |      S c c}	w )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   Nr   r   r   r+   r   )r(   r   stackr
   r$   r%   r&   r   flatten	transpose	unsqueezesincos)r*   r+   r    r,   r-   r   r   pos_dimr   sgridpos2	stack_dimpos_embs                 r   build_sincos2d_pos_embedrA   '   s   . 7a<ccc<QhGwKaOE"%
;;v 	QvU[[9<<U]]K   	
99Q? 	 >>" 22D (QIkk599T?EIIdO<)LTTUVWG::E:""s   A
Eseqc                 T    t        |       dk  r| S | d   | d   gt        | dd        z   S )Nr   r   r   )lenlist)rB   s    r   swap_shape_xyrF   P   s4    
3x!|
FCFd3qr7m++r              ijr   max_resinclude_grid	in_pixelsref_feat_shapegrid_offsetgrid_indexingc           
      |   |+|rt        |t        |      ||      }n,t        ||d|      }n||j                  }||j                  }|
dk(  rt        |       } |t        |      }|r6| D cg c]*  }t        j                  dd||t        j                        , }}nS| D cg c]H  }t        j                  ||t        j                        j                  t        j                        |	z   J }}|(t        || |      D cg c]  \  }}}||z  |z   }}}}t        j                  t        j                  ||
	      d
      }|j                  d
      }||z  }|j!                         j                  |      |j#                         j                  |      }}|r|||g}|S ||g}|S c c}w c c}w c c}}}w )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.
        grid_offset: Constant offset to add to grid for non-pixel freq.
        grid_indexing: Indexing mode for meshgrid ('ij' or 'xy')
        dtype: Output dtype.
        device: Output device.

    Returns:

    )r   r   r   r0   xyg      r   )stepsr   r   r2   indexingr1   r3   r4   )r   floatr(   r   r   rF   r   r   r   r$   r%   r&   zipr5   meshgridr8   r9   r:   )r*   r   r   rJ   r    r   rK   rL   rM   rN   rO   r   r   r<   txfrr=   pospos_sinpos_cosouts                          r   build_fourier_pos_embedr`   V   s   F }$g)	E '	E >\\F=KKE":.
%*>:N  
 NN3!F%--P
 
  
 LL6=@@OR]]
 

 !&)!Z&HII71aQUQYII;;u~~a-@bID>>"D
,Cwwy||%|0#''),,U,2KWG&24'
"CJ :A'8JCJ)


 Js   -/F-#AF2F7c                   T     e Zd Z	 	 	 	 	 	 ddedef fdZd	dZd	dZd	dZd Z xZ	S )
FourierEmbedrJ   r   c                     t         |           || _        || _        || _        || _        | j                  dt        j                  |||      d       | j                          y )Nr   r2   F
persistent)
super__init__rJ   r   concat_gridkeep_spatialregister_bufferr   emptyreset_parameters)selfrJ   r   rh   ri   r   r   	__class__s          r   rg   zFourierEmbed.__init__   s`     	"&(Wekk)FRW&Xejk 	r   c                 $    | j                          yz"Initialize parameters and buffers.N_init_buffersrm   s    r   rl   zFourierEmbed.reset_parameters       r   c                 v    | j                   j                  t        | j                  | j                               y.Compute and fill non-persistent buffer values.N)r   copy_r   r   rJ   rs   s    r   rr   zFourierEmbed._init_buffers   s$    

)$..$,,GHr   c                 $    | j                          yz"Initialize non-persistent buffers.Nrq   rs   s    r   init_non_persistent_buffersz(FourierEmbed.init_non_persistent_buffers   rt   r   c           	         |j                   d d \  }}|j                   dd  }t        || j                  | j                  |j                  |j
                        }t        j                  |d      }|j                  dd      j                  t        |            }|fd|j                  dz
  z  z   }| j                  rKt        j                  ||j                  d      j                  |      j                  dd	dd      gd      }|S t        j                  |j                  ddd	d      |j                  d      j                  |      gd      }|j!                  ||j#                         d      }|S )
Nr   )rK   r   r   r1   r3   r1   r   r      )shaper`   r   rh   r   r   r   catr7   r6   rD   ndimri   r8   expandpermutereshapenumel)rm   rY   BCr*   embbatch_expands          r   forwardzFourierEmbed.forward   sC   wwr{1WWQR[
%JJ))''88
 ii$mmB#++C
O<teqvvz22 		1cmmA.55lCKKAqRSUVWX^_`A
  		199Q1a0#--2B2I2I,2WX^`aA		!Z--/4Ar   )rG   r)   TFNNr"   N)
__name__
__module____qualname__intrg   rl   rr   r{   r   __classcell__rn   s   @r   rb   rb      sC         &Ir   rb   c                     t        j                  | ddd df    | dd d df   gd      j                  | j                        S )N.r   r   r1   )r   r5   r   r   )rY   s    r   rotr      sE     ;;319qcc{3R8@@IIr   rY   c                 b    | j                  dd      \  }}t        j                  | |gd      S )Nr   r1   r3   )chunkr   r   )rY   x1x2s      r   rope_rotate_halfr      s1     WWQBWFB99rc2YB''r   sin_embcos_embhalfc                 V    |r| |z  t        |       |z  z   S | |z  t        |       |z  z   S N)r   r   )rY   r   r   r   s       r   apply_rot_embedr      s=      7{-a07:::
 7{SVg---r   c                     t        | t        j                        r| g} |r"| D cg c]  }||z  t        |      |z  z    c}S | D cg c]  }||z  t	        |      |z  z    c}S c c}w c c}w r   )
isinstancer   Tensorr   r   )rY   r   r   r   rX   s        r   apply_rot_embed_listr     su     !U\\"C FGGG.q1G;;GG
 9::1Gc!fw..:: H
 ;s   A$A)r   c                     |j                  dd      \  }}|r| |z  t        |       |z  z   S | |z  t        |       |z  z   S )Nr   r1   )r   r   r   )rY   r   r   r   r   s        r   apply_rot_embed_catr     sR    
 yyB'GW 7{-a07:::
 7{SVg---r   	pos_embedkeep_indicespos_embed_has_batchc                    |rt        |j                  dk\  d       nYt        |j                  dk\  d       | j                  d   fd|j                  z  z   }|j                  d      j	                  |      }|j                  d   fd|j                  dz
  z  z   |j                  d   dfz   }|j                  |      }t        |j                        }d|d	<   |j	                  |      }|j                  d	|      S )
a   Apply keep indices to different ROPE shapes

    Expected pos_embed shapes:
    * [seq_len, pos_embed_dim] --> output [batch_size, seq_len, pos_embed_dim]
    * [num_heads, seq_len, pos_embed_dim] --> output [batch_size, num_heads, seq_len, pos_embed_dim]
    * [depth, num_heads, seq_len, pos_embed_dim] --> output [batch_size, depth, num_heads, seq_len, pos_embed_dim]

    And all of the above with leading batch dimension already present if `pos_embed_has_batch == True`

    r   zIncorrect number of dimensionsr   r   r~   )r   r   r1   r}   )r   r   r   r8   r   viewrE   gather)rY   r   r   r   expand_shape
keep_shapekeep_expands          r   apply_keep_indices_nlcr   ,  s      	!#%EF 		!#%EF
}uy~~'==''*11,?	 $$Q')DINNQ4F,GG<K]K]^_K`bcJddJ$$Z0L y'KKO&&{3LB--r   c                     t        | ||dz  |||||||	|
|      \  }}d}| D ]  }||z  }	 |j                  |d      j                  dd      }|j                  |d      j                  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.
        ref_feat_shape: Reference feature shape for resize / fine-tune.
        grid_offset: Constant offset to add to grid for non-pixel freq.
        grid_indexing: Indexing mode for meshgrid ('ij' or 'xy')
        device: Output device.
        dtype: Output dtype.

    Returns:

    r/   )r   r   rJ   r    r   rL   rM   rN   rO   r   r   r   r1   r   )r`   r   repeat_interleave)r*   r   r+   rJ   r    r   rL   rM   rN   rO   r   r   r   r   num_spatial_dimrY   s                   r   build_rotary_pos_embedr   Q  s    B /(!%#GW O 1ooor2DDQKGooor2DDQKGGr   c                        e Zd ZdZ	 	 	 	 	 	 	 	 	 	 ddedeee      deee      dede	f
 fdZ
dd	Zdd
ZddZdej                  fdee   fdZddZdee   fdZddeee      fdZd Z x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/
    Nr   r*   rM   rN   rO   c                 .   t         |           || _        || _        || _        || _        || _        || _        || _        || _	        |	| _
        |d u| _        |?|dz  f}| j                  dt        j                  ||
|      d       d | _        d | _        nmd | _        d}|D ]  }||z  }	 ||f}| j                  dt        j                  ||
|      d       | j                  dt        j                  ||
|      d       | j%                          y )	Nr/   r   r2   Frd   r   pos_embed_sinpos_embed_cos)rf   rg   r+   rJ   r    r   rL   r*   rM   rN   rO   _use_cached_embedrj   r   rk   r   r   r   rl   rm   r+   rJ   r    rL   r   r*   rM   rN   rO   r   r   bands_shapenum_posr<   	emb_shapern   s                   r   rg   zRotaryEmbedding.__init__  s/    	&("$,&* ",4!7!8+K  %++k&X]*^kp q!%D!%D DJG 1 #I  %++iPV^c2dqv w  %++iPV^c2dqv w 	r   c                 $    | j                          yrp   rq   rs   s    r   rl   z RotaryEmbedding.reset_parameters  rt   r   c                    | j                   s*| j                  j                  | j                                y| j	                  | j
                        \  }}| j                  j                  |       | j                  j                  |       yrv   )r   r   rx   _compute_bands_get_pos_embed_valuesr*   r   r   )rm   emb_sinemb_coss      r   rr   zRotaryEmbedding._init_buffers  sg    %%JJT0023#99$//JGW$$W-$$W-r   c                     | j                   r9t        | j                  dz  t        | j                        | j
                        }n%t        | j                  dz  | j                  d      }|j                  ||      S zCompute frequency bands.r/   )r   r   )r    r!   r2   	rL   r   r+   rU   rJ   r   r(   r    r&   rm   r   r   r   s       r   r   zRotaryEmbedding._compute_bands  k    >>$Adll#!..E A ,,E
 xxvUx33r   c                     t        || j                  | j                  | j                  | j                  | j
                  | j                  | j                  | j                  ||      \  }}||fS )Nr*   r+   rJ   r    r   rL   rM   rN   rO   r   r   )	r   r+   rJ   r    r   rL   rM   rN   rO   )rm   r*   r   r   r   r   s         r   r   z%RotaryEmbedding._get_pos_embed_values  si    1!LL((**nn..((,,
 r   c                 $    | j                          yrz   rq   rs   s    r   r{   z+RotaryEmbedding.init_non_persistent_buffers  rt   r   c                    | j                   }|| j                   k7  rm| j                  J | j                  J | j                  || j                  j                  | j                  j
                        \  | _        | _        || _         y y y Nr2   )r*   r   r   r   r   r   rm   r*   s     r   update_feat_shapez!RotaryEmbedding.update_feat_shape  s    ??&:+H%%111%%111595O5O))00((.. 6P 62D 2
 )DO ,I&r   r   c                    |O| j                   Ct        || j                   | j                  | j                  | j                  | j
                        S | j                  $| j                  | j                  | j                  fS J d       )NrL   rM   rN   rO   zQget_embed() requires pre-computed pos embeds or valid shape w/ pre-computed bands)r   r   rL   rM   rN   rO   r   r   )rm   r   s     r   	get_embedzRotaryEmbedding.get_embed   s    !7)

..#22 ,,"00  +0B0B0N%%t'9'999mmm5r   c                 ^    | j                  |j                  dd        \  }}t        |||      S Nr   )r   r   r   )rm   rY   r   r   s       r   r   zRotaryEmbedding.forward  s.    >>!''!"+6q'733r   
rG   i'  TFNNrH   rI   NNr   NNr   )r   r   r   __doc__boolr   r   r   rU   strrg   rl   rr   r   r   r   r   r{   r   r   r   r   r   s   @r   r   r     s     !&.226!#!%-  -  !c+-  %T#Y/-  -  - ^.4  CGemm  S	   
)DI 
)nxS	2 n 4r   r   c                   Z    e Zd ZdZ	 	 	 	 	 	 	 	 	 	 ddedededededeee      d	eee      d
ede	f fdZ
ddZddZddZdej                  fdee   fdZddZdee   fdZddeee      fdZ	 ddeeeef      dee   deej,                  eej,                     f   fdZd Z x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/
    Nr+   rJ   r    rL   r   r*   rM   rN   rO   c                    t         |           || _        || _        || _        || _        || _        || _        || _        || _	        |	| _
        |d u| _        |8|dz  f}| j                  dt        j                  ||
|      d       d | _        nFd | _        d}|D ]  }||z  }	 ||dz  f}| j                  dt        j                  ||
|      d       | j#                          y )	Nr/   r   r2   Frd   r   r   r   )rf   rg   r+   rJ   r    rL   r   r*   rM   rN   rO   r   rj   r   rk   r   r   rl   r   s                   r   rg   zRotaryEmbeddingCat.__init__  s    	&"($,&* ",4!7!8+K  %++k&X]*^kp q!DN DJG 1 #'*I  ekk)FZ_.`mr s 	r   r"   c                 $    | j                          yrp   rq   rs   s    r   rl   z#RotaryEmbeddingCat.reset_parametersK  rt   r   c                     | j                   s*| j                  j                  | j                                y| j                  j                  | j                  | j                               yrv   )r   r   rx   r   r   r   r*   rs   s    r   rr   z RotaryEmbeddingCat._init_buffersO  sG    %%JJT0023NN  !;!;DOO!LMr   c                     | j                   r9t        | j                  dz  t        | j                        | j
                        }n%t        | j                  dz  | j                  d      }|j                  ||      S r   r   r   s       r   r   z!RotaryEmbeddingCat._compute_bandsV  r   r   c                     t        || j                  | j                  | j                  | j                  | j
                  | j                  | j                  | j                  ||      }t        j                  |d      S )Nr   r1   )r   r+   rJ   r    r   rL   rM   rN   rO   r   r   )rm   r*   r   r   embedss        r   r   z(RotaryEmbeddingCat._get_pos_embed_valuesf  sj    '!LL((**nn..((,,
 yy$$r   c                 $    | j                          yrz   rq   rs   s    r   r{   z.RotaryEmbeddingCat.init_non_persistent_buffersv  rt   r   c                     | j                   g|| j                   k7  rW| j                  J | j                  || j                  j                  | j                  j                        | _        || _         y y y r   )r*   r   r   r   r   r   s     r   r   z$RotaryEmbeddingCat.update_feat_shapez  sm    ??&:+H>>---!77~~,,nn** 8 DN
 )DO ,I&r   r   c                    |e| j                   Yt        || j                   | j                  | j                  | j                  | j
                        }t        j                  |d      S | j                  | j                  S J d       )Nr   r1   zPget_embed() requires pre-computed pos embed or valid shape w/ pre-computed bands)	r   r   rL   rM   rN   rO   r   r   r   )rm   r   r   s      r   r   zRotaryEmbeddingCat.get_embed  sy    !7+

..#22 ,,"00F 99VR((^^'>>!lll5r   shapesseq_lenc                    |sg S | j                   t        d      t        d |D              }t        d |D              }t        ||f| j                   | j                  | j
                  | j                  | j                        \  }}t        j                  ||gd      j                  ||d      }|{t        j                  t        |      ||j                  d         j                  |      }t        |      D ]/  \  }	\  }
}|
|z  }|d|
d|f   j!                  |d      ||	d|f<   1 |S |D 
cg c]#  \  }
}|d|
d|f   j!                  |
|z  d      % }}
}|S c c}}
w )a  Generate ROPE embeddings for multiple grid shapes efficiently.

        Computes embeddings for the maximum grid size once, then extracts
        and flattens the relevant portions for each requested shape.

        Args:
            shapes: List of (H, W) tuples representing different grid sizes

        Returns:
            List of concatenated sin/cos embeddings for each shape,
            where each tensor has shape (H*W, dim)
        NzMBatch embedding generation requires cached bands, not pre-computed embeddingsc              3   &   K   | ]	  \  }}|  y wr    .0hws      r   	<genexpr>z6RotaryEmbeddingCat.get_batch_embeds.<locals>.<genexpr>       )$!QA)   c              3   &   K   | ]	  \  }}|  y wr   r   r   s      r   r   z6RotaryEmbeddingCat.get_batch_embeds.<locals>.<genexpr>  r   r   )r*   r   rL   rM   rN   rO   r1   r3   )r   RuntimeErrormaxr   rL   rM   rN   rO   r   r   r   zerosrD   r   type_as	enumerater   )rm   r   r   max_hmax_wr   r   rope_embed_2dflat_embedsir   r   src_lenflat_embeds_lists                 r   get_batch_embedsz#RotaryEmbeddingCat.get_batch_embeds  s   " I ::noo )&)))&)) 2u~**nn..((,,
 		7G"4"=BB5%QST++c&k7M<O<OPR<ST\\]deK&v. V	6Aqa%+8!RaR+@+H+HRT+UAxxK(V U[\TQPQbqb"1"f 5 = =a!eR H\\##  ]s   4(E!c                 V    | j                  |j                  dd        }t        ||      S r   r   r   r   rm   rY   r   s      r   r   zRotaryEmbeddingCat.forward  '    NN17712;/	"1i00r   r   r   r   r   )r   r   r   r   r   rU   r   r   r   r   rg   rl   rr   r   r   r   r   r{   r   r   r   r   r   r   r   r   r   s   @r   r   r     sE    !&"!&.226!#!%+ +  +  	+ 
 +  +  !c++  %T#Y/+  +  + ZN4  CGemm %S	 % 	)DI 	)mxS	2 m( &*3$sCx)3$ c]3$ 
u||T%,,//	0	3$j1r   r   r2   head_dimdepth	num_headsrotatec          	         d|t        j                  d| d||      | z  z  z  }|j                  d      j                  d      }|r/t        j                  ||d||      dz  t         j                  z  }nt        j
                  ||d||      }t        j                  |t        j                  |      z  |t        j                  |t         j                  dz  z         z  gd      }	t        j                  |t        j                  |      z  |t        j                  |t         j                  dz  z         z  gd      }
t        j                  |	|
gd      S )	z Vectorised 2D ROPE frequencies with random rotation for mixed mode ROPE.
    Returns:
         Tensor (2, depth, num_heads, head_dim//2)
    r   r   r/   r2   r   r   r1   r3   )
r   r$   r8   randr   r   r   r:   r9   r5   )r  r  r  r    r  r   r   maganglesfxfys              r   init_random_2d_freqsr    s!    a1VSX!Y\d!de
fC --

$
$Q
'C E9aeLqPSXS[S[[UIquM 
C%))F++S599VehhQRl=R3S-STZ\	]B	C%))F++S599VehhQRl=R3S-STZ\	]B ;;BxQ''r   r   c                 x   |dk(  rt        |       } t        j                  t        j                  | d   |t        j                        t        j                  | d   |t        j                        |      \  }}|j                  |      j                         }|j                  |      j                         }||fS )NrQ   r   r2   r   rS   )rF   r   rW   r$   r   r&   r6   )r   rO   r   r   x_posy_post_xt_ys           r   get_mixed_gridr    s     e$>>U1XfEMMBU1XfEMMBLE5
 ((5/
!
!
#C
((5/
!
!
#C8Or   freqsr  r  c                    | j                   }| j                         } |j                  d      | d   j                  d      z  }|j                  d      | d   j                  d      z  }||z   }t        j                  |      j                  dd      }t        j                  |      j                  dd      }t        j                  ||gd      }	|	j                  |      S )z&Compute mixed (learnable) frequencies.r1   r   r}   r   r   r3   )	r   rU   r8   r   r9   r   r:   r   r&   )
r  r  r  r   freqs_xfreqs_ycombinedr   r   rope_embedss
             r   get_mixed_freqsr    s     KKEKKME}}R 58#5#5b#99G}}R 58#5#5b#99G Hii!33Ar:Gii!33Ar:G))Wg.B7K>>%  r   c            	       v     e Zd ZdZ	 	 	 	 	 d
dedeeeef   dedef fdZde	e   de
j                  fd	Z xZS )RotaryEmbeddingMRopeaa  Interleaved multimodal RoPE (Qwen2-VL style) for vision, matching the reference GenLIP layout.

    Drop-in sibling of ``RotaryEmbeddingCat``: ``get_embed(shape) -> [N, 2*dim]``, consumed by
    ``apply_rot_embed_cat(..., half=True)`` (no new apply path / no separate sin/cos tensors). The ``dim // 2``
    frequency channels are assigned to height/width/temporal axes in a strided ``T,H,W,T,H,W,...`` interleave
    (the reference ``apply_interleaved_mrope``): channels ``1,4,7,...`` -> height, ``2,5,8,...`` -> width, and
    the remainder -> temporal. ``mrope_section`` sets the strided extent per axis; the actual per-axis channel
    *counts* equal ``mrope_section`` only for the standard equal-section configs that tile exactly
    (``3*section == dim // 2``, e.g. ``(12,12,12)`` -> 36 channels = 12/12/12), and are otherwise the clamped
    interleave (e.g. ``dim=64, (8,12,12)`` -> 11/11/10) -- this matches the reference, which also clamps.

    For an image encoder there is no text, so the temporal channels sit at position 0 (inert) and this reduces
    to a 2-axis ``(h, w)`` rope -- numerically identical to a checkpoint trained with the reference MRoPE.

    Only ``grid_indexing='ij'`` is supported (GenLIP / NaFlex ``(y, x)`` row-major patch order); ``'xy'`` would
    require mirroring the timm axial shape-swap and is intentionally not implemented here.
    r+   mrope_sectionr    rO   c                    t         |           |dz  dk(  sJ d       t        d |D              s
J d|        |dk(  sJ d       || _        || _        || _        || _        d|t        j                  d|d|	      j                         |z  z  z  }| j                  d
|d       |\  }}	}
t        j                  |dz  t        j                  |      }d|d|	dz  d<   d|d|
dz  d<   | j                  d|d       y )Nr   r   zdim (head_dim) must be evenc              3   &   K   | ]	  }|d k\    yw)r   Nr   )r   r<   s     r   r   z0RotaryEmbeddingMRope.__init__.<locals>.<genexpr>9  s     1a161r   z0mrope_section entries must be non-negative, got rI   zXRotaryEmbeddingMRope supports grid_indexing='ij' only (GenLIP/NaFlex (y,x) patch order).r   r   inv_freqFrd   r   r   r   axis)rf   rg   allr+   r  r    rO   r   r$   rU   rj   r   long)rm   r+   r  r    rO   r   r   r!  _sec_tsec_hsec_wr"  rn   s               r   rg   zRotaryEmbeddingMRope.__init__,  s1    	Qw!|:::| 1=11 	O>}oN	O1$ 	gf	g$*&* +%,,q#q*P*V*V*X[^*^_`ZeD
  -u{{3!85::fEQuqy]Quqy]VTe<r   r   r"   c                    |\  }}| j                   j                  }t        j                  t        j                  ||      t        j                  ||      | j
                        \  }}|j                  d      j                         |j                  d      j                         }}t        j                  |j                  d   | j                  dz  |      }|dddf   |dd| j                  dk(  f<   |dddf   |dd| j                  dk(  f<   || j                   z  }t        j                  ||gd      }	t        j                  |	j                         |	j                         gd      S )	zArgs:
            shape: ``(H, W)`` patch grid.

        Returns:
            Rope tensor of shape ``[H*W, 2*dim]`` for ``apply_rot_embed_cat(..., half=True)``.
        r   rS   r1   r   r   Nr   r3   )r!  r   r   rW   r$   rO   r   rU   r   r   r+   r"  r   r9   r:   )
rm   r   r   r   r   ysxsr\   r
  r   s
             r   r   zRotaryEmbeddingMRope.get_embedO  s'    1%%LL6*LL6*''
B
 B%%'B)=)=)?Bkk"((1+txx1}VD!#AtGAtyyA~!#AtGAtyyA~ t}}$ii(b1yy#'')SWWY/R88r   ))      r,  r   rI   NN)r   r   r   r   r   r   rU   r   rg   r   r   r   r   r   r   s   @r   r  r    sk    * 3>!'!%!=!= !c3/!= 	!=
 !=F9tCy 9U\\ 9r   r  c                   @    e Zd ZdZ	 	 	 	 	 ddededededeee      def fdZ	dd
Z
ddZdeee      fdZdeee      fdZddZddeee      d	ej                   fdZ	 ddeeeef      dee   d	eej                   eej                      f   fdZd Zd Z xZS )RotaryEmbeddingMixeda~  Rotary position embedding with depth-dependent learnable frequencies.

    This implementation supports mixed (learnable) ROPE. In mixed mode,
    each transformer block has its own set of learnable frequency parameters.

    Based on 'Rotary Position Embedding for Vision: https://arxiv.org/abs/2403.13298)'
    Compatible with original at https://github.com/naver-ai/rope-vit
    r+   r  r  r    r*   rO   c	           	         t         |           || _        || _        || _        || _        || _        || _        ||z  }	|	dz  dk(  s
J d|	        t        |	|||d||      }
t        j                  |
      | _        |sd}|D ]  }||z  }	 | j                  dt        j                  |||	      d
       | j                  dt        j                  |||	      d
       | j                          ydx| _        | _        y)a  Initialize rotary embeddings.

        Args:
            dim: Embedding dimension (should be divisible by 4)
            depth: Number of transformer blocks
            num_heads: Number of attention heads
            temperature: Base for frequency computation
            feat_shape: Spatial dimensions [H, W] if known in advance
            grid_indexing: How to index grid positions ('xy' or 'ij')
        r/   r   z%head_dim must be divisible by 4, got T)r    r  r   r   Nr   r  r2   Frd   r  )rf   rg   r+   r  r  r    r*   rO   r  r   	Parameterr  rj   r   rk   rr   r  r  )rm   r+   r  r  r    r*   rO   r   r   r  r  r   r<   rn   s                r   rg   zRotaryEmbeddingMixed.__init__r  s   * 	
"&$*)#!|q T$I("TT $#
 \\%(
!G 1  GFRW(Xej k  GFRW(Xej k "&&DHtxr   r"   c                     | j                   U| j                  | j                         \  }}| j                  j                  |       | j                  j                  |       yyrv   )r*   _get_grid_valuesr  rx   r  )rm   r  r  s      r   rr   z"RotaryEmbeddingMixed._init_buffers  sI    ??&,,T__=HCHHNN3HHNN3 'r   c                 $    | j                          yrp   rq   rs   s    r   rl   z%RotaryEmbeddingMixed.reset_parameters  rt   r   c                 h    t        || j                  | j                  j                        \  }}||fS )NrO   r   )r  rO   r  r   rm   r*   r  r  s       r   r2  z%RotaryEmbeddingMixed._get_grid_values  s4    !,,::$$
S
 Cxr   c                    | j                   || j                   k7  r| j                  J | j                  J | j                  |      \  }}|j	                  | j                  j
                  | j                  j                        | _        |j	                  | j                  j
                  | j                  j                        | _        || _         y y y r   )r*   r  r  r2  r&   r   r   r6  s       r   r   z&RotaryEmbeddingMixed.update_feat_shape  s    ??&:+H88'''88''',,Z8HCvvdhhootxx~~>DHvvdhhootxx~~>DH(DO ,I&r   c                 $    | j                          yrz   rq   rs   s    r   r{   z0RotaryEmbeddingMixed.init_non_persistent_buffers  rt   r   r   c                    |0t        || j                  | j                  j                        \  }}n8| j                  %| j
                  | j                  | j
                  }}nJ d       t        | j                  ||      S )zGenerate rotary embeddings for the given spatial shape.

        Args:
            shape: Spatial dimensions [H, W]

        Returns:
            Tensor of shape (depth, H*W, dim) containing concatenated sin/cos embeddings
        r5  z8get_embed() requires pre-computed t_x/t_y or valid shape)r  rO   r  r   r  r  r  )rm   r   r  r  s       r   r   zRotaryEmbeddingMixed.get_embed  su     %"00zz((HC
 XX!dhh&:xxCTTT5tzz344r   r   r   c           	         |sg S t        d |D              }t        d |D              }t        ||g| j                  | j                  j                        \  }}t        | j                  ||      }|j                  \  }}	}
}|j                  ||	|||      }|t        |      }t        j                  |||	||| j                  j                  | j                  j                        }t        |      D ]B  \  }\  }}|ddddd|d|f   j                  ||	||z  |      }||z  }|||ddddd|f<   D |S g }|D ]<  \  }}|ddddd|d|f   j                  ||	||z  |      }|j                  |       > |S )ai  Generate ROPE embeddings for multiple grid shapes efficiently.

        Computes embeddings for the maximum grid size once, then extracts
        and flattens the relevant portions for each requested shape.

        Args:
            shapes: List of (H, W) tuples representing different grid sizes
            seq_len: If provided, return padded tensor of this length. Otherwise return list.

        Returns:
            If seq_len is provided: Padded tensor of shape (len(shapes), depth, num_heads, seq_len, dim)
            Otherwise: List of tensors with shape (depth, num_heads, H*W, dim) for each shape
        c              3   &   K   | ]	  \  }}|  y wr   r   r   s      r   r   z8RotaryEmbeddingMixed.get_batch_embeds.<locals>.<genexpr>  r   r   c              3   &   K   | ]	  \  }}|  y wr   r   r   s      r   r   z8RotaryEmbeddingMixed.get_batch_embeds.<locals>.<genexpr>  r   r   r5  Nr2   )r   r  rO   r  r   r  r   r   rD   r   r   r   r   r   append)rm   r   r   r   r   r  r  	max_embedr  r  _r+   max_embed_2dr   paddedr   r   r   embed_slice
actual_lenresultss                        r   r   z%RotaryEmbeddingMixed.get_batch_embeds  s   $ I )&)))&)) "EN,,::$$
S
 $DJJS9	 $-?? y!S ~~eYucJFA[[E9gs4::K\K\dhdndndtdtuF&v. ;	6Aq*1a!RaR<8@@	STWXSXZ]^U
/:q!Q+,	;
 M G ,1*1a!RaR<8@@	STWXSXZ]^{+, Nr   c                 V    | j                  |j                  dd        }t        ||      S r   r   r  s      r   r   zRotaryEmbeddingMixed.forward  r  r   c                     dhS )z/Exclude frequency parameters from weight decay.r  r   rs   s    r   no_weight_decayz$RotaryEmbeddingMixed.no_weight_decay  s
    yr   )      $@NrQ   NNr   r   )r   r   r   r   r   rU   r   r   r   rg   rr   rl   r2  r   r{   r   r   r   r   r   r   r   rG  r   r   s   @r   r.  r.  i  s    "&.2!%5'5' 5' 	5'
 5' !c+5' 5'n 8DI+> )HT#Y,? )5xS	2 5ell 52 &*6sCx)6 c]6 
u||T%,,//	0	6p1
r   r.  separatecpuheightwidthnormalize_coordsc                    t        j                  d| |t         j                        |z   }t        j                  d||t         j                        |z   }|dk(  rt        t	        | |            }	|	}
|	}nI|dk(  rt        t        | |            }	|	}
|	}n*|dk(  rt        |       }
t        |      }nt        d|       ||
z  }||z  }|j                  |      }|j                  |      }|dk(  r5t        j                  ||d      \  }}t        j                  ||gd	
      }n-t        j                  t        j                  ||d      d	
      }|j                  dd      }d|z  dz
  }|S )z{Make coordinate grid matching offset and normalization of original.
    Returns: coords with shape (HW, 2) in [-1, 1].
    g      ?r2   r   minrI  zUnknown normalize_coords: rQ   rS   r1   r3   rI   r   r          @r   )r   r$   r   rU   r   rO  
ValueErrorr&   rW   r5   r6   )rK  rL  rM  rO   rN   r   r   coords_hcoords_wdenomh_denomw_denomgrid_wgrid_hcoordss                  r   make_coords_dinov3rZ  !  s]    ||CemmL{ZH||CvU]]KkYH 5 c&%()	U	"c&%()	Z	'-,56F5GHII '!H'!H{{5!H{{5!H (TJff-26U^^HhNTVW^^Aq!F6\CFMr   c                   p    e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 d"dedee   dee   dee   deee      deded	ed
e	dee   dee   dee   f fdZ
d#dZd#dZdej                  fdej                  dej                   dej"                  fdZdej"                  dej"                  fdZdej"                  deej"                  ej"                  f   fdZ	 d$dee   de	dej"                  fdZdee   fdZdee   fdZd#dZd%deee      dej"                  fdZd ej"                  dej"                  fd!Z xZS )&RotaryEmbeddingDinoV3a  RoPE for timm DinoV3 port, numerically matching original.

    Math is aligned to original DinoV3 RopePositionEmbedding at https://github.com/facebookresearch/dinov3:
      - 0.5-centered coords normalized by H/W (or min/max), mapped to [-1,1]
      - training-time augmentations (shift/jitter/rescale)
      - periods schedule equals Rope's temperature (base) or min/max period
    r+   r    
min_period
max_periodr*   rM  rN   rO   rotate_halfshift_coordsjitter_coordsrescale_coordsc                 t   t         |           || _        |	| _        t	        |      | _        || _        || _        || _        |
| _	        || _
        || _        t        | j                  | j                  | j                  fD cg c]  }|d u c}      | _        || _        || _        || _        |dz  f}| j#                  dt%        j&                  |||      d       |=|d   |d   z  }||dz  f}| j#                  d	t%        j&                  |||      d       nd | _        | j+                          y c c}w )
Nr/   periodsr2   Frd   r   r   r   pos_embed_cached)rf   rg   r+   r_  rU   r    r]  r^  rM  r`  ra  rb  any
aug_activer*   rN   rO   rj   r   rk   re  rl   )rm   r+   r    r]  r^  r*   rM  rN   rO   r_  r`  ra  rb  r   r   aperiods_shaper   r   rn   s                      r   rg   zRotaryEmbeddingDinoV3.__init__\  sG   " 	 & !-$$ !1(*,t7H7H$J\J\^b^q^q6rsq}st %&* YM&X](^kpq! mjm3G #'*I  !3U[[SYaf5gty z$(D! 	) ts   D5r"   c                 $    | j                          yrp   rq   rs   s    r   rl   z&RotaryEmbeddingDinoV3.reset_parameters  rt   r   c                     | j                   j                  | j                                | j                  F| j                  9| j                  | j                  d      }| j                  j                  |       yyy)rw   NTno_aug)rd  rx   _compute_periodsr*   re  _create_embed)rm   
rope_embeds     r   rr   z#RotaryEmbeddingDinoV3._init_buffers  sh    40023??&4+@+@+L++DOOD+IJ!!''
3 ,M&r   rJ  r   r   c                    | j                   dz  }| j                  ^| j                  Rt        j                  dd|dt        j
                        }| j                  | j                  | j                  z  |z  z  }n_| j                  t        d      dt        j                  |dt        j
                        z  | j                   dz  z  }| j                  |z  }|j                  ||      S )	z5Construct periods from either min/max or temperature.r/   r   r   rJ  r2   z0Provide either min/max periods or `temperature`.rP  r   )
r+   r]  r^  r   r   r   r    rQ  r$   r&   )rm   r   r   r+   	exponentsrd  s         r   rn  z&RotaryEmbeddingDinoV3._compute_periods  s    hh!m??&4??+Fq!SemmTIoo$//DOO*KPY)YZG' !STTell3uEMMRRVZV^V^bcVcdI&&)3G zzuz55r   rY  c                 ,   | j                   r| j                  s|S |j                  }|j                  }| j                  Jt        | j                        }t        j                  d||      j                  | |      }||dddf   z   }| j                  }t        | j                        }|dk  rt        d      t        j                  |      }t        j                  d||      j                  | |      j                         }||dddf   z  }| j                  vt        | j                        }	|	dk  rt        d      t        j                  |	      }
t        j                  d||      j                  |
 |
      j                         }||z  }|S )z4Apply shift/jitter/rescale train time augmentations.Nr   r2   r   zAjitter_coords must be > 0 (interpreted as multiplicative factor).zBrescale_coords must be > 0 (interpreted as multiplicative factor).r   )trainingrg  r   r   r`  rU   r   rk   uniform_ra  rQ  r   r   r'   rb  )rm   rY  r   r   shiftshift_hwjitter_factor
jitter_max	jitter_hwrescale_factorrescale_maxrescales               r   _apply_coord_augsz'RotaryEmbeddingDinoV3._apply_coord_augs  sq   }}DOOM ($++,E{{1V5AJJE6SXYHhtQw//F )!$"4"45M! !dee-0JAfEBKKZKYcdhhjIia00F *"4#6#67N" !eff((>2Kkk!F%@II;,XcdhhjGg%Fr   c                 &   | j                   dz  }| j                  j                  }| j                  j                  }| j                  j	                         |k(  sJ |dddddf   j                  ||      }dt        j                  z  |z  | j                  ddddf   z  }|j                  d      }| j                  r|j                  d      }n|j                  dd      }t        j                  |      }t        j                  |      }||fS )zEReturn sin/cos embeddings with either 'half' or 'interleaved' layout.r/   Nr2   r   r   r1   r3   )r+   rd  r   r   r   r&   r   r   r6   r_  tiler   r   r9   r:   )rm   rY  r+   r   r   r
  r9   r:   s           r   _get_pos_embed_from_coordsz0RotaryEmbeddingDinoV3._get_pos_embed_from_coords  s     hh!m$$""||!!#s*** 1d
#&&fE&BTWWv%T4](CC"[[^F --aR-8FiiiiCxr   rm  c                     |\  }}t        ||| j                  | j                  | j                        }|s| j	                  |      }| j                  |      \  }}t        j                  ||gd      }|S )N)rM  rO   rN   r1   r3   )rZ  rM  rO   rN   r~  r  r   r   )	rm   r*   rm  HWrY  r9   r:   rp  s	            r   ro  z#RotaryEmbeddingDinoV3._create_embed  sz    
 1#q!22,,((	
 ++F3F226:SYYSzr2
r   c                 `    | j                  |d      }| j                  d|d       || _        y )NTrl  re  Frd   )ro  rj   r*   )rm   r*   rp  s      r   _cache_embedz"RotaryEmbeddingDinoV3._cache_embed  s4    ''
4'@
/N$r   c                 `    | j                   "|| j                   k7  r| j                  |       y y y r   )r*   r  r   s     r   r   z'RotaryEmbeddingDinoV3.update_feat_shape  s.    ??&:+Hj) ,I&r   c                 $    | j                          yrz   rq   rs   s    r   r{   z1RotaryEmbeddingDinoV3.init_non_persistent_buffers  rt   r   r   c                    || j                  |      }|S | j                  du xs | j                  xr | j                  }|r0| j                  J d       | j                  | j                        }|S | j                  J | j                  }|S )zGenerate rope_embed matching DINOv3 RopePositionEmbedding numerics.

        Returns: (HW, num_heads, 2 * head_dim) with last dim = [sin, cos] cat.
        Nz&feature shape must be cached on create)ro  re  rt  rg  r*   )rm   r   rp  need_creates       r   r   zRotaryEmbeddingDinoV3.get_embed
  s    
 ++E2J  //47^DMM<]dooK2\4\\2!//@

  ,,888!22
r   rY   c                 n    | j                  |j                  dd       }t        ||| j                        S )z$Get and apply rotary embeddings to xr   N)r   )r   r   r   r_  r  s      r   r   zRotaryEmbeddingDinoV3.forward  s1     NN17712;/	"1id6F6FGGr   )g      Y@NNNrI  rH   rI   TNNNNNr   Fr   )r   r   r   r   r   r   rU   r   r   r   rg   rl   rr   r   r   r   r   r   rn  r~  r   r  ro  r  r   r{   r   r   r   r   s   @r   r\  r\  S  s    ,1*.*..2$.!$!% $,0-1.25 5  "%5  !	5 
 !5  !c+5  "5  5  5  5  #5/5  $E?5  %UO5 n4 7<RWR_R_ 6u|| 6EKK 6didpdp 6"     D %V[VbVbHbBc 6 !S	  
	$%tCy %*DI *
xS	2 ell $H H%,, Hr   r\  	rope_typec                    | dk(  r"|j                  dd       t        dd||z  i|S | dk(  r"|j                  dd       t        dd||z  i|S | dk(  r2|j                  dd       |j                  dd       t        d||d	|S | d
k(  r4|j                  dd       |j                  dd       t	        dd||z  i|S | dk(  r)dD ]  }|j                  |d        t        dd||z  i|S t        d|        )a  Factory function for creating rotary position embeddings.

    Args:
        rope_type: Type of RoPE to create. Options:
            - 'base': Basic RotaryEmbedding
            - 'cat': RotaryEmbeddingCat (concatenated sin/cos)
            - 'mixed': RotaryEmbeddingMixed (learnable per-depth frequencies)
            - 'dinov3': RotaryEmbeddingDinoV3 (with coordinate transforms)
            - 'mrope': RotaryEmbeddingMRope (interleaved multimodal RoPE; requires `mrope_section`)
        dim: Total embedding dimension
        num_heads: Number of attention heads
        **kwargs: Additional arguments passed to the specific RoPE class

    Returns:
        Rotary embedding module
    baser_  Nr+   r   mixedrL   rM   )r+   r  dinov3mrope)rL   rM   r_  zUnknown RoPE type: r   )popr   r   r.  r\  r  rQ  )r  r+   r  kwargsks        r   create_rope_embedr  #  s'   , F

=$'>3)#3>v>>	e	

=$'!AcY&6A&AA	g	

;%

#T*#KyKFKK	h	

;%

#T*$D	)9DVDD	g	? 	 AJJq$	 #Cy(8CFCC.yk:;;r   )g      l@TN)r   r   Nr  )rH  T)r   i   r,  )1r   r   typingr   r   r   r   r   r   _fxr	   r=   r
   trace_utilsr   r   rU   r   r   r   r   r(   r   r   rA   rF   r   r`   Modulerb   r   r   r   r   r   r   r   r   r   r  r  wrapr  r  r  r.  rZ  r\  r  r   r   r   <module>r     s    / /   *    !)-	


 
 &	
 $)-	  &	
 \\ ###()-"]]&#I&#&# &# 	&#
 !&# &&# {{&# \\&#R,tCy ,T#Y , )-#"".2!)-"]]RIR%R R 	R
 R R R R !c+R R R &R {{R 
%,,Rj6299 6rJ( ( ( 	.<<.. . 	.
 \\.. 	;;; ; 	;
 
%,,;0 .<<.\\. . \\	.. %*	".<<".<<". ll". "	".
 \\".N )-#".2!)-"]]5I5%5 5 	5
 5 5 5 !c+5 5 5 &5 {{5pJ4bii J4Zx1 x1~ "( mm((( ( 	(
 ( \\(D  ")-"]]	Cy & {{	
 5<<%&  $!||!\\! \\! \\	!$M9299 M9`u299 up  !+!$"]]--- - 	-
 - - {{- \\-  -`MHBII MHb *<*<*< *<
 YY*<r   