+
    UV-j‹  ã                   ó¬   € R t ^ RIt^ RIHtHt ^ RIHt ^ RIH	t	  ! R R]	P                  4      tRR R lltR R ltR	 R
 ltRR R lltRR R lltR# )zJPosition encodings: Sinusoidal 2D and Rotary Position Embeddings for SAM3.N)ÚOptionalÚTuplec                   ó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	# )ÚPositionEmbeddingSinezSSinusoidal 2D position embedding (used in DETR encoder/decoder and memory encoder).c          	      óB   <€ V ^8„  d   QhRS[ RS[RS[RS[S[,          /# )é   Únum_pos_featsÚtemperatureÚ	normalizeÚscale)ÚintÚfloatÚboolr   )ÚformatÚ__classdict__s   "€Úm/Volumes/fast/ai/experiments/ui-tars-smoke/.venv/lib/python3.14/site-packages/mlx_vlm/models/sam3/position.pyÚ__annotate__Ú"PositionEmbeddingSine.__annotate__   s=   ø€ ÷ Añ AáðAñ ðAñ ð	Añ
 ™�ñAó    c                ó˜   <€ \         SV `  4        Wn        W n        W0n        Ve	   W@n        R # ^\
        P                  ,          V n        R # )N)ÚsuperÚ__init__r   r	   r
   ÚmathÚpir   )Úselfr   r	   r
   r   Ú	__class__s   &&&&&€r   r   ÚPositionEmbeddingSine.__init__   s9   ø€ ô 	‰ÑÔØ*ÔØ&ÔØ"ŒØ#Ò/�UŽ
°Q¼¿¹µ[ˆŽ
r   c                óN   <€ V ^8„  d   QhRS[ P                  RS[ P                  /# ©r   ÚxÚreturn©ÚmxÚarray)r   r   s   "€r   r   r      s#   ø€ ÷ $ñ $™"Ÿ(™(ð $¡r§x¡xñ $r   c                óœ  € VP                   w  r#rE\        P                  ! \        P                  ! V4      ^,           P	                  ^V^4      W#V34      P                  \        P                  4      p\        P                  ! \        P                  ! V4      ^,           P	                  ^^V4      W#V34      P                  \        P                  4      pV P                  '       d[   RpWfRRR1R3,          V,           ,          V P                  ,          pWwRRRR13,          V,           ,          V P                  ,          p\        P                  ! V P                  4      P                  \        P                  4      p	V P                  ^V	^,          ,          V P                  ,          ,          p	VR,          V	,          p
VR,          V	,          p\        P                  ! \        P                  ! V
R,          4      \        P                  ! V
R,          4      .RR7      p
V
P                  ! . V
P                   RR	 ORN5!  p
\        P                  ! \        P                  ! VR,          4      \        P                  ! VR,          4      .RR7      pVP                  ! . VP                   RR	 ORN5!  p\        P                  ! Wº.RR7      pV# )
z�
Args:
    x: (B, H, W, C) feature map in MLX channel-last format
Returns:
    pos: (B, H, W, num_pos_feats*2) position encoding
g�íµ ÷Æ°>ºNNNN©Úaxiséÿÿÿÿ).N©.:é    Nr   ©.:é   Nr   éþÿÿÿ)Úshaper"   Úbroadcast_toÚarangeÚreshapeÚastypeÚfloat32r
   r   r   r	   ÚstackÚsinÚcosÚconcatenate)r   r   ÚBÚHÚWÚ_Úy_embedÚx_embedÚepsÚdim_tÚpos_xÚpos_yÚposs   &&           r   Ú__call__ÚPositionEmbeddingSine.__call__   s  € ð —W‘W‰
ˆˆaô —/’/Ü�YŠY�q‹\˜AÕ×&Ñ& q¨!¨QÓ/°!¸°ó
ç
‰&”—‘Ó
ð 	ô —/’/Ü�YŠY�q‹\˜AÕ×&Ñ& q¨!¨QÓ/°!¸°ó
ç
‰&”—‘Ó
ð 	ð �>�>ˆ>ØˆCØ¨¨B©C°¨Õ!3°cÕ!9Õ:¸T¿Z¹ZÕGˆGØ¨¨A¨r©s¨Õ!3°cÕ!9Õ:¸T¿Z¹ZÕGˆGä—	’	˜$×,Ñ,Ó-×4Ñ4´R·Z±ZÓ@ˆØ× Ñ  Q¨%°1­*Õ%5¸×8JÑ8JÕ%JÕKˆà˜	Õ" UÕ*ˆØ˜	Õ" UÕ*ˆô —’œ"Ÿ&š&  yÕ!1Ó2´B·F²F¸5ÀÕ;KÓ4LÐMÐTVÔWˆØ—’Ð4˜uŸ{™{¨3¨BÐ/Ð4°Ó4ˆÜ—’œ"Ÿ&š&  yÕ!1Ó2´B·F²F¸5ÀÕ;KÓ4LÐMÐTVÔWˆØ—’Ð4˜uŸ{™{¨3¨BÐ/Ð4°Ó4ˆä�nŠn˜e˜^°"Ô5ˆØˆ
r   )r
   r   r   r	   )é   ç     ˆÃ@TN)
Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   rC   Ú__static_attributes__Ú__classdictcell__Ú__classcell__)r   r   s   @@r   r   r   
   s#   ù‡ € Ù]÷Aõ A÷$÷ $ð $r   r   c                óž   € V ^8„  d   QhR\         R\         R\         R\        R\        \        P                  \        P                  3,          /# )r   ÚdimÚend_xÚend_yÚthetar    ©r   r   r   r"   r#   )r   s   "r   r   r   A   sL   € ÷ .ñ .Ü	ð.äð.ô ð.ô ð	.ô
 Œ2�8‰8”R—X‘XÐÕñ.r   c                ó¼  € RV\         P                  ! ^ V ^4      P                  \         P                  4      V ,          ,          ,          p\         P                  ! W,          4      pWQ,          P                  \         P                  4      pWQ,          P                  \         P                  4      pVR,          VR,          ,          pVR,          VR,          ,          p	\         P                  ! W‰.RR7      p
\         P
                  ! Wª.RR7      P                  V
P                  ^ ,          R4      p
\         P                  ! V
4      \         P                  ! V
4      3# )z´Compute 2D axial rotary position embeddings matching HF Sam3ViTRotaryEmbedding.

Returns:
    cos: (end_x*end_y, dim) cosine embeddings
    sin: (end_x*end_y, dim) sine embeddings
ç      ð?r&   )r%   N)Nr%   r(   )
r"   r0   r2   r3   r7   r4   r1   r.   r6   r5   )rP   rQ   rR   rS   ÚfreqsÚflat_idxÚx_positionsÚy_positionsÚfreqs_xÚfreqs_yÚinv_freqs   &&&&       r   Úcompute_axial_cisr^   A   sú   € ð �5œRŸYšY q¨#¨qÓ1×8Ñ8¼¿¹ÓDÀsÕJÕKÕL€Eô �yŠy˜�Ó'€HØÕ#×+Ñ+¬B¯J©JÓ7€KØÕ$×,Ñ,¬R¯Z©ZÓ8€Kð ˜'Õ" U¨7¥^Õ3€GØ˜'Õ" U¨7¥^Õ3€Gô �~Š~˜wÐ0°rÔ:€Hô �xŠx˜Ð,°2Ô6×>Ñ>¸x¿~¹~ÈaÕ?PÐRTÓU€Hä�6Š6�(ÓœRŸVšV HÓ-Ð-Ð-r   c                óX   € V ^8„  d   QhR\         P                  R\         P                  /# r   r!   )r   s   "r   r   r   b   s"   € ÷ 4ñ 4”r—x‘xð 4¤B§H¡Hñ 4r   c                óè   € V P                   ! . V P                  RR ORN^N5!  p V R,          pV R,          p\        P                  ! V) V.RR7      pVP                   ! . VP                  RR ORN5!  # )z;Pairwise rotation: (x0,x1,x2,x3,...) -> (-x1,x0,-x3,x2,...)Nr&   r(   ).r*   ).r,   r-   )r1   r.   r"   r4   )r   Úx1Úx2Úrotateds   &   r   Úrotate_pairwiserd   b   sq   € à	�	Š	Ð'�1—7‘7˜3˜B�<Ð' Ð' QÓ'€AØ	
ˆ6�€BØ	
ˆ6�€BÜ�hŠh˜˜˜R�y rÔ*€GØ�?Š?Ð3˜GŸM™M¨#¨2Ð.Ð3°Ó3Ð3r   c                óî   € V ^8„  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\        \         P                  \         P                  3,          /# )r   ÚxqÚxkr6   r5   r    )r"   r#   r   )r   s   "r   r   r   k   s\   € ÷ ñ Ü
�‰ðä
�‰ðô 
�‰ðô 
�‰ð	ô
 Œ2�8‰8”R—X‘XÐÕñr   c                ó„   € W,          \        V 4      V,          ,           pW,          \        V4      V,          ,           pWE3# )aG  Apply 2D rotary position encoding matching HF implementation.

Formula: q_out = q * cos + rotate_pairwise(q) * sin

Args:
    xq: (B, H, N, D) queries (already transposed for SDPA)
    xk: (B, H, N, D) keys
    cos: (N, D) cosine embeddings
    sin: (N, D) sine embeddings
Returns:
    xq_out, xk_out: rotated queries and keys
)rd   )rf   rg   r6   r5   Úxq_outÚxk_outs   &&&&  r   Úapply_rotary_encrk   k   s8   € ð$ �Xœ¨Ó+¨cÕ1Õ1€FØ�Xœ¨Ó+¨cÕ1Õ1€FØˆ>Ðr   c                óú   € V ^8„  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\        R\        \         P                  \         P                  3,          /# )r   rf   rg   Ú	freqs_cosÚ	freqs_sinÚrepeat_freqs_kr    )r"   r#   r   r   )r   s   "r   r   r   ‚   sf   € ÷ ,ñ ,Ü
�‰ð,ä
�‰ð,ô �x‰xð,ô �x‰xð	,ô
 ð,ô Œ2�8‰8”R—X‘XÐÕñ,r   c                óà  € VRRV P                   ^,          1RR3,          pVRRV P                   ^,          1RR3,          pV'       dˆ   VP                   ^,          pVP                   ^ ,          pWx,           ^,
          V,          p	\        P                  ! W)^34      RRV1RR3,          p
\        P                  ! W9^34      RRV1RR3,          pM@VRRVP                   ^,          1RR3,          p
VRRVP                   ^,          1RR3,          pV R,          V R,          rÜVR,          VR,          rþWÅ,          WÖ,          ,
          pWÆ,          WÕ,          ,           pWê,          Wû,          ,
          pWë,          Wú,          ,           p\        P                  ! VV.RR7      P	                  V P                   4      p\        P                  ! VV.RR7      P	                  VP                   4      pVV3# )a,  Apply 1D RoPE for tracker memory attention (RoPEAttention).

Args:
    xq: (B, N_q, H, D) queries
    xk: (B, N_k, H, D) keys
    freqs_cos: (N, D//2) cosine frequencies
    freqs_sin: (N, D//2) sine frequencies
    repeat_freqs_k: if True, tile freqs to match key length
Returns:
    xq_out, xk_out
Nr%   r&   r)   r+   r(   )r.   r"   Útiler4   r1   )rf   rg   rm   rn   ro   Úcos_qÚsin_qÚN_kÚN_fÚrepeatsÚcos_kÚsin_kÚxq_rÚxq_iÚxk_rÚxk_iÚxq_out_rÚxq_out_iÚxk_out_rÚxk_out_iri   rj   s   &&&&&                 r   Úapply_rotary_enc_1dr�   ‚   s   € ð& �d˜M˜bŸh™h q�k˜M¨4°Ð2Õ3€EØ�d˜M˜bŸh™h q�k˜M¨4°Ð2Õ3€EçØ�h‰h�q�kˆØ�o‰o˜aÕ ˆØ•9˜q•= SÕ(ˆÜ—’˜	¨Q <Ó0°°t¸°t¸TÀ1Ð1DÕEˆÜ—’˜	¨Q <Ó0°°t¸°t¸TÀ1Ð1DÕE‰à˜$  "§(¡(¨1¥+ ¨t°QÐ6Õ7ˆØ˜$  "§(¡(¨1¥+ ¨t°QÐ6Õ7ˆà�I•  9¥ˆ$Ø�I•  9¥ˆ$à�|˜d�lÕ*€HØ�|˜d�lÕ*€HØ�|˜d�lÕ*€HØ�|˜d�lÕ*€Hô �XŠX�x Ð*°Ô4×<Ñ<¸R¿X¹XÓF€FÜ�XŠX�x Ð*°Ô4×<Ñ<¸R¿X¹XÓF€Fà�6ˆ>Ðr   c                óž   € V ^8„  d   QhR\         R\         R\         R\        R\        \        P                  \        P                  3,          /# )r   rP   Úfeat_hÚfeat_wrS   r    rT   )r   s   "r   r   r   ±   sL   € ÷ 0ñ 0Ü	ð0äð0ô ð0ô ð	0ô
 Œ2�8‰8”R—X‘XÐÕñ0r   c                óÎ  € V ^,          pRV\         P                  ! ^ V^4      P                  \         P                  4      V,          ,          ,          p\         P                  ! V4      P                  \         P                  4      p\         P                  ! V4      P                  \         P                  4      p\         P                  ! WgRR7      w  r‰VP                  R4      pV	P                  R4      p	\         P                  ! W…4      p
\         P                  ! W•4      p\         P                  ! W«.RR7      p\         P                  ! V4      \         P                  ! V4      3# )z�Initialize 2D RoPE frequencies for memory attention.

Returns:
    freqs_cos: (feat_h*feat_w, dim//2)
    freqs_sin: (feat_h*feat_w, dim//2)
rV   Úij)Úindexingr&   r(   )
r"   r0   r2   r3   Úmeshgridr1   Úouterr7   r6   r5   )rP   rƒ   r„   rS   ÚhalfrW   Út_yÚt_xÚgrid_yÚgrid_xr\   r[   Ú	freqs_alls   &&&&         r   Úinit_2d_freqsr�   ±   só   € ð �!�8€DØ�5œRŸYšY q¨$°Ó2×9Ñ9¼"¿*¹*ÓEÈÕLÕMÕN€Eä
�)Š)�FÓ
×
"Ñ
"¤2§:¡:Ó
.€CÜ
�)Š)�FÓ
×
"Ñ
"¤2§:¡:Ó
.€Cä—[’[ °DÔ9�N€FØ�^‰^˜BÓ€FØ�^‰^˜BÓ€Fä�hŠh�vÓ%€GÜ�hŠh�vÓ%€Gô —’ Ð1¸Ô;€Iä�6Š6�)ÓœbŸfšf YÓ/Ð/Ð/r   )rF   )F)rK   r   Útypingr   r   Úmlx.coreÚcorer"   Úmlx.nnÚnnÚModuler   r^   rd   rk   r�   r�   © r   r   Ú<module>r˜      sE   ðÙ Pã ß "å Ý ô4˜BŸI™Iô 4÷n.õB4õ÷.,÷^0ñ 0r   