+
    QV-jý*  ã                   óø   € ^ RI t ^RIHt ^RIHt ]P
                  ! ]4      tRt^€t	^€t
RsR R lt] P                  P                  RRR7      R R	 l4       t]P                   R
 4       tR R ltR tRR R lltR# )é    N)Úlogging)Úsdpa_attention_forwardc                ó$   € V ^8„  d   QhR\         /# )é   Úattn_implementation)Ústr)Úformats   "Úx/Volumes/fast/ai/experiments/ui-tars-smoke/.venv/lib/python3.14/site-packages/transformers/integrations/msa_attention.pyÚ__annotate__r       s   € ÷ ñ ´cñ ó    c                ó<  € \         e   \         # ^RIHp V P                  R4      R
,          pVP	                  R4      w  r#pT! Y$;'       g    RV'       d   RM^ RR7      pR F0  p\        \        WVR4      4      '       d   K   \        RV RV R	24      h	  Vs \         # )ag  Load the MSA hub kernel once and verify the expected callables are present.

The ``attn_implementation`` string may carry a ``paged|`` prefix and/or an ``@<revision>`` pin
(e.g. ``kernels-staging/msa@v0``); the build currently lives on the repo's ``v0`` branch. The
loaded module is cached in a module-level global so registration happens once, not per call.
N)Ú
get_kernelÚ|Ú@T)ÚrevisionÚversionÚallow_all_kernelszThe MSA kernel loaded from `z` does not expose a callable `zK`. Make sure you request a compatible build, e.g. `kernels-staging/msa@v0`.éÿÿÿÿ)Úsparse_atten_funcÚbuild_k2q_csr)Ú_MSA_KERNELÚhub_kernelsr   ÚsplitÚ	partitionÚcallableÚgetattrÚImportError)r   r   Úrepo_idÚ_ÚrevÚkernelÚfn_names   &      r
   Úload_and_register_msa_kernelr#       s©   € ô ÒÜÐå'à!×'Ñ'¨Ó,¨RÕ0€GØ×'Ñ'¨Ó,�O€G�Ù˜¯+¨+°Çs¹tÐPQÐeiÔj€Fã9ˆÜœ °Ó6×7Ô7ÜØ.¨w¨iÐ7UÐV]ÐU^ð _[ð [óð ñ :ð €KÜÐr   ztransformers_msa::sparse_atten)Úmutates_argsc                óX  € V ^8„  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\        R\        R	\        R
\        R\        R\        R\        R\        R\         P                  /# )r   ÚqÚkÚvÚq2kÚcu_seqlens_qÚcu_seqlens_kÚtopkÚ
block_sizeÚtotal_kÚmax_seqlen_qÚmax_seqlen_kÚqheads_per_kvÚscalingÚimplÚreturn)ÚtorchÚTensorÚintÚfloatr   )r	   s   "r
   r   r   =   sº   € ÷ 5$ñ 5$Ü‡|�|ð5$ä‡|�|ð5$ô ‡|�|ð5$ô 
�‰ð	5$ô
 —,‘,ð5$ô —,‘,ð5$ô ð5$ô ð5$ô ð5$ô ð5$ô ð5$ô ð5$ô ð5$ô ð5$ô ‡\�\ñ5$r   c                ód  € \        V4      p\        P                  P                  V P                  4      ;_uu_ 4        VP	                  VVVVVV
V	VR7      w  ppVP                  V VVVVVVVV	V
VRVR7      pRRR4       VP                  4       #   + '       g   i     XP                  4       # ; i)a  Opaque wrapper around the CuTe-DSL CSR build + block-sparse kernel.

Registered as a ``torch.library`` custom op so ``torch.compile(fullgraph=True)`` treats the
whole CSR-build + attention as a single opaque node (no graph break) and ``reduce-overhead``
CUDA graphs can capture it. The internal ``build_k2q_csr`` output is data-dependent in shape,
but it never escapes this op (only the fixed-shape ``[total_q, Hq, D]`` attention output does),
so the fake/meta impl below is exact. The op is functional (no input mutation).
)r.   r0   r/   Úqhead_per_kvT)r*   r+   r/   r0   Úblk_kvÚcausalÚsoftmax_scaleN)r#   r5   ÚcudaÚdevicer   r   Ú
contiguous)r&   r'   r(   r)   r*   r+   r,   r-   r.   r/   r0   r1   r2   r3   ÚmsaÚk2q_row_ptrÚk2q_q_indicesÚattn_outputs   &&&&&&&&&&&&&&    r
   Ú_msa_sparse_atten_oprE   <   sÉ   € ô2 ' tÓ
