+
    QV-jB9  ã                   óh  € R t ^ RIHtHt ^ RIt^ RIHt ^RIHtH	t	 ^RI
HtHtHtHt ]! R4      t]! 4       '       d"   ^ RIHt ^ RIHtHtHt ]'       d   ^ R	IHt MRt]	P.                  ! ]4      t ! R
 R4      tR R ltRR R llt]P:                  ],          tRR R llt R R lt!RR R llt"R# )a7  
Partially inspired by torchtune's flex attention implementation

Citation:
@software{torchtune,
  title = {torchtune: PyTorch's finetuning library},
  author = {torchtune maintainers and contributors},
  url = {https//github.com/pytorch/torchtune},
  license = {BSD-3-Clause},
  month = apr,
  year = {2024}
}
)ÚOptionalÚUnionN)Úversion)Úis_torch_flex_attn_availableÚlogging)Úget_torch_versionÚis_torch_greater_or_equalÚis_torch_less_or_equalÚis_torchdynamo_compilingz2.9.0)Ú_DEFAULT_SPARSE_BLOCK_SIZE)Ú	BlockMaskÚcreate_block_maskÚflex_attention)Ú
AuxRequestc                   óŒ   a a€ ] tR t^;t oRtRtRtRtV 3R lt]	P                  P                  RR7      R 4       tR tRtVtV ;t# )	ÚWrappedFlexAttentionz`
We are doing a singleton class so that flex attention is compiled once when it's first called.
NFc                ó`   <€ V P                   f   \        SV `	  V 4      V n         V P                   # ©N)Ú	_instanceÚsuperÚ__new__)ÚclsÚargsÚkwargsÚ	__class__s   &*,€Úy/Volumes/fast/ai/experiments/ui-tars-smoke/.venv/lib/python3.14/site-packages/transformers/integrations/flex_attention.pyr   ÚWrappedFlexAttention.__new__D   s'   ø€ Ø�=‰=Ò ä!™G™O¨CÓ0ˆCŒMØ�}‰}Ðó    )Ú	recursivec                ó¼  € V P                   '       d   WP                  8w  dº   Wn        \        R4      '       d#   \        P                  ! \
        RR7      V n        Mw\        P                  ! \        4       4      P                  R8X  d,   V'       d$   \        P                  ! \
        RRR7      V n        M\        P                  ! \
        4      V n        RV n         R# R# )	z.
Initialize or update the singleton instance.
ú2.5.1F)Údynamicz2.6.0zmax-autotune-no-cudagraphs)r!   ÚmodeTN)Ú_is_flex_compiledÚtrainingr	   ÚtorchÚcompiler   Ú_compiled_flex_attentionr   Úparser   Úbase_version)Úselfr$   s   &&r   Ú__init__ÚWrappedFlexAttention.__init__J   s—   € ð
 ×%×%Ð%¨·]±]Ô)BØ$ŒMÜ% g×.Ò.Ü05·²¼nÐV[Ô0\�Õ-ô —’Ô0Ó2Ó3×@Ñ@ÀGÔK×PXÜ05·²Ü"¨EÐ8Tô1�Õ-ô
 16·²¼nÓ0M�Ô-à%)ˆDÖ"ñ *Cr   c                ó   € V P                   # r   )r'   )r*   s   &r   Ú__call__ÚWrappedFlexAttention.__call__`   s   € Ø×,Ñ,Ð,r   )r'   r#   r$   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   r#   r'   r   r%   ÚcompilerÚdisabler+   r.   Ú__static_attributes__Ú__classdictcell__Ú__classcell__)r   Ú__classdict__s   @@r   r   r   ;   sR   ù‡ € ñð €IØÐØ#Ðõð ‡^�^×Ñ eÐÓ,ñ*ó -ð*÷*-ò -r   r   c                óv   € V ^8„  d   QhR\         R\        \        \         \        R,          ,          3,          /# )é   Ú
return_lseÚreturnr   )ÚboolÚdictÚstrr   )Úformats   "r   Ú__annotate__rC   d   s/   € ÷ &ñ &¬dð &´t¼CÄÌÐQ]ÕH^ÕA^Ð<^Õ7_ñ &r   c                óT   € \         '       d   RV '       d   \        RR7      /# R/# RV /# )aA  
Requests the LSE from flex_attention in a version-agnostic fashion.

