ó
    !Eñi(  ã                   ó0  • % S r SSKrSSKr/ r\\   \S'   S\R                  S\4S jr	S\R                  S\
S\R                  4S	 jrS
\
S\S-  S\4S jr     SS\R                  S\R                  S\R                  S\R                  S-  SS4
S jjrg)zCDefines utilities for interacting with scaled_dot_product_attentioné    NÚ__all__ÚtensorsÚreturnc                  ó&   • [        S U  5       5      $ )z0Returns True if any of the tensors requires gradc              3   ó8   #   • U  H  oR                   v •  M     g 7f)N)Úrequires_grad)Ú.0Úts     ÚV/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/nn/attention/_utils.pyÚ	<genexpr>Ú'_input_requires_grad.<locals>.<genexpr>   s   é € Ð0ª 1�Žªùs   ‚)Úany)r   s    r   Ú_input_requires_gradr      s   € äÑ0©Ó0Ó0Ð0ó    Úinpt_tensorÚog_sizec                 óB   • U R                  S5      U:w  a	  U SSU24   $ U $ )z'Handles the unpad of the last dimensionéÿÿÿÿ.N)Úsize)r   r   s     r   Ú_postprocess_flash_outputr      s.   € à×Ñ˜Ó˜wÓ&Ø˜3   ˜=Ñ)Ð)ØÐr   Úhead_dim_sizeÚscalec                 ó>   • Ub  U$ S[         R                  " U 5      -  $ )z‘
For FlashAttention we pad the head dimension to be a multiple of 8 so we need to scale the output
by the original head size and not the padded.
g      ð?)ÚmathÚsqrt)r   r   s     r   Ú_calculate_scaler      s#   € ð
 ÑØˆØ”—’˜=Ó)Ñ)Ð)r   ÚqueryÚkeyÚvalueÚ	attn_maskc           	      ó¤  • U(       dg  U R                   UR                   :w  d  U R                   UR                   :w  a3  [        SU R                    SUR                    SUR                    S35      eU R                  UR                  :w  d  U R                  UR                  :w  a3  [        SU R                   SUR                   SUR                   S35      eU R                  5       S:  d(  UR                  5       S:  d  UR                  5       S:  a?  [        S	U R                  5        S
UR                  5        SUR                  5        S35      eg )NzLExpected query, key, and value to have the same dtype, but got query.dtype: z, key.dtype: z, and value.dtype: z	 instead.zSExpected query, key, and value to have the same device type, but got query.device: z, key.device: z, and value.device: é   zUExpected query, key, and value to all be  at least 2 dimensional, but got query.dim: z, key.dim: z and value.dim: )ÚdtypeÚ
ValueErrorÚdeviceÚdim)r   r   r   r    Ú	dropout_pÚ	is_causalr   Úallow_lowp_kvs           r   Ú_validate_sdpa_inputr*   "   s%  € ö Ø�;‰;˜#Ÿ)™)Ó# u§{¡{°e·k±kÓ'AÜð(Ø(-¯© }°MÀ#Ç)Á)Àð M$Ø$)§K¡K =°	ð;óð ð
 ‡|�|�s—z‘zÓ! U§\¡\°U·\±\Ó%AÜð%Ø%*§\¡\ N°.ÀÇÁÀð M!Ø!&§¡ ¨ið9ó
ð 	
ð
 ‡y�yƒ{�Qƒ˜#Ÿ'™'›) a›-¨5¯9©9«;¸«?ÜØcØ�y‰y‹{ˆm˜; s§w¡w£y kÐ1AÀ%Ç)Á)Ã+ÀÈiðYó
ð 	
ð ,;r   )Ng        FNF)Ú__doc__r   Útorchr   ÚlistÚstrÚ__annotations__ÚTensorÚboolr   Úintr   Úfloatr   r*   © r   r   Ú<module>r5      sÎ   ðâ Iã ã ð €ˆˆc‰Ó ð1 5§<¡<ð 1°Dô 1ð
¨5¯<©<ð À#ð È%Ï,É,ô ð* Cð *°¸±ð *Àô *ð &*ØØØ
Øñ
Ø�<‰<ð
à	�‰ð
ð �<‰<ð
ð �|‰|˜dÑ"ð	
ð 
ö
r   