Ë
    S^(hL$  ã                   óè  — d Z ddlmZmZmZ ddlZddlmZ ddlm	Z	 ddl
mZ  e	«       rddlmZmZ dd	lmZ  G d
„ d«      Zeej$                  ef   Z	 	 	 	 d!dej$                  dee   deeeef      ddfd„Zej,                  j/                  d¬«      	 d"dej$                  dej$                  dej$                  dej$                  fd„«       Zdej$                  dedej$                  fd„Z	 	 	 d#dej4                  j6                  dej$                  dej$                  dej$                  deej$                  df   dee   dee   deej$                     deej$                  ej$                  f   fd „Zy)$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ÚTupleÚUnionN)Úversioné   )Úis_torch_flex_attn_available)Ú_torch_version)Ú	BlockMaskÚflex_attention)Úcreate_block_maskc                   óx   ‡ — e Zd ZdZdZdZdZˆ fd„Zej                  j                  d¬«      d„ «       Zd„ Zˆ xZS )ÚWrappedFlexAttentionzh
    We are doing a singleton class so that flex attention is compiled once when it's first called.
    NFc                 ó\   •— | j                   €t        ‰| �	  | «      | _         | j                   S ©N)Ú	_instanceÚsuperÚ__new__)ÚclsÚargsÚkwargsÚ	__class__s      €úf/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/integrations/flex_attention.pyr   zWrappedFlexAttention.__new__6   s'   ø€ Ø�=‰=Ð ä!™G™O¨CÓ0ˆCŒMØ�}‰}Ðó    ©Ú	recursivec                 ó(  — | j                   r|| j                  k7  rw|| _        t        j                  t        «      j
                  dk(  r$|r"t        j                  t        dd¬«      | _	        nt        j                  t        «      | _	        d| _         yy)z>
        Initialize or update the singleton instance.
        z2.6.0Fzmax-autotune-no-cudagraphs)ÚdynamicÚmodeTN)
Ú_is_flex_compiledÚtrainingr   Úparser	   Úbase_versionÚtorchÚcompiler   Ú_compiled_flex_attention)Úselfr    s     r   Ú__init__zWrappedFlexAttention.__init__<   st   € ð
 ×%Ò%¨°T·]±]Ò)Bð %ˆDŒMÜ�}‰}œ^Ó,×9Ñ9¸WÒDÉÜ05·±Ü"¨EÐ8Tô1�Õ-ô 16·±¼nÓ0M�Ô-Ø%)ˆDÕ"ð *Cr   c                 ó   — | j                   S r   )r%   )r&   s    r   Ú__call__zWrappedFlexAttention.__call__N   s   € Ø×,Ñ,Ð,r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r%   r   r#   ÚcompilerÚdisabler'   r)   Ú__classcell__)r   s   @r   r   r   -   sK   ø„ ñð €IØÐØ#Ðôð ‡^�^×Ñ eÐÓ,ñ*ó -ð*ö"-r   r   Úattention_mask_2dÚattention_chunk_sizeÚoffsetsÚreturnr
   c           	      óz  ‡ ‡	‡
‡‡— ‰ j                   \  }}|s|}|s|}t        j                  j                  j	                  ‰ dd|f¬«      Š ‰ j
                  }‰ j                  «       Š
|�&‰
j                  d«      j                  d«      dz
  |z  Š
ˆ ˆ
fd„Š	|�|d   Š|d   Šˆ	ˆˆfd„}n‰	}t        ||d|||d¬	«      S )
a  
    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 block causal
    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   )ÚvalueÚpadNé   éÿÿÿÿc                 óT   •— ||k\  }‰	| |f   ‰	| |f   k(  }‰| |f   dkD  }||z  |z  }|S )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_maskr1   Údocument_idss
           €€r   Úcausal_mask_modz4make_flex_block_causal_mask.<locals>.causal_mask_mod„   sV   ø€ ð ˜v‘oˆØ$ Y°Ð%5Ñ6¸,ÀyÐRXÐGXÑ:YÑYˆØ(¨°EÐ)9Ñ:¸QÑ>ˆØ  <Ñ/°-Ñ?ˆ
ØÐr   c                 ó.   •— |‰z   }|‰z   } ‰| |||«      S r   r;   )	r<   r=   r>   r?   Úoffset_qÚ	offset_kvrE   Ú	kv_offsetÚq_offsets	         €€€r   Úmask_modz-make_flex_block_causal_mask.<locals>.mask_mod–   s(   ø€ Ø˜xÑ'ˆHØ Ñ*ˆIÙ" 9¨h¸À)ÓLÐLr   T)rK   ÚBÚHÚQ_LENÚKV_LENÚdeviceÚ_compile)
Úshaper#   ÚnnÚ
functionalr7   rP   ÚcloneÚfill_ÚcumsumÚcreate_block_causal_mask_flex)r1   r2   Úquery_lengthÚ
key_lengthr3   Ú
batch_sizeÚtotal_seq_lenrP   rK   rE   rD   rI   rJ   s   `        @@@@r   Úmake_flex_block_causal_maskr]   U   sâ   ü€ ð: !2× 7Ñ 7Ñ€J�ÙØ"ˆ
ÙØ$ˆÜŸ™×+Ñ+×/Ñ/Ð0AÈÐQRÐT^ÐP_Ð/Ó`ÐØ×%Ñ%€FØ$×*Ñ*Ó,€LàÐ'à$×*Ñ*¨1Ó-×4Ñ4°RÓ8¸1Ñ<ÐBVÑWˆõð ÐØ˜1‘:ˆØ˜A‘Jˆ	÷	Mð
 #ˆÜ(ØØ
