ó
    pyüiB9  ã                   ó”  • S r SSKJrJr  SSKrSSKJr  SSKJrJ	r	  SSK
JrJrJrJr  \" S5      r\" 5       (       a   SS	KJr  SS
KJrJrJr  \(       a  SSKJr  OSr\	R.                  " \5      r " S S5      rS\S\\\\S   -  4   4S jr S%S\R>                  S\R>                  S\R>                  S\R>                  \ \R>                  \R>                  4   -  4S jjr!\R>                  \"-  r#     S&S\R>                  S\"S-  S\ \#\#4   S-  S\S-  SS4
S jjr$S\R>                  S\"S\R>                  4S jr%   S'S\RL                  RN                  S\R>                  S\R>                  S\R>                  S \\R>                  S4   S!\(S-  S"\(S-  S#\R>                  S-  S\ \R>                  \R>                  S-  4   4S$ jjr)g)(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                   ó|   ^ • \ rS rSrSrSrSrSrU 4S jr\	R                  R                  SS9S 5       rS rS	rU =r$ )
ÚWrappedFlexAttentioné;   z`
We are doing a singleton class so that flex attention is compiled once when it's first called.
NFc                 ó^   >• U R                   c  [        TU ]	  U 5      U l         U R                   $ ©N)Ú	_instanceÚsuperÚ__new__)ÚclsÚargsÚkwargsÚ	__class__s      €Úe/home/mande/repo/quber/.venv/lib/python3.13/site-packages/transformers/integrations/flex_attention.pyr   ÚWrappedFlexAttention.__new__D   s'   ø€ Ø�=‰=Ñ ä!™G™O¨CÓ0ˆCŒMØ�}‰}Ðó    )Ú	recursivec                 ó¢  • U R                   (       a  XR                  :w  a¯  Xl        [        S5      (       a  [        R                  " [
        SS9U l        Or[        R                  " [        5       5      R                  S:X  a'  U(       a   [        R                  " [
        SSS9U l        O[        R                  " [
        5      U l        SU l         gg)	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                 ó   • U R                   $ 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.   r1   Ú__static_attributes__Ú__classcell__)r   s   @r   r   r   ;   sP   ø† ñð €IØÐØ#Ðõð ‡^�^×Ñ eÐÐ,ñ*ó -ð*÷*-ð -r    r   Ú
return_lseÚreturnr   c                 óJ   • [         (       a  SU (       a
  [        SS90$ S0$ SU 0$ )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_kwargsrB   d   s/   € ÷ ÒØ¶jœj¨TÑ2ÐKÐKÀdÐKÐKà˜*Ð%Ð%r    ÚqueryÚkeyÚvaluec                 ób   • [        5       (       d  [        U5      " 5       O[        nU" U UU40 UD6$ r   )r   r   r   )rC   rD   rE   r'   r   Úflex_attention_compileds         r   Úcompile_friendly_flex_attentionrH   r   s@   € ô G_×F`ÑF`Ô2°8Ô<Ô>ÔftÐÙ"ØØØñð ñ	ð r    Úattention_mask_2dÚattention_chunk_sizeÚoffsetsÚ	is_causalr   c                 óh  ^ ^^^^^^• T R                   u  pgU(       d  UnU(       d  UnU[        -  S-   [        -  n[        R                  R                  R                  T SSXƒ-
  4S9m T R                  n	T R                  5       mUb4  TR                  5       R                  S5      R                  S5      S-
  U-  mU U4S jmUU4S jn
U U4S jnU(       d  UmOUc  TOU
mUb1  US   R                  U	5      mUS   R                  U	5      mUUU4S	 jnOTn[        UUSUUU	[        S
5      (       + S9$ )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
é   r   )rE   ÚpadNéÿÿÿÿc                 óJ   >• X#:¬  nT	X4   T	X4   :H  nTX4   S:„  nXF-  U-  nU$ )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.
r   © )
Ú	batch_idxÚhead_idxÚq_idxÚkv_idxÚcausal_maskÚdocument_maskÚpadding_maskÚ
final_maskrI   Ú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                 ó8   >• TX4   TX4   :H  nT" XX#5      nXE-  $ )zE
Combines the chunk mask with the causal mask for chunked attention.
rR   )rS   rT   rU   rV   Ú
chunk_maskÚcausal_doc_maskr\   Ú
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                 ó<   >• TX4   TX4   :H  nTX4   S:„  nXT-  nU$ )zX
Utilizes default attention mask to enable encoder and encoder-decoder
attention masks.
r   rR   )	rS   rT   rU   rV   rX   rY   rZ   rI   r[   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                 ó*   >• UT-   nUT-   nT" XXE5      $ r   rR   )	rS   rT   rU   rV   Ú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#   )rm   ÚBÚHÚQ_LENÚKV_LENÚdeviceÚ_compile)ÚshapeÚflex_default_block_sizer(   ÚnnÚ
functionalrO   rs   ÚcloneÚfill_ÚcumsumÚtor   r   )rI   rJ   Úquery_lengthÚ
key_lengthrK   rL   Ú
batch_sizeÚtotal_seq_lenÚpad_lenrs   rb   re   rm   r\   ra   r[   rj   rk   rl   s   `            @@@@@@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    Úhidden_statesÚn_repc                 ó    • U R                   u  p#pEUS:X  a  U $ U SS2SS2SSS2SS24   R                  X#XU5      n U R                  X#U-  XE5      $ )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)
rN   N)ru   ÚexpandÚreshape)rƒ   r„   ÚbatchÚnum_key_value_headsÚslenÚhead_dims         r   Ú	repeat_kvrŒ   ú   s_   € ð
 2?×1DÑ1DÑ.€E Ø�ƒzØÐØ!¢!¢Q¨ªa²Ð"2Ñ3×:Ñ:¸5ÐW\ÐdlÓm€MØ× Ñ  ¸eÑ(CÀTÓTÐTr    ÚmoduleÚattention_maskÚscalingÚsoftcapÚs_auxc           
      ó¶  ^^• UR                  SS5      S:”  a  [        S5      eS n	S m[        U[        5      (       a  Un	OUmTb  TS S 2S S 2S S 2S UR                  S   24   mUU4S jn
SnUR                  S   nXÌS-
  -  S:w  aR  [        X!R                  S   UR                  S   -  5      n[        X1R                  S   UR                  S   -  5      nS	nUR                  S
5      nUR                  R                  S:g  nU(       d  Ub  [        S5      e[        UUU4U
U	UUUU R                  S.[        U5      D6nU(       aí  [        (       a  Uu  nnUR                  nOUu  nnUR                  UR                  5      nUb¬  UR                  u  nnnnUR                  SSSS5      R!                  UUUS5      nUR#                  S5      n[$        R&                  " [$        R(                  " UU/SS9SSS9n[$        R*                  " UU-
  5      nUU-  nUR                  UR                  5      nOUnS nUR-                  SS5      R/                  5       nUU4$ )NÚdropoutg        r   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.éþÿÿÿc                 ón   >• Tb  T[         R                  " U T-  5      -  n Tb  U TU   S   U   U   -   n U $ )Nr   )r(   Útanh)ÚscorerS   rT   rU   rV   Ú
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    TrN   FÚkernel_optionsÚcpuzhAttention sinks cannot be run on CPU with flex attention. Please switch to a different device, e.g. CUDA)r™   Ú
block_maskÚ
enable_gqaÚscaler›   r'   rP   )Údim)r    Úkeepdimr   )ÚgetÚ
ValueErrorÚ
isinstancer   ru   rŒ   rs   ÚtyperH   r'   rB   rA   r@   r|   ÚdtypeÚviewr†   Ú	unsqueezer(   Ú	logsumexpÚcatÚexpÚ	transposeÚ
contiguous)r�   rC   rD   rE   rŽ   r�   r�   r‘   r   r�   r™   rž   Únum_local_query_headsr›   r<   Úflex_attention_outputÚattention_outputÚauxr@   r   Ú	num_headsÚ	seq_len_qÚ_ÚsinksÚlse_expandedÚcombined_lseÚrenorm_factorr˜   s         `                    @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)*r7   Útypingr   r   r(   Ú	packagingr   Úutilsr   r   Úutils.import_utilsr	   r
   r   r   rA   Ú!torch.nn.attention.flex_attentionr   rv   r   r   r   r   Ú
get_loggerr3   Úloggerr   ÚboolÚdictÚstrrB   ÚTensorÚtuplerH   ÚintÚOffsetr‚   rŒ   rw   ÚModuleÚfloatr¹   rR   r    r   Ú<module>rÊ      s=  ðñ÷8 #ã Ý ç 9÷ó ñ 0°Ó8Ð ñ  ×!Ñ!Ýgß^Ñ^æÞ@àˆ
ð 
×	Ò	˜HÓ	%€÷&-ñ &-ðR&¨dð &°t¸CÀÈÐQ]ÑH^ÑA^Ð<^Ñ7_ô &ð$ ñ	Ø�<‰<ðà	�‰ðð �<‰<ðð ‡\�\�E˜%Ÿ,™,¨¯©Ð4Ñ5Ñ5õð$ 
�‰˜Ñ	€ð (,ØØØ,0Ø!ñoØ—|‘|ðoà ™*ðoð
 �6˜6�>Ñ" TÑ)ðoð �d‰{ðoð õoðd	U˜UŸ\™\ð 	U°#ð 	U¸%¿,¹,ô 	Uð$ !Ø Ø!%ñg!Ø�H‰H�O‰Oðg!à�<‰<ðg!ð 
�‰ðg!ð �<‰<ð	g!ð
 ˜%Ÿ,™,¨Ð3Ñ4ðg!ð �T‰\ðg!ð �T‰\ðg!ð �<‰<˜$Ñðg!ð ˆ5�<‰<˜Ÿ™¨Ñ,Ð,Ñ-ög!r    