Before torch 2.9, the LSE was requested via the boolean return_lse field. However, starting with
torch 2.9, an AuxRequest object must be passed via the aux_request field. This method conditionally
returns the correct form based on the python version.
Ú
return_auxT)ÚlseNr=   )Ú_TORCH_FLEX_USE_AUXr   )r=   s   &r   Úget_flex_attention_lse_kwargsrH   d   s/   € ÷ ÓØ·jœj¨TÔ2ÐKÐKÀdÐKÐKà˜*Ð%Ð%r   c                óø   € V ^8„  d   QhR\         P                  R\         P                  R\         P                  R\         P                  \        \         P                  \         P                  3,          ,          /# )r<   ÚqueryÚkeyÚvaluer>   )r%   ÚTensorÚtuple)rB   s   "r   rC   rC   r   sW   € ÷ ñ Ü�<‰<ðä	�‰ðô �<‰<ðô ‡\�\”Eœ%Ÿ,™,¬¯©Ð4Õ5Õ5ñr   c                 ód   € \        4       '       g   \        V4      ! 4       M\        pV! V VV3/ VB # r   )r
   r   r   )rJ   rK   rL   r$   r   Úflex_attention_compileds   &&&&, r   Úcompile_friendly_flex_attentionrQ   r   s@   € ô G_×F`ÒF`Ô2°8Ô<Ô>ÔftÐÙ"ØØØñð ñ	ð r   c          
      ó¬   € V ^8„  d   QhR\         P                  R\        R,          R\        \        \        3,          R,          R\
        R,          RR/# )r<   Úattention_mask_2dÚattention_chunk_sizeNÚoffsetsÚ	is_causalr>   r   )r%   rM   ÚintrN   ÚOffsetr?   )rB   s   "r   rC   rC   ˆ   s[   € ÷ oñ oÜ—|‘|ðoä �*ðoô
 ”6œ6�>Õ" TÕ)ðoô �d�{ðoð ñor   c                óÀ  a aaaaaa€ S P                   w  rgV'       g   TpV'       g   TpV\        ,          ^,           \        ,          p\        P                  P                  P                  S ^ ^ Wƒ,
          3R7      o S P                  p	S P                  4       oVe=   SP                  4       P                  ^4      P                  R	4      ^,
          V,          oV V3R loVV3R lp
V V3R lpV'       g   VoMVf   SMT
oVe:   V^ ,          P                  V	4      oV^,          P                  V	4      oVVV3R lpMSp\        VVRVVV	\        R4      '       * R7      # )
a÷  
IMPORTANT NOTICE: This function is deprecated in favor of using the mask primitives in `masking_utils.py`,
and will be removed in a future version without warnings. New code should not use it. It is only kept here
for BC for now, while models using it are being patched accordingly.

Create a block (causal) document mask for a batch of sequences, both packed and unpacked.
Create Block (causal) logic and passing it into :func:`torch.nn.attention.flex_attention.create_block_mask`.
The resultant BlockMask is a compressed representation of the full (causal) block
mask. BlockMask is essential for performant computation of flex attention.
See: https://pytorch.org/blog/flexattention/

Args:
    attention_mask_2d (torch.Tensor): Attention mask for packed and padded sequences
    of shape (batch_size, total_seq_len). e.g.

    For unpacked sequence:
    [[1, 1, 1, 1, 0, 0, 0],
     [1, 1, 1, 1, 1, 0, 0]]

    For packed sequence:
    [[1, 1, 1, 2, 2, 2, 0],
     [1, 1, 2, 2, 2, 3, 3]]

Returns:
    BlockMask
)rL   ÚpadNc                ór   <€ W#8¬  pS	W3,          S	W3,          8H  pSW3,          ^ 8„  pWF,          V,          pV# )zÔ
Defines the logic of a block causal mask by combining both a standard causal mask
and a block diagonal document mask.
See :func:`~torchtune.modules.attention_utils.create_block_causal_mask`
for an illustration.
© )
Ú	batch_idxÚhead_idxÚq_idxÚkv_idxÚcausal_maskÚdocument_maskÚpadding_maskÚ
final_maskrS   Údocument_idss
   &&&&    €€r   Úcausal_mask_modÚ4make_flex_block_causal_mask.<locals>.causal_mask_mod¾   sK   ø€ ð ‘oˆØ$ YÐ%5Õ6¸,ÀyÐGXÕ:YÑYˆØ(¨Ð)9Õ:¸QÑ>ˆØ Õ/°-Õ?ˆ
ØÐr   c                óP   <€ SW3,          SW3,          8H  pS! WW#4      pWE,          # )zE
Combines the chunk mask with the causal mask for chunked attention.
r\   )r]   r^   r_   r`   Ú
chunk_maskÚcausal_doc_maskrf   Ú
chunk_idxss   &&&&  €€r   Úchunk_causal_mask_modÚ:make_flex_block_causal_mask.<locals>.chunk_causal_mask_modË   s4   ø€ ð   	Ð 0Õ1°ZÀ	Ð@QÕ5RÑRˆ
Ù)¨)¸uÓMˆØÕ+Ð+r   c                ó\   <€ SW3,          SW3,          8H  pSW3,          ^ 8„  pWT,          pV# )zX
Utilizes default attention mask to enable encoder and encoder-decoder
attention masks.
r\   )	r]   r^   r_   r`   rb   rc   rd   rS   re   s	   &&&&   €€r   Údefault_mask_modÚ5make_flex_block_causal_mask.<locals>.default_mask_modÓ   s?   ø€ ð
 % YÐ%5Õ6¸,ÀyÐGXÕ:YÑYˆà(¨Ð):Õ;¸aÑ?ˆØ!Õ1ˆ