Ø
ØØØØôð r   Fr   ÚqueryÚkeyr6   c                 ó8   —  t        |«      «       } || ||fi |¤ŽS r   )r   )r^   r_   r6   r    r   Úflex_attention_compileds         r   Úcompile_friendly_flex_attentionrb   §   s5   € ð =Ô2°8Ó<Ó>ÐÙ"ØØØñð ñ	ð r   Úhidden_statesÚn_repc                 óª   — | j                   \  }}}}|dk(  r| S | dd…dd…ddd…dd…f   j                  |||||«      } | j                  |||z  ||«      S )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)
    r8   N)rR   ÚexpandÚreshape)rc   rd   ÚbatchÚnum_key_value_headsÚslenÚhead_dims         r   Ú	repeat_kvrl   ¹   so   € ð
 2?×1DÑ1DÑ.€EÐ  hØ�‚zØÐØ!¢!¢Q¨ªa²Ð"2Ñ3×:Ñ:¸5ÐBUÐW\Ð^bÐdlÓm€MØ× Ñ  Ð(;¸eÑ(CÀTÈ8ÓTÐTr   ÚmoduleÚattention_maskÚscalingÚsoftcapÚ	head_maskc                 óN  ‡‡‡— d }	d Št        |t        «      r|}	n|Š‰�‰d d …d d …d d …d |j                  d   …f   Šˆˆˆfd„}
d}|j                  d   }||dz
  z  dk(  sTt        ||j                  d   |j                  d   z  «      }t        ||j                  d   |j                  d   z  «      }d}|j	                  dd «      }t        ||||
|	|||d| j                  ¬«
      \  }}|j                  |j                  «      }|j                  dd	«      j                  «       }||fS )
Néþÿÿÿc                 óŽ   •— ‰�‰t        j                  | ‰z  «      z  } ‰�| ‰|   d   |   |   z   } ‰�| ‰|   |   d   d   z   } | S )Nr   )r#   Útanh)Úscorer<   r=   r>   r?   r@   rq   rp   s        €€€r   Ú	score_modz)flex_attention_forward.<locals>.score_modÚ   sm   ø€ ØÐØœeŸj™j¨°©Ó9Ñ9ˆEØÐ"Ø˜K¨	Ñ2°1Ñ5°eÑ<¸VÑDÑDˆEØÐ Ø˜I iÑ0°Ñ:¸1Ñ=¸aÑ@Ñ@ˆEØˆr   Tr8   r   FÚkernel_options)rw   Ú
block_maskÚ
enable_gqaÚscalerx   Ú
return_lser    r   )Ú
isinstancer
   rR   rl   Úgetrb   r    ÚtoÚdtypeÚ	transposeÚ
contiguous)rm   r^   r_   r6   rn   ro   rp   rq   r   ry   rw   rz   Únum_local_query_headsrx   Úattn_outputÚattention_weightsr@   s         ``        @r   Úflex_attention_forwardr†   Å   sA  ú€ ð €JØ€KÜ�.¤)Ô,Ø#‰
à$ˆàÐØ!¢!¢Qª¨?¨S¯Y©Y°r©]¨?Ð":Ñ;ˆöð €JØ!ŸK™K¨™NÐð #Ð&;¸aÑ&?Ñ@ÀQÒFÜ˜˜UŸ[™[¨™^¨s¯y©y¸©|Ñ;Ó<ˆÜ˜% §¡¨Q¡°5·;±;¸q±>Ñ!AÓBˆØˆ
à—Z‘ZÐ 0°$Ó7€NÜ%DØØØØØØØØ%ð Ø—‘ô&Ñ"€KÐ"ð *×,Ñ,¨U¯[©[Ó9ÐØ×'Ñ'¨¨1Ó-×8Ñ8Ó:€KàÐ)Ð)Ð)r   )NNNN)F)NNN)r-   Útypingr   r   r   r#   Ú	packagingr   Úutilsr   Úutils.import_utilsr	   Ú!torch.nn.attention.flex_attentionr
   r   r   rX   r   ÚTensorÚintÚOffsetr]   r.   r/   rb   rl   rS   ÚModuleÚfloatr†   r;   r   r   ú<module>r‘      sÑ  ðñ÷8 *Ñ )ã Ý å 0Ý /ñ  Ô!ßKõ÷
"-ñ "-ðJ 
ˆu�|‰|˜SÐ Ñ	!€ð
 +/ØØØ/3ñOØ—|‘|ðOà" 3™-ðOð
 �e˜F F˜NÑ+Ñ,ðOð óOðd ‡�×Ñ %ÐÓ(ð
 ñ	Ø�<‰<ðà	�‰ðð �<‰<ðð ‡\�\òó )ðð"	U˜UŸ\™\ð 	U°#ð 	U¸%¿,¹,ó 	Uð$  $Ø#Ø(,ñ:*Ø�H‰H�O‰Oð:*à�<‰<ð:*ð 
�‰ð:*ð �<‰<ð	:*ð
 ˜%Ÿ,™,¨Ð3Ñ4ð:*ð �e‰_ð:*ð �e‰_ð:*ð ˜Ÿ™Ñ%ð:*ð ˆ5�<‰<˜Ÿ™Ð%Ñ&ô:*r   