+
    UV-jpA  ã                   ó:  € ^ RI HtHtHt ^ RIHt ^ RIHt R t	 ! R R]P                  4      t ! R R]P                  4      t ! R R	]P                  4      t ! R
 R]P                  4      t ! R R]P                  4      tR R ltR R ltR R ltR R ltR# )é    )ÚOptionalÚTupleÚTypeNc                ó$  € V P                   pV P                  ^,          pW18w  dk   T P                  ^ ^^^4      pVP                  \        P
                  4      p^RIHp V! WAV3RR7      P                  V4      pVP                  ^ ^^^4      pV# V # )z:Interpolate absolute positional embeddings to target size.)Úbicubic_interpolateT)ÚsizeÚ	antialias)ÚdtypeÚshapeÚ	transposeÚastypeÚmxÚfloat32Úkernelsr   )Úabs_posÚtgt_sizer
   Úsrc_sizeÚold_pos_embedr   Únew_pos_embeds   &&     Úo/Volumes/fast/ai/experiments/ui-tars-smoke/.venv/lib/python3.14/site-packages/mlx_vlm/models/deepseekocr/sam.pyÚget_abs_pos_samr      s‘   € à�M‰M€EØ�}‰}˜QÕ€HàÔà×)Ñ)¨!¨Q°°1Ó5ˆØ%×,Ñ,¬R¯Z©ZÓ8ˆõ 	2á+Ø¨8Ð 4Àô
ç
‰&�‹-ð 	ð
 &×/Ñ/°°1°a¸Ó;ˆØÐàˆó    c                   ón   a a€ ] tR t^t oRt]P                  3V3R lV 3R llltV3R lR ltRt	Vt
V ;t# )ÚMLPBlockzMLP block with GELU activation.c                óT   <€ V ^8„  d   QhRS[ RS[ RS[S[P                  ,          RR/# )é   Úembedding_dimÚmlp_dimÚactÚreturnN)Úintr   ÚnnÚModule)ÚformatÚ__classdict__s   "€r   Ú__annotate__ÚMLPBlock.__annotate__"   s;   ø€ ÷ 	ñ 	áð	ñ ð	ñ ‘"—)‘)�_ð		ð
 
ñ	r   c                ó¨   <€ \         SV `  4        \        P                  ! W4      V n        \        P                  ! W!4      V n        V! 4       V n        R # ©N)ÚsuperÚ__init__r"   ÚLinearÚlin1Úlin2r   )Úselfr   r   r   Ú	__class__s   &&&&€r   r+   ÚMLPBlock.__init__"   s9   ø€ ô 	‰ÑÔÜ—I’I˜mÓ5ˆŒ	Ü—I’I˜gÓ5ˆŒ	Ù“5ˆŽr   c                óN   <€ V ^8„  d   QhRS[ P                  RS[ P                  /# ©r   Úxr    ©r   Úarray)r$   r%   s   "€r   r&   r'   -   s#   ø€ ÷ 1ñ 1™"Ÿ(™(ð 1¡r§x¡xñ 1r   c                ó`   € V P                  V P                  V P                  V4      4      4      # r)   )r.   r   r-   ©r/   r4   s   &&r   Ú__call__ÚMLPBlock.__call__-   s"   € Ø�y‰y˜Ÿ™ $§)¡)¨A£,Ó/Ó0Ð0r   )r   r-   r.   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r"   ÚGELUr+   r9   Ú__static_attributes__Ú__classdictcell__Ú__classcell__©r0   r%   s   @@r   r   r      s*   ù‡ € Ù)ð  "Ÿw™w÷		õ 	÷1÷ 1ð 1r   r   c                   óX   a a€ ] tR t^1t oRtRV3R lV 3R llltV3R lR ltRtVtV ;t	# )Ú	Attentionz=Multi-head Attention block with relative position embeddings.c                ób   <€ V ^8„  d   QhRS[ RS[ RS[RS[RS[S[S[ S[ 3,          ,          RR/# )r   ÚdimÚ	num_headsÚqkv_biasÚuse_rel_posÚ
input_sizer    N)r!   Úboolr   r   )r$   r%   s   "€r   r&   ÚAttention.__annotate__4   s\   ø€ ÷  Iñ  Iáð Iñ ð Iñ ð	 Iñ
 ð Iñ ™U¡3© 8�_Õ-ð Ið 
