Ë
    S^(hX#  ã                   óŒ  — d dl Z d dlZd dlmc mZ ddlmZ  e«       rd dlZd dl	m
Z
mZ dZdZ e e j                  de¬«      «      Zeeefvr ed«      ‚d	„ Z G d
„ dej&                  j(                  «      Zej,                  Z G d„ dej&                  j(                  «      Zej,                  Zd„ Zdd„Z	 	 	 dd„Z	 	 	 dd„Zy)é    Né   )Úis_torch_npu_available)Ú	rearrangeÚrepeaté   ÚNPU_FA2_SPARSE_MODE)Údefaultz…Environment variable `NPU_FA2_SPARSE_MODE` can only be set as 2 (top-left aligned causal mask) or 3 (down-right aligned causal mask).c                  ó4   — t        «       rt        t        k(  S dS )NF)r   ÚSPARSE_MODEÚ!TOP_LEFT_ALIGNED_CAUSAL_MASK_MODE© ó    úk/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/integrations/npu_flash_attention.pyÚ'is_npu_fa2_top_left_aligned_causal_maskr   '   s   € Ü?UÔ?WŒ;Ô;Ñ;ÐbÐ]bÐbr   c                   ó,   — e Zd Zed„ «       Zed„ «       Zy)ÚIndexFirstAxisc           
      ó*  — | j                  |«       |j                  dk\  sJ ‚|j                  d   |j                  dd  c| _        }|j	                  «       } t        j                  t        |d«      dt        |d|¬«      «      j                  dg|¢­Ž S )Nr   r   é   úb ... -> b (...)úz -> z d©Údéÿÿÿÿ)
Úsave_for_backwardÚndimÚshapeÚfirst_axis_dimÚnumelÚtorchÚgatherr   r   Úreshape)ÚctxÚinputÚindicesÚother_shapeÚ
second_dims        r   ÚforwardzIndexFirstAxis.forward-   s‘   € à×Ñ˜gÔ&Ø�z‰z˜QŠÐˆØ*/¯+©+°a©.¸%¿+¹+ÀaÀb¸/Ð'ˆÔ˜KØ ×&Ñ&Ó(ˆ
ðŒu�|‰|Ü�eÐ/Ó0°!´V¸GÀZÐS]Ô5^ó
ç
‰'�"ð$à"ò$ð 	$r   c           	      ó–  — | j                   \  }|j                  dk\  sJ ‚|j                  dd  }t        |d«      }t	        j
                  | j                  |j                  d   g|j                  |j                  ¬«      }|j                  dt        |d|j                  d   ¬«      |«        |j                  | j                  g|¢­Ž d fS )Nr   r   r   ©ÚdeviceÚdtyper   r   r   )Úsaved_tensorsr   r   r   r   Úzerosr   r*   r+   Úscatter_r   r!   )r"   Úgrad_outputr$   r%   Ú