ØÐr   c                 ó:   <€ VS,           pVS,           pS! WWE4      # r   r\   )	r]   r^   r_   r`   Úoffset_qÚ	offset_kvÚ	kv_offsetÚmask_mod_maybe_combinedÚq_offsets	   &&&&  €€€r   Úmask_modÚ-make_flex_block_causal_mask.<locals>.mask_modç   s$   ø€ Ø˜xÕ'ˆHØ Õ*ˆIÙ*¨9ÀÓTÐTr   r    )rw   ÚBÚHÚQ_LENÚKV_LENÚdeviceÚ_compileéÿÿÿÿ)ÚshapeÚflex_default_block_sizer%   ÚnnÚ
functionalrZ   r}   ÚcloneÚfill_ÚcumsumÚtor   r	   )rS   rT   Úquery_lengthÚ
key_lengthrU   rV   Ú
batch_sizeÚtotal_seq_lenÚpad_lenr}   rl   ro   rw   rf   rk   re   rt   ru   rv   s   f&&&&&       @@@@@@r   Úmake_flex_block_causal_maskr�   ˆ   sD  þ€ ðD !2× 7Ñ 7Ñ€JßØ"ˆ
ßØ$ˆàÔ5Õ5¸Õ:Ô>UÕU€GÜŸ™×+Ñ+×/Ñ/Ð0AÈÐQRÐT[ÕThÐPiÐ/ÓjÐØ×%Ñ%€FØ$×*Ñ*Ó,€LàÒ'à"×(Ñ(Ó*×0Ñ0°Ó3×:Ñ:¸2Ó>ÀÕBÐH\Õ]ˆ
öö,ö	÷ Ø"2Ñà5IÒ5Q¡/ÐWlÐàÒØ˜1•:—=‘= Ó(ˆØ˜A•J—M‘M &Ó)ˆ	÷	Uð 	Uð
 +ˆäØØ
Ø
ØØØä+¨GÓ4Ô4ô	ð 	r   c                ód   € V ^8„  d   QhR\         P                  R\        R\         P                  /# )r<   Úhidden_statesÚn_repr>   )r%   rM   rW   )rB   s   "r   rC   rC   ú   s.   € ÷ 	Uñ 	UœUŸ\™\ð 	U´#ð 	U¼%¿,¹,ñ 	Ur   c                ó˜   € V P                   w  r#rEV^8X  d   V # V R,          P                  W#WV4      p V P                  W#V,          WE4      # )zÈ
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
)ºNNNr’   Nr’   r’   )r€   ÚexpandÚreshape)r�   r�   ÚbatchÚnum_key_value_headsÚslenÚhead_dims   &&    r   Ú	repeat_kvr™   ú   sU   € ð
 2?×1DÑ1DÑ.€E Ø�„zØÐØ!Ð"2Õ3×:Ñ:¸5ÐW\ÐdlÓm€MØ× Ñ  ¸eÕ(CÀTÓTÐTr   c                ó¬  € V ^8„  d   QhR\         P                  P                  R\         P                  R\         P                  R\         P                  R\        \         P                  R3,          R\
        R,          R	\
        R,          R
\         P                  R,          R\        \         P                  \         P                  R,          3,          /	# )r<   ÚmodulerJ   rK   rL   Úattention_maskr   ÚscalingNÚsoftcapÚs_auxr>   )r%   r‚   ÚModulerM   r   ÚfloatrN   )rB   s   "r   rC   rC     s·   € ÷ g!ñ g!Ü�H‰H�O‰Oðg!ä�<‰<ðg!ô 
�‰ðg!ô �<‰<ð	g!ô
 œ%Ÿ,™,¨Ð3Õ4ðg!ô �T�\ðg!ô �T�\ðg!ô �<‰<˜$Õðg!ô Œ5�<‰<œŸ™¨Õ,Ð,Õ-ñg!r   c                 ó:  aa€ VP                  R R4      ^ 8”  d   \        R4      hRp	Ro\        V\        4      '       d   Tp	MVoSe!   SRRRRVP                  R,          13,          oVV3R lp
RpVP                  ^,          pWÌ^,
          ,          ^ 8w  dk   \        W!P                  ^,          VP                  ^,          ,          4      p\        W1P                  ^,          VP                  ^,          ,          4      pRpVP                  R4      pVP                  P                  R	8g  pV'       g   Ve   \        R