ñ Ir   c                óì  <€ \         SV `  4        W n        W,          pVR,          V n        \        P
                  ! W^,          VR7      V n        \        P
                  ! W4      V n        W@n        V P                  '       dr   Vf   Q R4       h\        P                  ! ^V^ ,          ,          ^,
          V34      V n        \        P                  ! ^V^,          ,          ^,
          V34      V n        R# R# )a~  
Args:
    dim (int): Number of input channels.
    num_heads (int): Number of attention heads.
    qkv_bias (bool): If True, add a learnable bias to query, key, value.
    use_rel_pos (bool): If True, add relative positional embeddings to the attention map.
    input_size (tuple(int, int) or None): Input resolution for calculating the relative
        positional parameter size.
)ÚbiasNzBInput size must be provided if using relative positional encoding.g      à¿)r*   r+   rI   Úscaler"   r,   ÚqkvÚprojrK   r   ÚzerosÚ	rel_pos_hÚ	rel_pos_w)r/   rH   rI   rJ   rK   rL   Úhead_dimr0   s   &&&&&& €r   r+   ÚAttention.__init__4   sÅ   ø€ ô" 	‰ÑÔØ"ŒØÕ#ˆØ˜t•^ˆŒ
ä—9’9˜S¨¥'°Ô9ˆŒÜ—I’I˜cÓ'ˆŒ	à&ÔØ××ÐàÒ&ðTàSóTØ&ô  ŸXšX q¨:°a­=Õ'8¸1Õ'<¸hÐ&GÓHˆDŒNÜŸXšX q¨:°a­=Õ'8¸1Õ'<¸hÐ&GÓHˆDŽNñ r   c                óN   <€ V ^8„  d   QhRS[ P                  RS[ P                  /# r3   r5   )r$   r%   s   "€r   r&   rN   V   s#   ø€ ÷ 3ñ 3™"Ÿ(™(ð 3¡r§x¡xñ 3r   c                óú  € VP                   w  r#rEV P                  V4      P                  W#V,          ^V P                  R4      P	                  ^^ ^^^4      pVP                  ^W P                  ,          W4,          R4      pV^ ,          V^,          V^,          r©pRRrËV P
                  '       d'   \        W€P                  V P                  W43W434      w  r¼VP                  W P                  W4,          R4      pV	P                  W P                  W4,          R4      p	V
P                  W P                  W4,          R4      p
V P
                  '       Ed-   VP                  W P                  VP                   ^,          VP                   ^,          VP                   ^,          4      pVP                  W P                  VP                   ^,          VP                   ^,          VP                   ^,          4      pW¼,           P                  W P                  VP                   ^,          VP                   ^,          VP                   ^,          ,          4      p\        P                  P                  W‰W P                  VR7      pM+\        P                  P                  W‰W P                  R7      pVP                  W P                  W4R4      P	                  ^ ^^^^4      P                  W#VR4      pV P                  V4      pV# )é   N)rQ   Úmask)rQ   éÿÿÿÿ)r   rR   ÚreshaperI   r   rK   Úadd_decomposed_rel_posrU   rV   r   ÚfastÚscaled_dot_product_attentionrQ   rS   )r/   r4   ÚBÚHÚWÚ_rR   Úqkv_reshapedÚqÚkÚvÚrel_hÚrel_wÚ	attn_biass   &&            r   r9   ÚAttention.__call__V   sH  € Ø—W‘W‰
ˆˆað �H‰H�Q‹Kß‰W�Q˜A�˜q $§.¡.°"Ó5ß‰Y�q˜!˜Q  1Ó%ð 	ð —{‘{ 1 a¯.©.Õ&8¸!½%ÀÓDˆØ˜q•/ <°¥?°LÀµOˆaˆð ˜TˆuØ××ÐÜ1Ø—>‘> 4§>¡>°A°6¸A¸6ó‰LˆEð
 �I‰I�aŸ™¨­°Ó3ˆØ�I‰I�aŸ™¨­°Ó3ˆØ�I‰I�aŸ™¨­°Ó3ˆð ××ÑØ—M‘MØ—>‘> 5§;¡;¨q¥>°5·;±;¸qµ>À5Ç;Á;ÈqÅ>óˆEð —M‘MØ—>‘> 5§;¡;¨q¥>°5·;±;¸qµ>À5Ç;Á;ÈqÅ>óˆEð �×/Ñ/Ø—>‘> 5§;¡;¨q¥>°5·;±;¸qµ>ÀEÇKÁKÐPQÅNÕ3RóˆIô —‘×4Ñ4Ø�aŸz™z°	ð 5ó ‰Aô —‘×4Ñ4°Q¸1ÇJÁJÐ4ÓOˆAð �I‰I�aŸ™¨¨rÓ2ß‰Y�q˜!˜Q  1Ó%ß‰W�Q˜1˜bÓ!ð 	
ð �I‰I�a‹LˆØˆr   )rI   rS   rR   rU   rV   rQ   rK   )é   TFN©
r;   r<   r=   r>   r?   r+   r9   rA   rB   rC   rD   s   @@r   rF   rF   1   s$   ù‡ € ÙG÷ Iõ  I÷D3÷ 3ð 3r   rF   c                   óŽ   a a€ ] tR t^Œt oRtRR]P                  ]P                  R^ R3V3R lV 3R llltV3R lR	 lt	R
t
VtV ;t# )ÚBlockzMTransformer blocks with support of window attention and residual propagation.ç      @TFNc                óÂ   <€ V ^8„  d   QhRS[ RS[ RS[RS[RS[S[P
                  ,          RS[S[P
                  ,          RS[RS[ R	S[S[S[ S[ 3,          ,          R
R/
# )r   rH   rI   Ú	mlp_ratiorJ   Ú
norm_layerÚ	act_layerrK   Úwindow_sizerL   r    N)r!   ÚfloatrM   r   r"   r#   r   r   )r$   r%   s   "€r   r&   ÚBlock.__annotate__�   sŒ   ø€ ÷ )'ñ )'áð)'ñ ð)'ñ ð	)'ñ
 ð)'ñ ™Ÿ™•Oð)'ñ ™Ÿ	™	•?ð)'ñ ð)'ñ ð)'ñ ™U¡3© 8�_Õ-ð)'ð 
ñ)'r   c
                óî   <€ \         S
V `  4        V! VRR7      V n        \        TTTTV^ 8X  d   T	MWˆ3R7      V n        V! VRR7      V n        \        V\        W,          4      VR7      V n        W€n	        R# )a¢  
Args:
    dim (int): Number of input channels.
    num_heads (int): Number of attention heads in each ViT block.
    mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
    qkv_bias (bool): If True, add a learnable bias to query, key, value.
    norm_layer (nn.Module): Normalization layer.
    act_layer (nn.Module): Activation layer.
    use_rel_pos (bool): If True, add relative positional embeddings to the attention map.
    window_size (int): Window size for window attention blocks. If it equals 0, then
        use global attention.
    input_size (tuple(int, int) or None): Input resolution for calculating the relative
        positional parameter size.