grad_inputs        r   ÚbackwardzIndexFirstAxis.backward9   sÇ   € à×&Ñ&‰
ˆØ×Ñ 1Ò$Ð$Ð$Ø!×'Ñ'¨¨Ð+ˆÜ Ð-?Ó@ˆÜ—[‘[Ø×Ñ ×!2Ñ!2°1Ñ!5Ð6Ø×%Ñ%Ø×#Ñ#ô
ˆ
ð 	×Ñ˜Aœv g¨z¸[×=NÑ=NÈqÑ=QÔRÐT_Ô`Ø!ˆz×!Ñ! #×"4Ñ"4ÐC°{ÒCÀTÐIÐIr   N©Ú__name__Ú
__module__Ú__qualname__Ústaticmethodr'   r1   r   r   r   r   r   ,   s*   „ Øñ	$ó ð	$ð ñJó ñJr   r   c                   ó,   — e Zd Zed„ «       Zed„ «       Zy)ÚIndexPutFirstAxisc                 óì   — | j                  |«       |j                  dk(  sJ ‚|j                  dk\  sJ ‚t        j                  |g|j                  dd  ¢­|j
                  |j                  dœŽ}|||<   |S )Nr   r   r)   )r   r   r   r-   r   r*   r+   )r"   Úvaluesr$   r   Úoutputs        r   r'   zIndexPutFirstAxis.forwardO   sr   € à×Ñ˜gÔ&Ø�|‰|˜qÒ Ð Ð Ø�{‰{˜aÒÐÐÜ—‘˜^Ði¨f¯l©l¸1¸2Ð.>ÑiÀvÇ}Á}Ð\b×\hÑ\hÒiˆà ˆˆw‰àˆr   c                 ó2   — | j                   \  }||   }|d d fS ©N)r,   )r"   r/   r$   Úgrad_valuess       r   r1   zIndexPutFirstAxis.backwardZ   s&   € à×&Ñ&‰
ˆà! 'Ñ*ˆà˜D $Ð&Ð&r   Nr2   r   r   r   r8   r8   N   s(   „ Øñó ðð ñ'ó ñ'r   r8   c                 ó>   — t        | |||z  «      }t        |d|¬«      S )a«  
    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, ...)
    z(b s) ... -> b s ...)Úb)Úindex_put_first_axisr   )Úhidden_statesr$   ÚbatchÚseqlenr;   s        r   Ú	pad_inputrE   g   s&   € ô " -°¸%À&¹.ÓI€FÜ�VÐ3°uÔ=Ð=r   c                 óä  — |�||z   n|}|j                  dt        j                  ¬«      }|j                  dt        j                  ¬«      }t        j                  |j	                  «       d¬«      j	                  «       }|j                  «       j                  «       }t        j                  t        j                  |dt        j                  ¬«      d«      }t        t        | d«      |«      ||||fS )a¿  
    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   )Údimr+   F)Úas_tupler   )r   r   zb s ... -> (b s) ...)Úsumr   Úint32ÚnonzeroÚflattenÚmaxÚitemÚFÚpadÚcumsumÚindex_first_axisr   )	rB   Úattention_maskÚunused_maskÚ	all_masksÚseqlens_in_batchÚused_seqlens_in_batchr$   Úmax_seqlen_in_batchÚ
cu_seqlenss	            r   Úunpad_inputrZ   y   sÍ   € ð 3>Ð2I� +Ò-È~€IØ —}‘}¨´5·;±;�}Ó?ÐØ*×.Ñ.°2¼U¿[¹[Ð.ÓIÐÜ�m‰m˜I×-Ñ-Ó/¸%Ô@×HÑHÓJ€GØ*×.Ñ.Ó0×5Ñ5Ó7ÐÜ—‘”u—|‘|Ð$4¸!Ä5Ç;Á;ÔOÐQWÓX€Jô 	œ =Ð2HÓIÈ7ÓSØØØØðð r   c                 ó‚  — d|z
  }|s0| j                   d   }t        j                  | |||d||¬«      d   }	|	S t        j                  t        j
                  ddg«      d¬«      j                  «       j                  | j                  «      }
| j                   d   }t        j                  | |||d|||
t        ¬	«	      d   }	|	S )
Nç      ð?r   ÚBSND)Ú	keep_probÚscaler   é   r   ©Údiagonal)r^   r_   Ú
atten_maskÚsparse_mode)
r   Ú	torch_npuÚnpu_fusion_attentionr   ÚtriuÚonesÚboolÚtor*   r   )ÚqÚkÚvÚ	dropout_pÚsoftmax_scaleÚcausalÚkwargsr^   Úhead_numr;   Úattn_mask_npus              r   Únpu_flash_attn_funcrt   š   sË   € ð �i‘€IáØ—7‘7˜1‘:ˆÜ×/Ñ/°°1°a¸À6ÐU^ÐfsÔtÐuvÑwˆð  €Mô Ÿ
™
¤5§:¡:¨t°T¨lÓ#;ÀaÔH×MÑMÓO×RÑRÐST×S[ÑS[Ó\ˆØ—7‘7˜1‘:ˆÜ×/Ñ/ØØØØØØØØ$Ü#ô

ð ñ
ˆð €Mr   c                 óB  — d|z
  }	|s | j                   d   }
t        j                  | |||
d d ||	dt        |dd  j	                  «       j                  «       j                  «       «      t        |dd  j	                  «       j                  «       j                  «       «      ¬«      d   }|S t        j                  t        j                  ddg«      d¬«      j                  «       j                  | j                  «      }| j                   d   }
t        j                  | |||
d d |||	dt        |dd  j	                  «       j                  «       j                  «       «      t        |dd  j	                  «       j                  «       j                  «       «      t        ¬«      d   }|S )	Nr\   r   ÚTND)Úpserc   r_   r^   Úinput_layoutÚactual_seq_qlenÚactual_seq_kvlenr   r`   ra   )	rw   Úpadding_maskrc   r_   r^   rx   ry   rz   rd   )r   re   rf   ÚtupleÚcpuÚnumpyÚtolistr   rg   rh   ri   rj   r*   r   )rk   rl   rm   Úcu_seqlens_qÚcu_seqlens_krn   ro   rp   rq   r^   rr   r;   rs   s                r   Únpu_flash_attn_varlen_funcr‚   º   s‹  € ð �i‘€IáØ—7‘7˜1‘:ˆÜ×/Ñ/ØØØØØØØØØÜ! ,¨q¨rÐ"2×"6Ñ"6Ó"8×">Ñ">Ó"@×"GÑ"GÓ"IÓJÜ" <°°Ð#3×#7Ñ#7Ó#9×#?Ñ#?Ó#A×#HÑ#HÓ#JÓKô
ð ñˆð@ €Mô% Ÿ
™
¤5§:¡:¨t°T¨lÓ#;ÀaÔH×MÑMÓO×RÑRÐST×S[ÑS[Ó\ˆØ—7‘7˜1‘:ˆÜ×/Ñ/ØØØØØØØ$ØØØÜ! ,¨q¨rÐ"2×"6Ñ"6Ó"8×">Ñ">Ó"@×"GÑ"GÓ"IÓJÜ" <°°Ð#3×#7Ñ#7Ó#9×#?Ñ#?Ó#A×#HÑ#HÓ#JÓKÜ#ô
ð ñˆð  €Mr   r=   )g        NF)Úosr   Útorch.nn.functionalÚnnÚ
functionalrO   Úutils.import_utilsr   re   Úeinopsr   r   r   Ú#DOWN_RIGHT_ALIGNED_CAUSAL_MASK_MODEÚintÚgetenvr   Ú
ValueErrorr   ÚautogradÚFunctionr   ÚapplyrR   r8   rA   rE   rZ   rt   r‚   r   r   r   ú<module>r�      sè   ðó 
ã ß Ð å 7ñ ÔÛß(ð
 %&Ð !Ø&'Ð #á�)�"—)‘)Ð1Ð;^Ô_Ó`€ØÐ8Ð:]Ð^Ñ^Ù
ð	1óð òcô
J�U—^‘^×,Ñ,ô Jð< "×'Ñ'Ð ô'˜Ÿ™×/Ñ/ô 'ð* )×.Ñ.Ð ò>ó$ðJ ØØóðL ØØô/r   