,€Cô 
�‰×	Ñ	˜1Ÿ8™8×	$Õ	$Ø%(×%6Ñ%6ØØØØØØ%Ø%Ø&ð &7ó 	&
Ñ"ˆ�]ð ×+Ñ+ØØØØØØØ%Ø%Ø%Ø%ØØØ!ð ,ó 
ˆ÷ 
%ð4 ×!Ñ!Ó#Ð#÷5 
%Ö	$ð4 ×!Ñ!Ó#Ð#ús   ½;BÂB/	c                 ó.   € \         P                  ! V 4      # ©N)r5   Ú
empty_like)r&   r'   r(   r)   r*   r+   r,   r-   r.   r/   r0   r1   r2   r3   s   &&&&&&&&&&&&&&r
   Ú_msa_sparse_atten_fakerI   u   s   € ô$ ×Ò˜AÓÐr   c                óH   € V ^8„  d   QhR\         P                  R\        RR/# )r   ÚqueryÚdropoutr4   N)r5   r6   r8   )r	   s   "r
   r   r   Š   s%   € ÷ 
ñ 
¤e§l¡lð 
¼Uð 
Àtñ 
r   c                ó  € VP                   P                  R8w  g6   \        P                  P	                  VP                   4      ^ ,          ^
8w  d   \        R4      hVP                  R,          \        8w  d   \        R\         R24      hV P                  P                  \        8w  d   \        R\         R24      hVR8w  d   \        R4      hV P                  P                  pV\        9  d   \        R\         R	V R
24      hR# )a  Validate kernel capability, dropout and configured topk once per attention module.

Mirrors the flash-attention integration, which checks capability/dropout at model init rather
than on every forward. The check is cached on the module so the hot path never re-runs it.

There is no SDPA fallback: a sparse layer either runs the MSA kernel or this raises. Serves both
prefill (q_len > 1) and single-token decode (q_len == 1) -- decode is just a varlen call with one
query slot, so there is no context-length threshold.
r>   z‡MSA block-sparse attention requires an SM100 / Blackwell CUDA device. Select a different `attn_implementation` on unsupported hardware.z2MSA block-sparse attention only supports head_dim Ú.z4MSA block-sparse attention only supports block_size ç        zYMSA block-sparse attention does not support attention dropout; set `attention_dropout=0`.z1MSA block-sparse attention only supports topk in z, got `z0`. Set `index_topk_blocks` to a supported value.Nr   )r?   Útyper5   r>   Úget_device_capabilityÚRuntimeErrorÚshapeÚMSA_SUPPORTED_HEAD_DIMÚ
ValueErrorÚindexerr-   ÚMSA_SUPPORTED_BLOCK_SIZEÚtopk_blocksÚMSA_SUPPORTED_TOPK)ÚmodulerK   rL   r,   s   &&& r
   Ú_validate_msa_initr[   Š   sú   € ð ‡|�|×Ñ˜FÔ"¤e§j¡j×&FÑ&FÀuÇ|Á|Ó&TÐUVÕ&WÐ[]Ô&]ÜðPó
ð 	
ð ‡{�{�2…Ô0Ô0ÜÐMÔNdÐMeÐefÐgÓhÐhØ‡~�~× Ñ Ô$<Ô<ÜÐOÔPhÐOiÐijÐkÓlÐlØ�#„~ÜÐtÓuÐuØ�>‰>×%Ñ%€DØÔ%Ô%ÜØ?Ô@RÐ?SÐSZÐ[_ÐZ`ð a<ð <ó
ð 	
ñ &r   c                 ó(  a€ VP                   w  r‰r«VP                   ^,          VP                   ^,          rÜWœ,          pVP                   R,          o\        V3R l\         4       4      pVS8w  dH   VP                  . VP                   RR OVS,
          N5R4      p\        P
                  ! VV.RR7      pVoVP                  ^^4      P                  WŠ,          W›4      P                  4       pVP                  ^^4      P                  W�,          WË4      P                  4       pVP                  ^^4      P                  W�,          WË4      P                  4       p\        P                  ! ^ V^,           V
,          V
VP                  \        P                  R7      pV^8X  d‰   Ve…   VR,          ^,           P                  \        P                  4      P                  ^4      p\        P
                  ! \        P                  ! ^VP                  \        P                  R7      V.4      pMA\        P                  ! ^ V^,           V,          VVP                  \        P                  R7      pVP                  \        P                  4      pVP                  WŠ,          S4      P                  ^ 4      P                  VRR4      P                  4       p\!        VVVVVVSVW�,          V
VVVV P"                  P$                  4      pVP                  WŠW›4      # )é   c              3   ó8   <"  € T F  qS8¼  g   K  Vx € K  	  R # 5irG   © )Ú.0Útr,   s   & €r
   Ú	<genexpr>Ú$_sparse_attention.<locals>.<genexpr>³   s   øé € ÐBÑ"4˜Q¸T¹	—q’qÓ"4ùs   ƒ�
N)Údim)r?   Údtyper   )rS   ÚnextrY   Únew_fullr5   ÚcatÚ	transposeÚreshaper@   Úaranger?   Úint32ÚtoÚzerosÚ	unsqueezeÚexpandrE   ÚconfigÚ_attn_implementation)rZ   rK   ÚkeyÚvaluer2   Úblock_indicesr-   Úcache_positionÚbszÚnum_q_headsÚq_lenÚhead_dimÚnum_kv_headsÚk_lenr1   Úpadded_topkÚpadr&   r'   r(   r*   Úvalid_kr+   r)   rD   r,   s   &&&&&&&&                 @r
   Ú_sparse_attentionr€   §   ss  ø€ Ø(-¯©Ñ%€C�eØŸ)™) A�,¨¯	©	°!­�%ØÕ/€MØ×Ñ˜rÕ"€Dô ÔBÕ"4ÓBÓB€KØ�dÔØ×$Ñ$Ð%T }×':Ñ':¸3¸BÐ'?Ð%TÀÈtÕASÑ%TÐVXÓYˆÜŸ	š	 =°#Ð"6¸BÔ?ˆØˆð
 	�‰˜˜1Ó×%Ñ% c¥k°;ÓI×TÑTÓV€AØ�‰�a˜Ó×#Ñ# C¥K°ÓH×SÑSÓU€AØ�‰˜˜1Ó×%Ñ% c¥k°<ÓJ×UÑUÓW€AÜ—<’<  C¨!¥G¨uÕ#4°eÀAÇHÁHÔTY×T_ÑT_Ô`€Lð ˆa„x�NÒ.Ø! "Õ%¨Õ)×-Ñ-¬e¯k©kÓ:×BÑBÀ1ÓEˆÜ—y’y¤%§+¢+¨a¸¿¹ÌÏÉÔ"TÐV]Ð!^Ó_‰ä—|’| A¨¨a­°5Õ'8¸%ÈÏÉÔX]×XcÑXcÔdˆà