ç�íµ ÷Æ°>©Úeps)rI   rJ   rK   rL   )r   r   r   N)
r*   r+   Únorm1rF   ÚattnÚnorm2r   r!   Úmlprw   )r/   rH   rI   rt   rJ   ru   rv   rK   rw   rL   r0   s   &&&&&&&&&&€r   r+   ÚBlock.__init__�   sw   ø€ ô4 	‰ÑÔÙ ¨Ô.ˆŒ
ÜØØØØ#Ø%0°AÔ%5‘z¸KÐ;Uô
ˆŒ	ñ   ¨Ô.ˆŒ
ÜØ¤s¨3­?Ó';Àô
ˆŒð 'Ör   c                óN   <€ V ^8„  d   QhRS[ P                  RS[ P                  /# r3   r5   )r$   r%   s   "€r   r&   ry   º   s#   ø€ ÷ ñ ™"Ÿ(™(ð ¡r§x¡xñ r   c                ó˜  € TpV P                  V4      pV P                  ^ 8”  d=   VP                  ^,          VP                  ^,          rC\        WP                  4      w  rV P	                  V4      pV P                  ^ 8”  d   \        WP                  XXX34      pW!,           pWP                  V P                  V4      4      ,           pV# )r   )r~   rw   r   Úwindow_partitionr   Úwindow_unpartitionr�   r€   )r/   r4   Úshortcutrc   rd   Úpad_hws   &&    r   r9   ÚBlock.__call__º   s§   € ØˆØ�J‰J�q‹Mˆð ×Ñ˜aÔØ—7‘7˜1•:˜qŸw™w q�zˆqÜ(¨×,<Ñ,<Ó=‰IˆAà�I‰I�a‹Lˆð ×Ñ˜aÔÜ" 1×&6Ñ&6¸ÀÀAÀÓGˆAà�LˆØ—‘˜Ÿ™ A›Ó'Õ'ˆàˆr   )r   r�   r~   r€   rw   ©r;   r<   r=   r>   r?   r"   Ú	LayerNormr@   r+   r9   rA   rB   rC   rD   s   @@r   rq   rq   Œ   sA   ù‡ € ÙWð ØØ&(§l¡lØ%'§W¡WØ!ØØ04÷)'õ )'÷V÷ ð r   rq   c                   óX   a a€ ] tR t^Ït oRtRV3R lV 3R llltV3R lR ltRtVtV ;t	# )Ú
PatchEmbedzImage to Patch Embedding.c          
      ób   <€ V ^8„  d   QhRS[ S[S[3,          RS[ S[S[3,          RS[RS[RR/# )r   Úkernel_sizeÚstrideÚin_chansÚ	embed_dimr    N)r   r!   )r$   r%   s   "€r   r&   ÚPatchEmbed.__annotate__Ò   sM   ø€ ÷ 
ñ 
á™3¡˜8•_ð
ñ ‘c™3�h•ð
ñ ð	
ñ
 ð
ð 
ñ
r   c                ó^   <€ \         SV `  4        \        P                  ! W4WR7      V n        R# )zÝ
Args:
    kernel_size (Tuple): kernel size of the projection layer.
    stride (Tuple): stride of the projection layer.
    in_chans (int): Number of input image channels.
    embed_dim (int): Patch embedding dimension.
)r�   r�   N)r*   r+   r"   ÚConv2drS   )r/   r�   r�   r‘   r’   r0   s   &&&&&€r   r+   ÚPatchEmbed.__init__Ò   s%   ø€ ô 	‰ÑÔÜ—I’IØ¨[ô
ˆŽ	r   c                óN   <€ V ^8„  d   QhRS[ P                  RS[ P                  /# r3   r5   )r$   r%   s   "€r   r&   r“   å   s#   ø€ ÷ ñ ™"Ÿ(™(ð ¡r§x¡xñ r   c                ó(   € V P                  V4      pV# r)   ©rS   r8   s   &&r   r9   ÚPatchEmbed.__call__å   s   € Ø�I‰I�a‹LˆØˆr   r™   )©é   rœ   r›   r[   é   ro   rD   s   @@r   r�   r�   Ï   s!   ù‡ € Ù#÷
õ 
÷&÷ ð r   r�   c                   ó    a a€ ] tR t^êt oRtR^^R^^RRR]P                  ]P                  RR^RR3V3R lV 3R llltV3R	 lR
 lt	Rt
VtV ;t# )Ú
SAMEncoderz5Vision Transformer encoder based on SAM architecture.i   r�   rr   é   Tc          "      óÚ   <€ V ^8„  d   QhRS[ RS[ RS[ RS[ RS[ RS[ RS[RS[ R	S[R
S[S[P
                  ,          RS[S[P
                  ,          RS[RS[RS[ RS[S[ R3,          RS[ RR/# )r   Úimg_sizeÚ
patch_sizer‘   r’   ÚdepthrI   rt   Ú	out_chansrJ   ru   rv   Úuse_abs_posrK   rw   Úglobal_attn_indexes.Úfinal_out_chansr    N)r!   rx   rM   r   r"   r#   r   )r$   r%   s   "€r   r&   ÚSAMEncoder.__annotate__í   sà   ø€ ÷ R
ñ R
áðR
ñ ðR
ñ ð	R
ñ
 ðR
ñ ðR
ñ ðR
ñ ðR
ñ ðR
ñ ðR
ñ ™Ÿ™•OðR
ñ ™Ÿ	™	•?ðR
ñ ðR
ñ ðR
ñ ðR
ñ  #¡3¨ 8�_ð!R
ñ" ð#R
ð$ 
ñ%R
r   c                óÜ  <€ \         SV `  4        Wn        \        W"3W"3VVR7      V n        WÀn        V'       d,   \        P                  ! ^W,          W,          V34      V n        . V n	        \        V4       FI  p\        TTTT	T
TTVV9  d   TM^ W,          W,          3R7	      pV P                  P                  V4       KK  	  \        P                  ! WH^RR7      \        P                  ! VRR7      \        P                  ! Wˆ^^RR7      \        P                  ! VRR7      .V n        \        P                  ! RR	^^^RR
7      V n        \        P                  ! R	V^^^RR
7      V n        R# )a²  
Args:
    img_size (int): Input image size.
    patch_size (int): Patch size.
    in_chans (int): Number of input image channels.
    embed_dim (int): Patch embedding dimension.
    depth (int): Depth of ViT.
    num_heads (int): Number of attention heads in each ViT block.
    mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
    out_chans (int): Output channels for neck.
    qkv_bias (bool): If True, add a learnable bias to query, key, value.
    norm_layer (nn.Module): Normalization layer.
    act_layer (nn.Module): Activation layer.
    use_abs_pos (bool): If True, use absolute positional embeddings.
    use_rel_pos (bool): If True, add relative positional embeddings to the attention map.
    window_size (int): Window size for window attention blocks.
    global_attn_indexes (tuple): Indexes for blocks using global attention.
    final_out_chans (int): Final output channels after net_3 (1024 for OCR, 896 for OCR-2).
)r�   r�   r‘   r’   )	rH   rI   rt   rJ   ru   rv   rK   rw   rL   F)r�   rP   r{   r|   )r�   ÚpaddingrP   r    i   )r�   r�   r«   rP   N)r*   r+   r¢   r�   Úpatch_embedr¦   r   rT   Ú	pos_embedÚblocksÚrangerq   Úappendr"   r•   r‹   ÚneckÚnet_2Únet_3)r/   r¢   r£   r‘   r’   r¤   rI   rt   r¥   rJ   ru   rv   r¦   rK   rw   r§   r¨   ÚiÚblockr0   s   &&&&&&&&&&&&&&&&&  €r   r+   ÚSAMEncoder.__init__í   sJ  ø€ ôL 	‰ÑÔØ Œä%Ø#Ð0ØÐ+ØØô	
ˆÔð 'ÔßäŸXšXØ�HÕ*¨HÕ,BÀIÐNóˆDŒNð ˆŒÜ�u–ˆAÜØØ#Ø#Ø!Ø%Ø#Ø'Ø+,Ð4GÔ+G™KÈQØ$Õ2°HÕ4JÐKô
ˆEð �K‰K×Ñ˜uÖ%ñ ô  �IŠI�i¸ÀÔFÜ�LŠL˜¨Ô-Ü�IŠI�i¸À1È5ÔQÜ�LŠL˜¨Ô-ð	
ˆŒ	ô —Y’Y˜s C°Q¸qÈ!ÐRWÔXˆŒ
Ü—Y’YØ�¨a¸À1È5ô
ˆŽ
r   c                óN   <€ V ^8„  d   QhRS[ P                  RS[ P                  /# r3   r5   )r$   r%   s   "€r   r&   r©   A  s#   ø€ ÷ ñ ™"Ÿ(™(ð ¡r§x¡xñ r   c                óX  € V P                  V4      pV P                  '       d/   V\        V P                  VP                  ^,          4      ,           pV P
                   F  pV! V4      pK  	  V P                   F  pV! V4      pK  	  V P                  V4      pV P                  V4      pV# )é   )	r¬   r¦   r   r­   r   r®   r±   r²   r³   )r/   r4   ÚblkÚns   &&  r   r9   ÚSAMEncoder.__call__A  s�   € à×Ñ˜QÓˆð ××ÐØ”O D§N¡N°A·G±G¸AµJÓ?Õ?ˆAð —;”;ˆCÙ�A“ŠAñ ð —”ˆAÙ�!“ŠAñ ð �J‰J�q‹MˆØ�J‰J�q‹Mˆàˆr   )r®   r¢   r±   r²   r³   r¬   r­   r¦   )r   é   rn   é   rŠ   rD   s   @@r   rŸ   rŸ   ê   s^   ù‡ € Ù?ð ØØØØØØØØØ&(§l¡lØ%'§W¡WØ Ø ØØ/<Ø#÷#R
õ R
÷h÷ ð r   rŸ   c          
      ó¨   € V ^8„  d   QhR\         P                  R\        R\        \         P                  \        \        \        3,          3,          /# )r   r4   rw   r    ©r   r6   r!   r   )r$   s   "r   r&   r&   [  s<   € ÷ ñ œŸ™ð ¬sð ´u¼R¿X¹XÄuÌSÔRUÈXÅÐ=VÕ7Wñ r   c                óŠ  € V P                   w  r#rEWV,          ,
          V,          pWV,          ,
          V,          pV^ 8”  g   V^ 8”  d    \        P                  ! V R^ V3^ V3R.4      p W6,           WG,           r˜V P                  W(V,          WV,          W4      p V P	                  ^ ^^^^^4      P                  RWV4      p
W¨V	33# )a8  
Partition into non-overlapping windows with padding if needed.

Args:
    x (mx.array): input tokens with [B, H, W, C].
    window_size (int): window size.

Returns:
    windows: windows after partition with [B * num_windows, window_size, window_size, C].
    (Hp, Wp): padded height and width before partition
)r   r   r]   )r   r   Úpadr^   r   )r4   rw   rb   rc   rd   ÚCÚpad_hÚpad_wÚHpÚWpÚwindowss   &&         r   r…   r…   [  s½   € ð —‘�J€Aˆ!à˜{�?Õ*¨kÕ9€EØ˜{�?Õ*¨kÕ9€Eàˆq„y�E˜A”IÜ�FŠF�1�v  5˜z¨A¨u¨:°vÐ>Ó?ˆà�Y˜�	ˆà	�	‰	�!˜;Õ&¨¸;Õ5FÈÓW€AØ�k‰k˜!˜Q  1 a¨Ó+×3Ñ3°B¸ÐRSÓT€Gà˜�HÐÐr   c          
      óÀ   € V ^8„  d   QhR\         P                  R\        R\        \        \        3,          R\        \        \        3,          R\         P                  /# )r   rÈ   rw   rˆ   Úhwr    rÀ   )r$   s   "r   r&   r&   w  sR   € ÷ ñ Ü�X‰Xðäðô ”#”s�(�Oðô 	Œc”3ˆh�ð	ô
 ‡X�Xñr   c                ó<  € Vw  rEVw  rgV P                   ^ ,          WE,          V,          V,          ,          pV P                  W„V,          WQ,          WR4      p	V	P                  ^ ^^^^^4      P                  W„VR4      p	WF8”  g   WW8”  d   V	RRV1RV1R3,          p	V	# )az  
Window unpartition into original sequences and removing padding.

Args:
    windows (mx.array): input tokens with [B * num_windows, window_size, window_size, C].
    window_size (int): window size.
    pad_hw (Tuple): padded height and width (Hp, Wp).
    hw (Tuple): original height and width (H, W) before padding.

Returns:
    x: unpartitioned sequences with [B, H, W, C].
ºNNNNr]   )r   r^   r   )
rÈ   rw   rˆ   rÊ   rÆ   rÇ   rc   rd   rb   r4   s
   &&&&      r   r†   r†   w  s    € ð$ �F€BØ�D€AØ�‰�aÕ˜R�W¨Õ3°{ÕBÕC€Aà�‰Ø	�Õ˜bÕ/°È2ó	€Að 	
�‰�A�q˜!˜Q  1Ó%×-Ñ-¨a°R¸Ó<€Aà	„v�”Øˆa��!��R�a�R˜ˆl�Oˆà€Hr   c                óp   € V ^8„  d   QhR\         R\         R\        P                  R\        P                  /# )r   Úq_sizeÚk_sizeÚrel_posr    )r!   r   r6   )r$   s   "r   r&   r&   ˜  s0   € ÷ *=ñ *=œð *=¤Sð *=´2·8±8ð *=ÄÇÁñ *=r   c                ó’  € \        ^\        W4      ,          ^,
          4      pVP                  ^ ,          V8w  Ed²   VP                  pVP	                  \
        P                  4      pVP                  ^VP                  ^ ,          R4      P                  ^ ^^4      pVP                  ^,          V,          p\
        P                  ! V\
        P                  R7      V,          p\
        P                  ! V4      P	                  \
        P                  4      p\
        P                  ! V^,           VP                  ^,          ^,
          4      p	WxP	                  \
        P                  4      ,
          p
\
        P                  ! WX^R7      ^V
,
          ,          \
        P                  ! WY^R7      V
,          ,           P	                  V4      pVP                  RV4      P                  ^^ 4      pMTp\
        P                  ! V \
        P                  R7      R,          \        W,          R4      ,          p\
        P                  ! V\
        P                  R7      R,          \        W,          R4      ,          pW¼,
          V^,
          \        W,          R4      ,          ,           pW]P	                  \
        P                  4      ,          # )a7  
Get relative positional embeddings according to the relative positions of
query and key sizes.

Args:
    q_size (int): size of query q.
    k_size (int): size of key k.
    rel_pos (mx.array): relative position embeddings (L, C).

Returns:
    Extracted positional embeddings according to relative positions.
)r
   )Úaxisg      ð?r]   )rÌ   N)NrÌ   )r!   Úmaxr   r
   r   r   r   r^   r   ÚarangeÚfloorÚint32ÚminimumÚtake)rÎ   rÏ   rÐ   Úmax_rel_distr
   Úrel_pos_resizedrQ   ÚindicesÚ	idx_floorÚidx_ceilÚweightÚq_coordsÚk_coordsÚrelative_coordss   &&&           r   Úget_rel_posrâ   ˜  sð  € ô �qœ3˜vÓ.Õ.°Õ2Ó3€Lð ‡}�}�QÕ˜<Õ'Ø—‘ˆØ—.‘.¤§¡Ó,ˆØ!Ÿ/™/¨!¨W¯]©]¸1Õ-=¸rÓB×LÑLÈQÐPQÐSTÓUˆð  ×%Ñ% aÕ(¨<Õ7ˆÜ—)’)˜L´·