4      h\        VVV3RV
RV	RVRVRVRV P                  /\        V4      B pV'       dþ   \        '       d   Vw  ppVP                  pMVw  ppVP                  VP                  4      pVe»   VP                  w  ppppVP                  ^R^^4      P!                  VVV^4      pVP#                  R4      p\$        P&                  ! \$        P(                  ! VV.RR7      RRR7      p\$        P*                  ! VV,
          4      pVV,          pVP                  VP                  4      pMTpRpVP-                  ^^4      P/                  4       pVV3# )Údropoutg        z›`flex_attention` does not support `dropout`. Please use it with inference only (`model.eval()`) or turn off the attention dropout in the respective config.Nr’   c                 óª   <€ Se%   S\         P                  ! V S,          4      ,          p Se&   V SV,          ^ ,          V,          V,          ,           p V # r   )r%   Útanh)Úscorer]   r^   r_   r`   Ú
score_maskrž   s   &&&&&€€r   Ú	score_modÚ)flex_attention_forward.<locals>.score_mod!  sK   ø€ ØÒØœeŸjšj¨°­Ó9Õ9ˆEØÒ!Ø˜J yÕ1°!Õ4°UÕ;¸FÕCÕCˆEð ˆr   TFÚkernel_optionsÚcpuzhAttention sinks cannot be run on CPU with flex attention. Please switch to a different device, e.g. CUDAr¨   Ú
block_maskÚ
enable_gqaÚscaler$   )Údim)r¯   Úkeepdiméþÿÿÿr   )ÚgetÚ
ValueErrorÚ
isinstancer   r€   r™   r}   ÚtyperQ   r$   rH   rG   rF   r‡   ÚdtypeÚviewr“   Ú	unsqueezer%   Ú	logsumexpÚcatÚexpÚ	transposeÚ
contiguous)r›   rJ   rK   rL   rœ   r�   rž   rŸ   r   r¬   r¨   r­   Únum_local_query_headsrª   r=   Úflex_attention_outputÚattention_outputÚauxrF   rŠ   Ú	num_headsÚ	seq_len_qÚ_ÚsinksÚlse_expandedÚcombined_lseÚrenorm_factorr§   s   &&&&&&f&,                  @r   Úflex_attention_forwardrÉ     sŸ  ù€ ð ‡z�z�)˜SÓ! AÔ%Üðaó
ð 	
ð
 €JØ€JÜ�.¤)×,Ò,Ø#‰
à#ˆ
àÒØ  1 a¨¨3¯9©9°R­=¨Ð 8Õ9ˆ
öð €JØ!ŸK™K¨�NÐð 	¸Õ!:Õ;ÀÔAÜ˜Ÿ[™[¨�^¨s¯y©y¸­|Õ;Ó<ˆÜ˜%§¡¨Q¥°5·;±;¸qµ>Õ!AÓBˆØˆ
à—Z‘ZÐ 0Ó1€Nà—‘×"Ñ" eÑ+€Jç˜%Ò+ÜØvó
ð 	
ô <ØØØñð ð	ð
 ðð ðð ðð &ðð —‘ðô (¨
Ó
3ñÐ÷  ÷ ÓØ$9Ñ!Ð˜cØ—'‘'‰Cà$9Ñ!Ð˜cð �f‰f�U—[‘[Ó!ˆàÒà2B×2HÑ2HÑ/ˆJ˜	 9¨aØ—J‘J˜q " a¨Ó+×2Ñ2°:¸yÈ)ÐUVÓWˆEð
 Ÿ=™=¨Ó,ˆLÜ Ÿ?š?¬5¯9ª9°lÀEÐ5JÐPRÔ+SÐY[ÐeiÔjˆLô "ŸIšI l°\Õ&AÓBˆMØ/°-Õ?ÐØ/×2Ñ2°5·;±;Ó?Ðøà0ÐØˆà'×1Ñ1°!°QÓ7×BÑBÓDÐØ˜SÐ Ð r   )F)NNNNT)NNN)#r4   Útypingr   r   r%   Ú	packagingr   Úutilsr   r   Úutils.import_utilsr   r   r	   r
   rG   Ú!torch.nn.attention.flex_attentionr   r�   r   r   r   r   Ú
get_loggerr0   Úloggerr   rH   rQ   rM   rW   rX   r�   r™   rÉ   r\   r   r   Ú<module>rÑ      sž   ðñ÷8 #ã Ý ç 9÷ó ñ 0°Ó8Ð ñ  ×!Ò!Ýgß^Ñ^çÞ@àˆ
ð 
×	Ò	˜HÓ	%€÷&-ñ &-õR&÷ð$ 
�‰˜Õ	€÷oõd	U÷g!ñ g!r   