×
Ñ
œ5Ÿ;™;Ó
'€CØ
�+‰+�c•k 4Ó
(×
2Ñ
2°1Ó
5×
<Ñ
<¸\È2ÈrÓ
R×
]Ñ
]Ó
_€Cô 'Ø	Ø	Ø	ØØØØØØ�ØØØØØ�‰×*Ñ*ó€Kð  ×Ñ˜s¨;ÓAÐAr   c                óh  € V ^8„  d   QhR\         P                  P                  R\         P                  R\         P                  R\         P                  R\         P                  R,          R\        R\        R,          R	\         P                  R,          R
\
        \         P                  R3,          /	# )r   rZ   rK   rs   rt   Úattention_maskNrL   r2   ru   r4   )r5   ÚnnÚModuler6   r8   Útuple)r	   s   "r
   r   r   ã   sš   € ÷  ñ  Ü�H‰H�O‰Oð ä�<‰<ð ô 
�‰ð ô �<‰<ð	 ô
 —L‘L 4Õ'ð ô ð ô �T�\ð ô —<‘< $Õ&ð ô Œ5�<‰<˜ÐÕñ r   c           
     ó(  € Vf   VP                   R,          R	,          pVf   \        WW#V3RVRV/VB # \        V RR4      '       g   \        WV4       RV n        V P
                  P                  p	VP                  R4      p
\        WW#WgWš4      pVR3# )
zf
TODO: this opens a door to per-layer attn implementation which is something we might want lalter on.
NrL   r2   Ú_msa_validatedFTrv   r   g      à¿)	rS   r   r   r[   r‡   rV   r-   Úgetr€   )rZ   rK   rs   rt   r‚   rL   r2   ru   Úkwargsr-   rv   rD   s   &&&&&&&&,   r
   Úmsa_attention_forwardrŠ   ã   s­   € ð ‚Ø—+‘+˜b•/ TÕ)ˆð ÒÜ%Ø˜3 ~ñ
Ø?Fð
ØPWð
Ø[añ
ð 	
ô �6Ð+¨U×3Ò3Ü˜6¨'Ô2Ø $ˆÔà—‘×*Ñ*€JØ—Z‘ZÐ 0Ó1€NÜ# F°3¸wÐWaÓr€KØ˜ÐÐr   )é   é   é   é    r_   )NrO   NN)r5   Úutilsr   Úsdpa_attentionr   Ú
get_loggerÚ__name__ÚloggerrY   rW   rT   r   r#   ÚlibraryÚ	custom_oprE   Úregister_fakerI   r[   r€   rŠ   r_   r   r
   Ú<module>r—      s›   ðó å Ý 2ð 
×	Ò	˜HÓ	%€ð $Ð àÐ àÐ à€õð8 ‡�×ÑÐ9ÈÐÓKô5$ó Lð5$ðp ×#Ñ#ñó $ðõ(
ò:9B÷x ñ  r   