±
Ô;¸eÕCˆÜ—H’H˜WÓ%×,Ñ,¬R¯X©XÓ6ˆ	Ü—:’:˜i¨!�m¨_×-BÑ-BÀ1Õ-EÈÕ-IÓJˆØ×+Ñ+¬B¯J©JÓ7Õ7ˆô �GŠG�O°QÔ7¸1¸v½:ÕFÜ�gŠg�o°aÔ8¸6ÕAõBç
‰&�‹-ð 	ð
 *×1Ñ1°"°lÓC×MÑMÈaÐQRÓS‰à!ˆô �yŠy˜¤r§z¡zÔ2°7Õ;¼cÀ&Å/ÐSVÓ>WÕW€HÜ�yŠy˜¤r§z¡zÔ2°7Õ;¼cÀ&Å/ÐSVÓ>WÕW€HØÕ*¨v¸­z¼SÀÅÐRUÓ=VÕ.VÕV€Oà×1Ñ1´"·(±(Ó;Õ<Ð<r   c                ó*  € V ^8„  d   QhR\         P                  R\         P                  R\         P                  R\        \        \        3,          R\        \        \        3,          R\        \         P                  \         P                  3,          /# )r   rg   rU   rV   rÎ   rÏ   r    )r   r6   r   r!   )r$   s   "r   r&   r&   Å  sr   € ÷ $ñ $Ü	‡x�xð$ä�x‰xð$ô �x‰xð$ô ”#”s�(�Oð	$ô
 ”#”s�(�Oð$ô Œ2�8‰8”R—X‘XÐÕñ$r   c                ót  € Vw  rVVw  rx\        WWV4      p	\        WhV4      p
V P                  w  r¼pV P                  WµWm4      p\        P                  ! RWé4      p\        P                  ! RWê4      pVR,          pVR,          pVP                  WµV,          V^4      pVP                  WµV,          ^V4      pVV3# )a  
Calculate decomposed Relative Positional Embeddings.

Args:
    q (mx.array): query q in the attention layer with shape (B, q_h * q_w, C).
    rel_pos_h (mx.array): relative position embeddings (Lh, C) for height axis.
    rel_pos_w (mx.array): relative position embeddings (Lw, C) for width axis.
    q_size (Tuple): spatial sequence size of query q with (q_h, q_w).
    k_size (Tuple): spatial sequence size of key k with (k_h, k_w).

Returns:
    Tuple of (rel_h, rel_w): relative position biases for height and width.
zbhwc,hkc->bhwkzbhwc,wkc->bhwk).N).NrÌ   )râ   r   r^   r   Úeinsum)rg   rU   rV   rÎ   rÏ   Úq_hÚq_wÚk_hÚk_wÚRhÚRwrb   re   rH   Úr_qrj   rk   s   &&&&&            r   r_   r_   Å  s³   € ð( �H€CØ�H€Cä	�S˜yÓ	)€BÜ	�S˜yÓ	)€Bà—‘�I€Aˆ#Ø
�)‰)�A˜CÓ
%€Cä�IŠIÐ&¨Ó0€EÜ�IŠIÐ&¨Ó0€EØ�)Õ€EØ�,Õ€EØ�M‰M˜! 3�Y¨¨QÓ/€EØ�M‰M˜! 3�Y¨¨3Ó/€Eà�%ˆ<Ðr   )Útypingr   r   r   Úmlx.coreÚcorer   Úmlx.nnr"   r   r#   r   rF   rq   r�   rŸ   r…   r†   râ   r_   © r   r   Ú<module>rò      s…   ðß (Ñ (å Ý òô01ˆr�y‰yô 1ô$X�—	‘	ô Xôv@ˆB�I‰Iô @ôF�—‘ô ô6k�—‘ô kõbõ8õB*=÷Z$r   