ó
    pyüi‘  ã            (       óœ  • S SK r S SKrS SKrS SKJr  S SKJr  S SKJr  S SK	r	S SK
Js  Jr  SSKJrJrJrJrJrJrJrJr  SSKJr  SSKJrJr  \R8                  " \5      rS	 rS
 r SSSS.r!S\S \S4\S4\S4\S44\S4\S\!S    S344S.S\S \S44SS.S\S \S44S S.S!.r"Sq#Sq$Sq%Sq&Sq'Sq(Sq)S"S#S$.r*S%S&0r+ S[S(\,S-  S)\S-  S*\-4S+ jjr.S, r/ S[S(\,S-  S)\S-  S*\-4S- jjr0S\S(\,S-  S*\-4S. jjr1S/ r2S]S0 jr3S1 r4S2\	Rj                  S3\6\	Rj                  \	Rj                  \74   4S4 jr8S5\	Rj                  S6\	Rj                  S7\	Rj                  S2\	Rj                  S8\74
S9 jr9S: r:S; r;S< r< S]S=\	Rj                  S>\	Rj                  S?\	Rj                  S@\	Rz                  S-  4SA jjr> " SB SC\S'SD9r?          S^S8\7SE\7SF\-SG\@SH\@S-  SI\7S-  SJ\-SK\@S-  SL\-S-  S%\	Rj                  S-  SM\7\	R‚                  -  S-  SN\7\	R‚                  -  S-  SO\B\,\-4   S-  4SP jjrC             S_SQ\	Rj                  SR\	Rj                  SS\	Rj                  S2\	Rj                  S-  S8\7SF\-SG\@ST\	Rj                  S-  SH\@S-  SI\7S-  SJ\-SK\@S-  SL\-S-  SU\	Rˆ                  S-  SV\	Rˆ                  S-  SW\7S-  SX\7S-  S@\	Rz                  S-  SY\,S-  4&SZ jjrEg)`é    N)ÚCallable)Úpartial)Ú	TypedDicté   )Úis_flash_attn_2_availableÚis_flash_attn_3_availableÚis_flash_attn_4_availableÚis_torch_cuda_availableÚis_torch_mlu_availableÚis_torch_npu_availableÚis_torch_xpu_availableÚlogging)Úsplit_attention_implementation)ÚPACKAGE_DISTRIBUTION_MAPPINGÚ
is_tracingc                  óx   • [        5       (       d  [        5       (       d  [        5       (       a  gSSKJn   U " 5       $ )NFr   ©Ú'is_npu_fa2_top_left_aligned_causal_mask)r   r   r	   Ú integrations.npu_flash_attentionr   r   s    Úh/home/mande/repo/quber/.venv/lib/python3.13/site-packages/transformers/modeling_flash_attention_utils.pyÚ!flash_attn_supports_top_left_maskr   *   s,   € Ü ×"Ñ"Ô&?×&AÑ&AÔE^×E`ÑE`ØåYá2Ó4Ð4ó    c                  óž   • [        5       =(       d=    [        5       =(       d,    [        5       =(       d    [        5       =(       d
    [	        5       $ ©N)r	   r   r   r   r   © r   r   Úis_flash_attn_availabler   4   s=   € ä!Ó#÷ 	$Ü$Ó&÷	$ä$Ó&÷	$ô "Ó#÷	$ô "Ó#ðr   zkernels-community/flash-attn2z"kernels-community/vllm-flash-attn3zkernels-community/flash-attn4)Úflash_attention_2Úflash_attention_3Úflash_attention_4é   c                  ó´   • [         R                  R                  S5      S L=(       a,    S[        S    Vs/ s H  o"R	                  SS5      PM     sn;   $ s  snf )NÚ
flash_attnz
flash-attnÚ_Ú-©Ú	importlibÚutilÚ	find_specr   Úreplace©ÚargsÚkwargsÚpkgs      r   Ú<lambda>r.   N   sR   € ¼)¿.¹.×:RÑ:RÐS_Ó:`ÐhlÐ:l÷ ;jØÔ>ZÐ[gÒ>hÓiÒ>h°sŸ[™[¨¨cÖ2Ñ>hÑiÑið;jùÚió   µAÚcudaÚmluÚnpuÚxpuz+Detect using FlashAttention2 on Ascend NPU.z*Detect using FlashAttention2 (via kernel `r   z
`) on XPU.)Úflash_attn_versionÚgeneral_availability_checkÚpkg_availability_checkÚsupported_devicesÚcustom_supported_devicesé   c                  ó´   • [         R                  R                  S5      S L=(       a,    S[        S    Vs/ s H  o"R	                  SS5      PM     sn;   $ s  snf )NÚflash_attn_interfacezflash-attn-3r#   r$   r%   r*   s      r   r.   r.   a   sR   € ¼)¿.¹.×:RÑ:RÐSiÓ:jÐrvÐ:v÷ ;vØÔ@\Ð]sÒ@tÓuÒ@t¸Ÿ{™{¨3°Ö4Ñ@tÑuÑuð;vùÚur/   é   )r4   r5   r6   r7   Úcuda_min_major_versioné   c                  ó´   • [         R                  R                  S5      S L=(       a,    S[        S    Vs/ s H  o"R	                  SS5      PM     sn;   $ s  snf )Nr"   zflash-attn-4r#   r$   r%   r*   s      r   r.   r.   i   sR   € ¼)¿.¹.×:RÑ:RÐS_Ó:`ÐhlÐ:l÷ ;lØÔ@\Ð]iÒ@jÓkÒ@j¸Ÿ{™{¨3°Ö4Ñ@jÑkÑkð;lùÚkr/   é	   )r    r9   r>   Ú	dropout_pÚwindow_size)ÚdropoutÚsliding_windowÚs_auxÚlearnable_sinkFÚimplementationÚattention_wrapperÚallow_all_kernelsc                 óê  • [        5       n[        5       n[        5       n[        [        pv[        U 5      u  p€U S:X  a  U(       d  U c)  U(       a"  U(       d  U(       d  SSKJn	Jn
J	n  SSK
JnJn  GO [        5       (       a  SSKJn	  SSKJn
  SS	KJn  OÞU S
:X  d  U c  U(       a  U(       d  SSKJn	Jn
J	n  O¼U S:X  d
  U c  U(       a  SSKJn	Jn
  SnO¡SSKJn  [,        R/                  X 5      nU(       a  SU  3OUnU" XáUS9n[1        USS5      n	[1        USS5      n
[1        USS5      nU
c  [3        SU  S35      eU	c  [4        R7                  SU  S35        Uc  [4        R7                  SU  S35        XšX¶U4$ )aÎ  
Lazy loads the respective flash attention implementations.

Return:
    flash_attn_func: The base flash attention function.
    flash_attn_varlen_func: The flash attention function supporting variable sequence lengths,
                            e.g. for padding-free training.
    pad_input: The function to pad inputs into one sequence and returning the respective kwargs.
    unpad_input: The function to unpad outputs based on the kwargs (from pad_input).
r   Nr   )Úflash_attn_funcÚflash_attn_varlen_funcÚflash_attn_with_kvcache)Ú	pad_inputÚunpad_inputr   )Únpu_flash_attn_func)Únpu_flash_attn_varlen_func)Únpu_flash_attn_with_kvcacher   r   )rK   rL   )Úload_and_register_attn_kernelzpaged|©rI   rK   rL   rM   zJCould not find the currently requested flash attention implementation at `z_`.Make sure that you request a valid kernel from the hub, e.g. `kernels-community/flash-attn2`.z.The loaded flash attention implementation at `z£` only supports varlen, i.e. it can only be used with continuous batching and does not support the full functionality for the base transformers generation methods.z‰` does not support block tables, so the full performances of continuous batching will not be achieved, only the varlen path will be used.)r   r   r	   Ú
_pad_inputÚ_unpad_inputr   r"   rK   rL   rM   Úflash_attn.bert_paddingrN   rO   r   r   rP   rQ   rR   r;   Úflash_attn.cuteÚintegrations.hub_kernelsrS   ÚFLASH_ATTN_KERNEL_FALLBACKÚgetÚgetattrÚ
ValueErrorÚloggerÚwarning)rG   rH   rI   Úis_fa2Úis_fa3Úis_fa4rN   rO   Úis_pagedrK   rL   rM   rS   Úkernel_repoÚkernel_implementationÚkernels                   r   Ú_lazy_importsrg   „   s’  € ô 'Ó(€FÜ&Ó(€FÜ&Ó(€Fä'¬ˆ{ä=¸nÓMÑ€HàÐ-Ó-¶&ØÑ¦6¶&Æç_Ñ_ßBÑBÜ	×	!Ñ	!õ 	]ÝjÞlàÐ0Ó0°^Ñ5KÖPVÖ_eßmÒmØÐ2Ó2°~Ñ7MÖRXßOà&*Ñ#õ Pô 5×8Ñ8¸ÓXˆKæAI f¨^Ð,<Ñ$=È{Ð!Ù2Ø%ÐL]ñˆFô & fÐ.?ÀÓFˆOÜ%,¨VÐ5MÈtÓ%TÐ"Ü&-¨fÐ6OÐQUÓ&VÐ#Ø%Ñ-Ü Ø`ÐaoÐ`pð qtð tóð ð Ñ&Ü—‘ØDÀ^ÐDTð U@ð @ôð
 'Ñ.Ü—‘ØDÀ^ÐDTð Uð ôð Ð4KÐXcÐcÐcr   c                 ó6  • [         R                  " U 5      R                  n[         R                  " [        5      R                  n0 nU H@  n[        R                  XD5      nXQ;   X5'   [        R                  XD5      =oe:w  d  M:  Xa;   X6'   MB     [        [        US9$ )a“  
Depending on the version and kernel some features are not supported. Due to limitations in
`torch.compile`, we opt to statically type which (optional) kwarg parameters are supported
within `_process_flash_attention_kwargs`.

NOTE: While all supported kwargs are marked as `True`, everything else is marked as `False`.
      This might be confusing for kwargs that we use in any case, e.g. `is_causal`.
)Úsupports_mapping)ÚinspectÚ	signatureÚ
parametersÚ_process_flash_attention_kwargsÚ_hf_api_to_flash_mappingr[   Ú_flash_api_alternative_namesr   )Úflash_functionÚflash_parametersÚprocess_parametersri   ÚparamÚfa_paramÚfa_alternative_names          r   Ú_lazy_define_process_functionrv   Ï   s“   € ô ×(Ò(¨Ó8×CÑCÐÜ ×*Ò*Ô+JÓK×VÑVÐàÐÛ#ˆÜ+×/Ñ/°Ó=ˆØ%-Ñ%AÐÑ"ä#?×#CÑ#CÀEÓ#QÐQÐÕ^Ø4GÑ4[ÐÓ1ñ $ô Ô2ÐEUÑVÐVr   c                 óÊ   • U c  [         c  [        S5      eU b+  [         U :w  a!  U q [        XUS9u  qqqqq[        [        5      q	[        [        [
        [        [        4[        4$ )zý
Lazily import flash attention and return the respective functions + flags.

NOTE: For fullgraph, this needs to be called before compile, while no fullgraph can
work without preloading. See `load_and_register_attn_kernel` in `integrations.hub_kernels`.
zGCould not find any flash attn implementation based on your environment.rT   )
Ú_loaded_implementationr]   rg   Ú	_flash_fnÚ_flash_varlen_fnÚ_flash_with_kvcache_fnÚ_pad_fnÚ	_unpad_fnrv   Ú_process_flash_kwargs_fn)rG   rH   rI   s      r   Úlazy_import_flash_attentionr   ç   sx   € ð ÑÔ"8Ñ"@ÜÐbÓcÐcð Ñ!Ô&<ÀÓ&NØ!/ÐäR_ØÐARñS
ÑOˆ	Ð#Ð%;¸WÀiô $AÔAQÓ#RÐ äÔ'Ô)?ÄÌ)ÐTÔVnÐnÐnr   c                 ó8   • SSK Jn  [        XUS9u  u  p4n    nXE4$ )za
Same as `lazy_import_flash_attention` but explicitly wrapping it with the paged implementation.
r   )Úpaged_attention_forward)rH   rI   )Úintegrations.flash_pagedr�   r   )rG   rI   r�   r#   rL   Úflash_attn_with_kvcache_fns         r   Ú!lazy_import_paged_flash_attentionr„      s3   € õ BäGbØÐUfñHÑDÑA€QÐ :¸A¸qÀ1ð "Ð=Ð=r   c                 óJ   • U R                   " S/U R                  SS Q76 nX!   $ )zø
A local implementation of the PyTorch indexing operation `tensor[indices]` on the first axis,
after flattening the first two dimensions of the tensor. This is functionally equivalent to
FA2's `index_first_axis` and replaces the need to import it.
éÿÿÿÿr    N)ÚreshapeÚshape)ÚtensorÚindicesÚreshaped_tensors      r   Ú_index_first_axisrŒ     s+   € ð —n’n RÐ;¨&¯,©,°q°rÐ*:Ò;€OØÑ#Ð#r   c                 ó   • Ub  X-   OUnUR                  S[        R                  S9nUR                  S[        R                  S9n[        R                  " UR	                  5       SS9R	                  5       nUR                  5       n[        R                  " [        R                  " US[        R                  S9S5      n[        X5      UUUU4$ )a  
unpad_input function for flash attention variants that do not have them within their pkg themselves, e.g. fa3.

Arguments:
    hidden_states: (batch, seqlen, ...)
    attention_mask: (batch, seqlen), bool / int, 1 means valid and 0 means not valid.
    unused_mask: (batch, seqlen), bool / int, 1 means the element is allocated but unused.

Return:
    hidden_states: (total_nnz, ...), where total_nnz = number of tokens selected in attention_mask + unused_mask.
    indices: (total_nnz), the indices of masked tokens from the flattened input sequence.
    cu_seqlens: (batch + 1), the cumulative sequence lengths, used to index into hidden_states.
    max_seqlen_in_batch: int
    seqused: (batch), returns the number of tokens selected in attention_mask + unused_mask.
r†   ©ÚdimÚdtypeF©Úas_tupler   ©r   r   )
ÚsumÚtorchÚint32ÚnonzeroÚflattenÚmaxÚFÚpadÚcumsumrŒ   )	Úhidden_statesÚattention_maskÚunused_maskÚ	all_masksÚseqlens_in_batchÚused_seqlens_in_batchrŠ   Úmax_seqlen_in_batchÚ
cu_seqlenss	            r   rV   rV     s¸   € ð  3>Ñ2I�Ò-È~€IØ —}‘}¨´5·;±;�}Ð?ÐØ*×.Ñ.°2¼U¿[¹[Ð.ÐIÐÜ�mŠm˜I×-Ñ-Ó/¸%Ñ@×HÑHÓJ€GØ*×.Ñ.Ó0ÐÜ—’”u—|’|Ð$4¸!Ä5Ç;Á;ÑOÐQWÓX€Jô 	˜-Ó1ØØØØðð r   c                 ó°   • U R                   SS n[        R                  " X#-  /UQ7U R                  U R                  S.6nXU'   UR
                  " X#/UQ76 $ )aú  
pad_input function for flash attention variants that do not have them within their pkg themselves, e.g. fa3.

Arguments:
    hidden_states: (total_nnz, ...), where total_nnz = number of tokens in selected in attention_mask.
    indices: (total_nnz), the indices that represent the non-masked tokens of the original padded input sequence.
    batch: int, batch size for the padded sequence.
    seqlen: int, maximum sequence length for the padded sequence.

Return:
    hidden_states: (batch, seqlen, ...)
r   N)Údevicer�   )rˆ   r•   Úzerosr¦   r�   Úview)r�   rŠ   ÚbatchÚseqlenr�   Úoutputs         r   rU   rU   8  sZ   € ð ×
Ñ
˜a˜bÐ
!€CÜ�[Š[˜%™.Ðh¨CÑh¸×8LÑ8LÐTa×TgÑTgÒh€FØ#ˆ7�OØ�;Š;�uÐ+ sÒ+Ð+r   rž   Úreturnc                 ó<  • U R                  S[        R                  S9n[        R                  " U R	                  5       SS9R	                  5       nUR                  5       n[        R                  " [        R                  " US[        R                  S9S5      nUUU4$ )aA  
Retrieves indexing data required to repad unpadded (ragged) tensors.

Arguments:
    attention_mask (`torch.Tensor`):
        Boolean or int tensor of shape (batch_size, sequence_length), 1 means valid and 0 means not valid.

Return:
    indices (`torch.Tensor`):
        The indices of non-masked tokens from the flattened input sequence.
    cu_seqlens (`torch.Tensor`):
        The cumulative sequence lengths, used to index into ragged (unpadded) tensors. `cu_seqlens` shape is (batch_size + 1,).
    max_seqlen_in_batch (`int`):
        Maximum sequence length in batch.
r†   rŽ   Fr‘   r   r“   )	r”   r•   r–   r—   r˜   r™   rš   r›   rœ   )rž   r¡   rŠ   r£   r¤   s        r   Ú_get_unpad_datar®   K  s…   € ð  &×)Ñ)¨b¼¿¹Ð)ÐDÐÜ�mŠm˜N×2Ñ2Ó4¸uÑE×MÑMÓO€GØ*×.Ñ.Ó0ÐÜ—’”u—|’|Ð$4¸!Ä5Ç;Á;ÑOÐQWÓX€JàØØðð r   Úquery_layerÚ	key_layerÚvalue_layerÚquery_lengthc                 ó   • [        U5      u  pgnUR                  S   UR                  S   =n	:”  a!  USS2SU	2SS2SS24   USS2SU	2SS2SS24   p!UR                  u  p«pÍ[        X5      n[        X&5      nXK:X  a  [        X5      n UnUnUnOhUS:X  aJ  Sn[        R                  " U
S-   [        R
                  U R                  S9nUSS nU R                  S5      n OUSS2U* S24   nU" X5      tn npïnU UUUXç4Xø44$ )a‡  
Unpads query, key, and values tensors, using a single dimension for all tokens even though they belong to different batches.
This function is used instead of `flash_attn.bert_padding.unpad_input` in order to avoid the recomputation of the same intermediary
tensors for query, key, value tensors.

Arguments:
    query_layer (`torch.Tensor`):
        Query state with padding. Shape: (batch_size, query_length, num_heads, head_dim).
    key_layer (`torch.Tensor`):
        Key state with padding. Shape: (batch_size, kv_seq_len, num_key_value_heads, head_dim).
    value_layer (`torch.Tensor`):
        Value state with padding. Shape: (batch_size, kv_seq_len, num_key_value_heads, head_dim).
    attention_mask (`torch.Tensor`):
        Boolean or int tensor of shape (batch_size, sequence_length), 1 means valid and 0 means not valid.
    query_length (`int`):
        Target length.
    unpad_input_func:
        The function to use for unpadding the input tensors.

Return:
    query_layer (`torch.Tensor`):
        Query state without padding. Shape: (total_target_length, num_heads, head_dim).
    key_layer (`torch.Tensor`):
        Key state with padding. Shape: (total_source_length, num_key_value_heads, head_dim).
    value_layer (`torch.Tensor`):
        Value state with padding. Shape: (total_source_length, num_key_value_heads, head_dim).
    indices_q (`torch.Tensor`):
        The indices of non-masked tokens from the flattened input target sequence.
    (cu_seqlens_q, cu_seqlens_k) (`tuple[int]`):
        The cumulative sequence lengths for the target (query) and source (key, value), used to index into ragged (unpadded) tensors. `cu_seqlens` shape is (batch_size + 1,).
    (max_seqlen_in_batch_q, max_seqlen_in_batch_k) (`tuple[int]`):
        Maximum sequence length in batch (`max_seqlen_in_batch_q` for the target sequence i.e. query, `max_seqlen_in_batch_k` for the source sequence i.e. key/value).
r   r†   N©r�   r¦   )r®   rˆ   rŒ   r•   Úaranger–   r¦   Úsqueeze)r¯   r°   r±   rž   r²   Úunpad_input_funcÚ	indices_kÚcu_seqlens_kÚmax_seqlen_in_batch_kÚseq_lenÚ
batch_sizeÚ
kv_seq_lenÚnum_key_value_headsÚhead_dimÚcu_seqlens_qÚmax_seqlen_in_batch_qÚ	indices_qr#   s                     r   Ú_upad_inputrÃ   f  sE  € ôR 6EÀ^Ó5TÑ2€IÐ2ð ‡��qÑ¨×(<Ñ(<¸RÑ(@Ð@˜WÓAØ!*ª1¨h¨w¨hºº1Ð+<Ñ!=¸{Ê1ÈhÈwÈhÒXYÒ[\ÐK\Ñ?]�;à<E¿O¹OÑ9€JÐ/ä! )Ó7€IÜ# KÓ;€KØÓ!Ü'¨Ó?ˆØ#ˆØ 5ÐØ‰	Ø	˜Ó	Ø !ÐÜ—|’|Ø˜‰N¤%§+¡+°k×6HÑ6Hñ
ˆð !  "Ð%ˆ	Ø!×)Ñ)¨!Ó,‰ð (ª¨L¨=©>Ð(9Ñ:ˆÙJZÐ[fÓJwÐGˆ�Y Àað 	ØØØØ	Ð$Ø	Ð6ðð r   c                 óˆ  • [         R                  U R                  S.nU R                  S5      n U S:H  R	                  5       R                  S5      n[         R                  " UR                  " S0 UD6[         R                  " U R                  5       40 UD645      nUnUR                  5       R                  5       nUnX44XV44$ )aë  
This function returns all the necessary kwargs to call `flash_attn_varlen_func` extracted from position_ids.

Arguments:
    position_ids (`torch.Tensor`):
        Boolean or int tensor of shape (batch_size, sequence_length), 1 means valid and 0 means not valid.

Return:
    (cu_seqlens_q, cu_seqlens_k) (`tuple[int]`):
        The cumulative sequence lengths for the target (query) and source (key, value), used to index into
        ragged (unpadded) tensors. `cu_seqlens` shape is (batch_size + 1,).
    (max_seqlen_in_batch_q, max_seqlen_in_batch_k) (`tuple[int]`):
        Maximum sequence length in batch (`max_seqlen_in_batch_q` for the target sequence i.e. query,
        `max_seqlen_in_batch_k` for the source sequence i.e. key/value).
r´   r†   r   r   )r•   r–   r¦   r‡   r—   r¨   ÚcatÚtor‰   ÚsizeÚdiffr™   )Úposition_idsÚtensor_kwargsrÂ   Úcu_seq_lens_qÚcu_seq_lens_kÚmax_length_qÚmax_length_ks          r   Ú#prepare_fa_kwargs_from_position_idsrÏ   µ  s¸   € ô  $Ÿk™k°\×5HÑ5HÑI€Mà×'Ñ'¨Ó+€LØ Ñ"×+Ñ+Ó-×2Ñ2°2Ó6€Iä—I’Ià�LŠLÑ)˜=Ñ)Ü�LŠL˜×*Ñ*Ó,Ñ>°Ñ>ð	
ó€Mð "€Mð
 !×%Ñ%Ó'×+Ñ+Ó-€LØ€LàÐ)¨LÐ+GÐGÐGr   c                 ó°  • U R                  5       R                  SU R                  S5      U R                  S5      5      n UR                  5       R                  SUR                  S5      UR                  S5      5      nUR                  5       R                  SUR                  S5      UR                  S5      5      n[        U5      u  u  pEu  pgXX$U4Xg44$ )ai  
This function returns necessary arguments to call `flash_attn_varlen_func`.
All three query, key, value states will be flattened.
Cumulative lengths of each examples in the batch will be extracted from position_ids.
NOTE: ideally cumulative lengths should be prepared at the data collator stage

Arguments:
    query (`torch.Tensor`):
        Query state with padding. Shape: (batch_size, query_length, num_heads, head_dim).
    key (`torch.Tensor`):
        Key state with padding. Shape: (batch_size, kv_seq_len, num_key_value_heads, head_dim).
    value (`torch.Tensor`):
        Value state with padding. Shape: (batch_size, kv_seq_len, num_key_value_heads, head_dim).
    position_ids (`torch.Tensor`):
        Boolean or int tensor of shape (batch_size, sequence_length), 1 means valid and 0 means not valid.

Return:
    query (`torch.Tensor`):
        Query state without padding. Shape: (total_target_length, num_heads, head_dim).
    key (`torch.Tensor`):
        Key state with padding. Shape: (total_source_length, num_key_value_heads, head_dim).
    value (`torch.Tensor`):
        Value state with padding. Shape: (total_source_length, num_key_value_heads, head_dim).
    (cu_seqlens_q, cu_seqlens_k) (`tuple[int]`):
        The cumulative sequence lengths for the target (query) and source (key, value), used to index into ragged (unpadded) tensors. `cu_seqlens` shape is (batch_size + 1,).
    (max_seqlen_in_batch_q, max_seqlen_in_batch_k) (`tuple[int]`):
        Maximum sequence length in batch (`max_seqlen_in_batch_q` for the target sequence i.e. query, `max_seqlen_in_batch_k` for the source sequence i.e. key/value).
r†   éþÿÿÿ)Ú
contiguousr¨   rÇ   rÏ   )ÚqueryÚkeyÚvaluerÉ   rË   rÌ   rÍ   rÎ   s           r   Ú_prepare_from_posidsrÖ   Û  s´   € ð: ×ÑÓ×#Ñ# B¨¯
©
°2«¸¿
¹
À2»ÓG€EØ
�.‰.Ó
×
Ñ
  C§H¡H¨R£L°#·(±(¸2³,Ó
?€CØ×ÑÓ×#Ñ# B¨¯
©
°2«¸¿
¹
À2»ÓG€EäCfÐgsÓCtÑ@Ñ"€]Ñ$@ \à˜¨}Ð=ÀÐ?[Ð\Ð\r   c                 óø   • U c  g[         R                  " U R                  S   U R                  S9U R	                  5       -   nUS:H  =(       a.    X -
  R                  5       R                  5       R                  5       $ )a  
Check the position ids whether packed sequences are indicated or not
    1. Position ids exist
    2. Flattened sequences only are supported
    3. Compile-friendly `not (torch.diff(position_ids, dim=-1) >= 0).all()`, i.e. we have multiple increasing sequences
Fr   )r¦   )r•   rµ   rˆ   r¦   ÚminÚabsr”   Úbool)rÉ   r¼   Úincreasing_position_sequencess      r   Ú_is_packed_sequencerÜ     sq   € ð ÑØô 	�Š�\×'Ñ'¨Ñ*°<×3FÑ3FÑGÈ,×JZÑJZÓJ\Ñ\ð "ð ˜‰?×`Ð =Ñ L×QÑQÓS×WÑWÓY×^Ñ^Ó`Ð`r   ÚqÚkÚvÚtarget_dtypec                 óê   • U(       ai  U R                   [        R                  :X  aK  [        R	                  SU S35        U R                  U5      UR                  U5      UR                  U5      p!n XU4$ )aM  
PEFT usually casts the layer norms in float32 for training stability reasons
therefore the input hidden states gets silently casted in float32. Hence, we need
cast them back in float16 / bfloat16 just to be sure everything works as expected.
This might slowdown training & inference so it is recommended to not cast the LayerNorms!
zCasting fp32 inputs back to z for flash-attn compatibility.)r�   r•   Úfloat32r^   Úwarning_oncerÆ   )rÝ   rÞ   rß   rà   s       r   Úfa_peft_integration_checkrä     s^   € ö ˜Ÿ™¤5§=¡=Ó0Ü×ÑÐ:¸<¸.ÐHfÐgÔhØ—$‘$�|Ó$ a§d¡d¨<Ó&8¸!¿$¹$¸|Ó:LˆaˆØ�ˆ7€Nr   c                   ó‚   • \ rS rSr% Sr\R                  S-  \S'   \R                  S-  \S'   \S-  \S'   \S-  \S'   Sr	g)	ÚFlashAttentionKwargsi#  aÄ  
Keyword arguments for Flash Attention with Compile.

Attributes:
    cu_seq_lens_q (`torch.LongTensor`, *optional*)
        Gets cumulative sequence length for query state.
    cu_seq_lens_k (`torch.LongTensor`, *optional*)
        Gets cumulative sequence length for key state.
    max_length_q (`int`, *optional*):
        Maximum sequence length for query state.
    max_length_k (`int`, *optional*):
        Maximum sequence length for key state.
NrË   rÌ   rÍ   rÎ   r   )
Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r•   Ú
LongTensorÚ__annotations__ÚintÚ__static_attributes__r   r   r   ræ   ræ   #  s?   ‡ ñð ×#Ñ# dÑ*Ó*Ø×#Ñ# dÑ*Ó*Ø˜‘*ÓØ˜‘*Ör   ræ   )ÚtotalÚ
key_lengthÚ	is_causalrC   Úsoftmax_scalerD   Úuse_top_left_maskÚsoftcapÚdeterministicÚmax_seqlen_qÚmax_seqlen_kri   c                 ó®  • U=(       a    U=(       a    U S:H  (       + US.nUS   (       a  X>S'   US   (       a  Ub  X:”  a  US-
  US-
  4US'   US   (       a"  Ub  UO[         R                  " SS5      S:H  US'   US	   (       a  Ub  X~S	'   US
   =n(       d
  US   (       a  U	b  U(       a  XžS
'   OXžS'   X«L nUS   (       a<  U
b9  [        U
[        5      (       d   [	        U
5      (       a  U
R                  5       n
X®S'   US   (       aO  UbL  U(       a  US   b  US   nO5[        U[        5      (       d   [	        U5      (       a  UR                  5       nX¾S'   U$ )a  
Returns a set of kwargs that are passed down to the according flash attention function based on
requested features and whether it is supported - depends on the version and kernel implementation
which is dynamically configured at `lazy_import_flash_attention`. The (un)supported features can be
inspected in `supports_mapping`, see `_lazy_define_process_function` for more details.

Args:
    query_length (`int`):
        Length of the query states
    key_length (`int`):
        Length of the key states
    is_causal (`bool`):
        Whether we perform causal (decoder) attention or full attention.
    dropout (`float`):
        Attention dropout.
    softmax_scale (`float`, *optional*):
        The scaling of QK^T before applying softmax. Default to `1 / sqrt(head_dim)`.
    sliding_window (`int`, *optional*):
        The size of the sliding window, i.e. we look at a max of `sliding_window` tokens back.
    use_top_left_mask (`bool`):
        Deprecated behavior of older versions of flash attention requiring different masking.
    softcap (`float`, *optional*):
        Softcap for the attention logits, used e.g. in gemma2.
    deterministic (`bool`, *optional*):
        Determines if the deterministic option introduced in flash_attn>=2.4.1 is enabled.
    s_aux (`torch.Tensor`, *optional*):
        Attention sink auxiliary that adds a `bias` to the attention calculation via an additional head.
    max_seqlen_q (`Union[int, torch.IntTensor]`, *optional*):
        The maximum sequence length in the query tensor during a varlen forward.
    max_seqlen_k (`Union[int, torch.IntTensor]`, *optional*):
        The maximum sequence length in the key/value tensor during a varlen forward.
Return:
    flash_kwargs (`dict`):
        A dict of kwargs that are requested and supported.
r   )Úcausalró   rA   rB   rö   ÚFLASH_ATTENTION_DETERMINISTICÚ0Ú1rõ   rE   rF   r÷   rø   )ÚosÚgetenvÚ
isinstancerî   r   Úitem)r²   rñ   rò   rC   ró   rD   rô   rõ   rö   rE   r÷   rø   ri   r,   Úflash_kwargsÚlegacy_sink_paramÚsame_max_seqlens                    r   rm   rm   8  sy  € ðh ×MÐ%6×%L¸<È1Ñ;LÔ MØ&ñ€Lð
 ˜×$Ø$+�[Ñ!à˜×&¨>Ñ+EÈ*ÓJeð
 (6¸Ñ'9¸>ÈAÑ;MÐ&Nˆ�]Ñ#à˜×(à*Ñ6‰M¼B¿IºIÐFeÐgjÓ<kÐorÑ<rð 	�_Ñ%ð ˜	×" wÑ':Ø")�YÑà.¨wÑ7Ð	7Ð	Õ	7Ð<LÐM]×<^ÐdiÑduÞØ$)˜Ò!à-2Ð)Ñ*ð #Ð2€OØ˜×'¨LÑ,DÜ˜,¬×,Ñ,´¸L×1IÑ1IØ'×,Ñ,Ó.ˆLØ'3�^Ñ$à˜×'¨LÑ,DÞ˜|¨NÑ;ÑGØ'¨Ñ7‰LÜ˜L¬#×.Ñ.´:¸l×3KÑ3KØ'×,Ñ,Ó.ˆLØ'3�^Ñ$àÐr   Úquery_statesÚ
key_statesÚvalue_statesrÉ   rË   rÌ   rÍ   rÎ   Úattn_implementationc                 ó  • [        U5      u  u  nnnnnn[        XUU5      u  pn[        U4UUR                  S5      UUUU	U
UUS.	UD6n[	        XpR                  S5      S9n[        S XÞUU4 5       5      nUbŒ  [        XX#UU5      u  nnnn u  pÞu  nnS[        UR                  5      ;   a  UR                  5       nU" UUU4UUS.U" UUS9D6n![        U![        5      (       a  U!S   n!U" U!U U R                  S5      U5      n"U"$ U(       d  U(       GaJ  Ub  Uc  [        XX'5      u  nnnu  pÞu  nnO“U R                  S	U R                  S
5      U R                  S	5      5      nUR                  S	UR                  S
5      UR                  S	5      5      nUR                  S	UR                  S
5      UR                  S	5      5      nS[        UR                  5      ;   a  UR                  5       nU" UUU4UUS.U" UUS9D6n"[        U"[        5      (       a  U"S   n"U"R                  U R                  S5      S	U"R                  S
5      U"R                  S	5      5      n"U"$ U" XU40 U" 5       D6n"[        U"[        5      (       a  U"S   n"U"$ )aß  
Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token
first unpad the input, then computes the attention scores and pad the final attention scores.

(Optional) kwargs are described further in `_process_flash_attention_kwargs` and `FlashAttentionKwargs`.

Args:
    query_states (`torch.Tensor`):
        Input query states to be passed to Flash Attention API
    key_states (`torch.Tensor`):
        Input key states to be passed to Flash Attention API
    value_states (`torch.Tensor`):
        Input value states to be passed to Flash Attention API
    attention_mask (`torch.Tensor`, *optional*):
        The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the
        position of padding tokens and 1 for the position of non-padding tokens.
    attn_implementation (`str`, *optional*):
        The attention implementation to use. If None, will default to the one based on the environment.
r   )	r²   rñ   rò   rC   ró   rD   rô   rõ   rö   r   )r¼   c              3   ó(   #   • U  H  oS Lv •  M
     g 7fr   r   )Ú.0Úkwargs     r   Ú	<genexpr>Ú+_flash_attention_forward.<locals>.<genexpr>è  s   é € ð #Ú'a˜e�TÕÒ'aùs   ‚Úmps)rÀ   r¹   )r÷   rø   r†   rÑ   )r   rä   r   rÇ   rÜ   ÚallrÃ   Ústrr¦   Úcloner   ÚtuplerÖ   r‡   r¨   )#r  r  r  rž   r²   rò   rC   rÉ   ró   rD   rô   rõ   rö   rË   rÌ   rÍ   rÎ   rà   r  r,   Úflash_fnÚflash_varlen_fnr#   Úpad_fnÚunpad_fnÚprocess_flash_kwargs_fnr  Úis_fa_with_position_idsÚis_fa_with_varlen_kwargsrÝ   rÞ   rß   rÂ   Ú	out_unpadÚouts#                                      r   Ú_flash_attention_forwardr  Ÿ  sü  € ôR QlØóQÑMÑ4€Xˆ  6¨8Ð6Mô
 .GØ ,°ó.Ñ*€L˜lô
 Øðà!Ø—?‘? 1Ó%ØØØ#Ø%Ø+ØØ#ñð ñ€Lô* 2°,×K\ÑK\Ð]^ÓK_Ñ`ÐÜ"ñ #Ø(5ÀlÐT`Ñ'aó#ó  Ðð
 Ñ!Ü[fØ lÀLÐRZó\
ÑXˆˆ1ˆa�Ñ:˜]Ñ<X¸\È<ð ”C˜Ÿ™“MÓ!Ø)×/Ñ/Ó1ˆMá#ØØØð
ð 'Ø&ñ
ñ ¨À<ÑPñ
ˆ	ô �i¤×'Ñ'Ø! !™ˆIá�Y 	¨<×+<Ñ+<¸QÓ+?ÀÓNˆðJ €JöE 
"×%<ØÑ  MÑ$9ÜThØ¨,óUÑQˆAˆq�!Ñ3�mÑ5Q°lÁLð ×$Ñ$ R¨×):Ñ):¸2Ó)>À×@QÑ@QÐRTÓ@UÓVˆAØ×"Ñ" 2 z§¡°rÓ':¸J¿O¹OÈBÓ<OÓPˆAØ×$Ñ$ R¨×):Ñ):¸2Ó)>À×@QÑ@QÐRTÓ@UÓVˆAð ”C˜Ÿ™“MÓ!Ø)×/Ñ/Ó1ˆMáØØØð
ð 'Ø&ñ
ñ ¨À<ÑPñ
ˆô �cœ5×!Ñ!Ø�a‘&ˆCà�h‰h�|×(Ñ(¨Ó+¨R°·±¸"³¸s¿x¹xÈ»|ÓLˆð €Jñ	 �|°ÑPÁÃÑPˆÜ�cœ5×!Ñ!Ø�a‘&ˆCà€Jr   )NF)Fr   )
ç        NNFNNNNNN)r  NNNFNNNNNNNN)Fr&   rj   rþ   Úcollections.abcr   Ú	functoolsr   Útypingr   r•   Útorch.nn.functionalÚnnÚ
functionalrš   Úutilsr   r   r	   r
   r   r   r   r   Úutils.genericr   Úutils.import_utilsr   r   Ú
get_loggerrç   r^   r   r   rZ   Ú$FLASH_ATTENTION_COMPATIBILITY_MATRIXrx   ry   rz   r{   r|   r}   r~   rn   ro   r  rÚ   rg   rv   r   r„   rŒ   rV   rU   ÚTensorr  rî   r®   rÃ   rÏ   rÖ   rÜ   r�   rä   ræ   ÚfloatÚ	IntTensorÚdictrm   rì   r  r   r   r   Ú<module>r.     s  ðó Û Û 	Ý $Ý Ý ã ß Ð ÷	÷ 	ó 	õ :ß Hð 
×	Ò	˜HÓ	%€ò5òð 9Ø=Ø8ñÐ ð  Ø&?ñ#jð % fÐ-Ø# UÐ+Ø# UÐ+Ø# UÐ+ð	
ð $Ð%RÐSà&Ø<Ð=WÐXkÑ=lÐ<mÐmwÐxðð%
ñð(  Ø&?ñ#và6¸Ð?ÐAØ"#ñð  Ø&?ñ#là6¸Ð?ÐAØ"#ññ9$(Ð $ðP Ð Ø€	ØÐ ØÐ Ø
€Ø€	ð  Ð ð Ø#ñÐ ð
 !(Ð)9Ð:Ð ð fkñHdØ˜$‘JðHdØ3;¸d±?ðHdØ^bõHdòVWð2 fkñoØ˜$‘JðoØ3;¸d±?ðoØ^bõoñ2	>°c¸D±jð 	>ÐUYõ 	>ò	$ôò@,ð& E§L¡Lð °U¸5¿<¹<ÈÏÉÐWZÐ;ZÑ5[ô ð6LØ—‘ðLà�|‰|ðLð —‘ðLð —L‘Lð	Lð
 ôLò^#HòL#]òLað( (,ñ	Ø‡|�|ðà‡|�|ðð ‡|�|ðð —+‘+ Ñ$õ	ô$˜9¨Eò ð2 Ø"&Ø!%Ø#Ø Ø!%Ø!%Ø15Ø15Ø/3ñdØðdàðdð ðdð ð	dð
 ˜4‘<ðdð ˜$‘Jðdð ðdð �T‰\ðdð ˜$‘;ðdð �<‰<˜$Ñðdð ˜Ÿ™Ñ'¨$Ñ.ðdð ˜Ÿ™Ñ'¨$Ñ.ðdð ˜3 ˜9‘o¨Ñ,õdð\ Ø(,Ø"&Ø!%Ø#Ø Ø!%Ø-1Ø-1Ø#Ø#Ø'+Ø&*ñ'HØ—,‘,ðHà—‘ðHð —,‘,ðHð —L‘L 4Ñ'ð	Hð
 ðHð ðHð ðHð —,‘, Ñ%ðHð ˜4‘<ðHð ˜$‘JðHð ðHð �T‰\ðHð ˜$‘;ðHð ×#Ñ# dÑ*ðHð ×#Ñ# dÑ*ðHð  ˜‘*ð!Hð" ˜‘*ð#Hð$ —+‘+ Ñ$ð%Hð& ˜t™ö'Hr   