§
    �ŠtjØ5  ã                   ó®  — d Z ddlmZmZ ddlmZ ddlZddlmc m	Z
 ddlmZmZmZmZ ddlmZ ddlmZmZmZmZ g d¢Zej                             e¦  «         ej                             e¦  «         ej                             e¦  «         ej                             e¦  «          G d	„ d
e¦  «        Z G d„ dej        ¦  «        Zdefd„Zdefd„ZdS )zCDefines bias subclasses that work with scaled_dot_product_attentioné    )ÚautoÚIntEnum)ÚwarnN)Úcan_use_efficient_attentionÚcan_use_flash_attentionÚis_flash_attention_availableÚ
SDPAParams)Ú_raise_kernel_warnings)Ú_calculate_scaleÚ_input_requires_gradÚ_postprocess_flash_outputÚ_validate_sdpa_input)Úcausal_upper_leftÚcausal_lower_rightÚCausalVariantÚ
CausalBiasc                   ó:   — e Zd ZdZ e¦   «         Z e¦   «         ZdS )r   a+  
    Enum for causal variants used in attention mechanisms.

    Defines two types of causal biases:

    ``UPPER_LEFT``: Represents upper-left triangular bias for standard causal attention.
    The equivalent pytorch code for constructing this bias is:

    .. code-block:: python

        torch.tril(torch.ones(size, dtype=torch.bool))

    For instance, with ``shape=(3,4)``, the materialized bias tensor will be:

    .. code-block:: text

        [[1, 0, 0, 0],
         [1, 1, 0, 0],
         [1, 1, 1, 0]]


    ``LOWER_RIGHT``: Represents lower-right triangular bias, the include values are aligned to the lower
    right corner of the matrix.

    The equivalent pytorch code for constructing this bias is:

    .. code-block:: python

        diagonal_offset = size[1] - size[0]
        torch.tril(
            torch.ones(size, dtype=torch.bool),
            diagonal=diagonal_offset,
        )

    For instance, with ``shape=(3,4)``, the materialized bias tensor will be:

    .. code-block:: text

        [[1, 1, 0, 0],
         [1, 1, 1, 0],
         [1, 1, 1, 1]]

    Note that these variants are equivalent to each other when the sequence lengths of the query and key/value
    tensors are equal since the triangular matrix is square.

    .. warning:: This enum is a prototype and subject to change.
    N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ú
UPPER_LEFTÚLOWER_RIGHT© ó    úU/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/nn/attention/bias.pyr   r   !   s1   € € € € € ð.ð .ð` �‘”€JØ�$‘&”&€K€K€Kr   r   c                   óD  ‡ — e Zd ZdZdedededdfˆ fd„Zdej        dej	        fd	„Z
dej        dej	        fd
„Zddej        dz  dej	        fd„Ze	 	 	 	 ddej	        dej	        dej	        dd dedededz  dedej	        fd„¦   «         Zedˆ fd„	¦   «         Zdefd„Zˆ xZS )r   aP  
    A bias representing causal attention patterns. For an overview of the bias structure, see the :class:`CausalVariant` enum.

    This class is used for defining causal (triangular) attention biases. For constructing the bias, there exist
    two factory functions: :func:`causal_upper_left` and :func:`causal_lower_right`.

    Example:

    .. code-block:: python

        from torch.nn.attention.bias import causal_lower_right

        bsz, num_heads, seqlen_q, seqlen_kv, head_dim = 32, 8, 4, 12, 8

        # Create a lower-right causal bias
        attn_bias = causal_lower_right(seqlen_q, seqlen_kv)

        q = torch.randn(
            bsz, num_heads, seqlen_q, head_dim, device="cuda", dtype=torch.float16
        )
        k = torch.randn(
            bsz, num_heads, seqlen_kv, head_dim, device="cuda", dtype=torch.float16
        )
        v = torch.randn(
            bsz, num_heads, seqlen_kv, head_dim, device="cuda", dtype=torch.float16
        )

        out = F.scaled_dot_product_attention(q, k, v, attn_bias)

    .. warning:: This class is a prototype and subject to change.
    ÚvariantÚ	seq_len_qÚ
seq_len_kvÚreturnNc                 ó:  •— t          |t          ¦  «        s$t          dt          |¦  «        j        › �¦  «        ‚t          ¦   «                              ¦   «          || _        || _        || _	        ||k    r#|t          j
        k    rt          dd¬¦  «         dS dS dS )aÞ  
        Initializes the CausalBias instance with a specified variant and sequence lengths.

        Args:
            variant (CausalVariant): The type of causal bias to use (either UPPER_LEFT or LOWER_RIGHT).
            seq_len_q (int): The sequence length of the query tensor.
            seq_len_kv (int): The sequence length of the key/value tensor.

        Raises a warning if the LOWER_RIGHT variant is used with seq_len_q > seq_len_kv, as it may produce NaNs.
        z%variant must be a CausalVariant, got zTLower right causal bias will produce NaNs in the output when seq_len_q > seq_len_kv!é   )Ú
stacklevelN)Ú
isinstancer   ÚAssertionErrorÚtyper   ÚsuperÚ__init__r   r   r    r   r   )Úselfr   r   r    Ú	__class__s       €r   r)   zCausalBias.__init__w   s·   ø€ õ ˜'¥=Ñ1Ô1ð 	Ý ØP½¸W¹¼Ô8NÐPÐPñô ð õ 	‰Œ×ÒÑÔÐØˆŒØ"ˆŒØ$ˆŒØ�zÒ!Ð! gµÔ1JÒ&JÐ&JÝØfØðñ ô ð ð ð ð "Ð!Ð&JÐ&Jr   Údevicec                 ó~   — t          j        t          j        | j        | j        |t           j        ¬¦  «        ¦  «        S )zUpper left causal bias©r,   Údtype)ÚtorchÚtrilÚonesr   r    Úbool©r*   r,   s     r   Ú_upper_leftzCausalBias._upper_left�   s2   € åŒzÝŒJ�t”~ t¤¸vÍUÌZÐXÑXÔXñ
ô 
ð 	
r   c                 ó    — | j         | j        z
  }t          j        t          j        | j        | j         |t          j        ¬¦  «        |¬¦  «        S )zLower right causal biasr.   )Údiagonal)r    r   r0   r1   r2   r3   )r*   r,   Údiagonal_offsets      r   Ú_lower_rightzCausalBias._lower_right–   sQ   € àœ/¨D¬NÑ:ˆÝŒzÝŒJØ” ¤¸ÅeÄjðñ ô ð %ð	
ñ 
ô 
ð 	
r   c                 óÚ   — |€t          j        d¦  «        }| j        t          j        k    r|                      |¦  «        S | j        t          j        k    r|                      |¦  «        S dS )a˜  
        Materializes the causal bias into a tensor form.

        Depending on the variant, this method generates either an upper-left or lower-right
        triangular matrix to represent the causal bias.

        Args:
            device (Optional[torch.device]): The device on which to create the tensor. Defaults to CPU.

        Returns:
            torch.Tensor: The materialized bias tensor.
        NÚcpu)r0   r,   r   r   r   r5   r   r9   r4   s     r   Ú_materializezCausalBias._materialize¡   sh   € ð ˆ>Ý”\ %Ñ(Ô(ˆFØŒ<�=Ô3Ò3Ð3Ø×#Ò# FÑ+Ô+Ð+ØŒ\�]Ô6Ò6Ð6Ø×$Ò$ VÑ,Ô,Ð,ð 7Ð6r   ç        FÚqueryÚkeyÚvalueÚ	attn_maskÚ	dropout_pÚ	is_causalÚscaleÚ
enable_gqac                 óx  — |rt          d¦  «        ‚|j        |j        k    s|j        t          j        k    rt          j        | ||d|d||¬¦  «        S |j        t          j        k    �r<t          | ||d|||¦  «         t          | ||d|||¦  «        }t          |¦  «        �r| j        j        dk    rdnd}	|                      d¦  «        }
t          |
|¦  «        }|
|	z  d	k    }|r}|	|
|	z  z
  }t           j        j                             | d	|f¦  «        } t           j        j                             |d	|f¦  «        }t           j        j                             |d	|f¦  «        }t           j        j                             | |||dd
|¬¦  «        d	         }t/          ||
¦  «        S t1          |¦  «        r®d
}t3          | ||¦  «        rd}t           j        j                             |                      dd¦  «        |                     dd¦  «        |                     dd¦  «        ddddd|t9          |j        ¦  «        ||d¬¦  «        d	                              dd¦  «        S t;          |¦  «         t          j        | |||                     | j        ¦  «        |d
||¬¦  «        S t          d|j        › �¦  «        ‚)a8  
        Handles the logic for computing attention with the specified causal bias.

        Args:
            query (Tensor): Query tensor; shape :math:`(N, ..., L, E)`.
            key (Tensor): Key tensor; shape :math:`(N, ..., S, E)`.
            value (Tensor): Value tensor; shape :math:`(N, ..., S, Ev)`.
            attn_mask (CausalBias): The type of causal attention to apply.
                A boolean mask where a value of True indicates that the element *should* take part in attention.
                A float mask of the same type as query, key, value that is added to the attention score.
            dropout_p (float): Dropout probability; if greater than 0.0, dropout is applied
            is_causal (bool): If true, assumes upper left causal attention masking and errors if both attn_mask and is_causal
                are set.
            scale (optional float): Scaling factor applied prior to softmax. If None, the default value is set
                to :math:`\frac{1}{\sqrt{E}}`.
            enable_gqa (optional bool): If set to True, Grouped Query Attention (GQA) is enabled, by default it is set to False.

        Returns:
            output (Tensor): Attention output; shape :math:`(N, ..., L, Ev)`.

        Raises:
            ValueError: If the causal bias variant is not a CausalVariant type.

        z.CausalBias should not be used with causal=TrueNT)rA   rB   rC   rD   rE   Úxpué@   é   éÿÿÿÿr   F)rC   Úreturn_debug_maskrD   é   r#   )
ÚbiasÚcu_seqlens_qÚcu_seqlens_kÚmax_seqlen_qÚmax_seqlen_krB   Úcustom_mask_typeÚcompute_log_sumexprD   Úseqlen_kz<CausalBias.variant must be a CausalVariant type, but found: )Ú
ValueErrorr   r    r   r   r   ÚFÚscaled_dot_product_attentionr   r   r	   r   r,   r'   Úsizer   r0   ÚnnÚ
functionalÚpadÚopsÚatenÚ#_scaled_dot_product_flash_attentionr   r   r   Ú_efficient_attention_forwardÚ	transposeÚintr
   r<   )r>   r?   r@   rA   rB   rC   rD   rE   Úsdpa_paramsÚ	alignmentÚog_head_sizeÚog_scaleÚneeds_paddingÚpad_lenÚoutrS   s                   r   Ú	_dispatchzCausalBias._dispatchµ   s  € ðF ð 	OÝÐMÑNÔNÐNð Ô 9Ô#7Ò7Ð7ØÔ ¥MÔ$<Ò<Ð<åÔ1ØØØØØ#ØØØ%ð	ñ 	ô 	ð 	ð Ô¥-Ô";Ò;Ñ;Ý  ¨¨U°D¸)ÀYÐPUÑVÔVÐVÝ$Ø�s˜E 4¨°I¸zñô ˆKõ ' {Ñ3Ô3ñ DØ"'¤,Ô"3°uÒ"<Ð"<˜B˜BÀ!�	Ø$Ÿzšz¨"™~œ~�Ý+¨L¸%Ñ@Ô@�Ø ,¨yÑ 8¸AÒ =�Ø ð IØ'¨<¸)Ñ+CÑD�GÝ!œHÔ/×3Ò3°E¸A¸w¸<ÑHÔH�EÝœ(Ô-×1Ò1°#¸¸7°|ÑDÔD�CÝ!œHÔ/×3Ò3°E¸A¸w¸<ÑHÔH�EÝ”i”n×HÒHØØØØØ"Ø&+Ø"ð Iñ ô ð ô�õ 1°°lÑCÔCÐCÝ*¨;Ñ7Ô7ð Ø%*Ð"Ý'¨¨s°EÑ:Ô:ð .Ø)-Ð&Ý”y”~×BÒBØ—O’O A qÑ)Ô)Ø—M’M ! QÑ'Ô'Ø—O’O A qÑ)Ô)ØØ!%Ø!%Ø!%Ø!%Ø'Ý%(¨Ô):Ñ%;Ô%;Ø'9ØØ!ð Cñ ô ð ô÷ ’Y˜q !‘_”_ð%õ  ' {Ñ3Ô3Ð3åÔ5ØØØØ'×4Ò4°U´\ÑBÔBØ'Ø#ØØ)ð	ñ 	ô 	ð 	õ ØbÈyÔO`ÐbÐbñô ð r   r   c                 óž   •— |€i }|t           j        j        j        u r | j        |i |¤ŽS t          ¦   «                              ||||¦  «        S )zjDefines the behavior of torch.nn.functional.scaled_dot_product_attention when the attn_bias is an AttnBias)r0   rY   rZ   rW   ri   r(   Ú__torch_function__)ÚclsÚfuncÚtypesÚargsÚkwargsr+   s        €r   rk   zCausalBias.__torch_function__'  sW   ø€ ð ˆ>ØˆFØ•5”8Ô&ÔCÐCÐCØ �3”= $Ð1¨&Ð1Ð1Ð1Ý‰wŒw×)Ò)¨$°°t¸VÑDÔDÐDr   c                 óN   — |                       ¦   «                              ¦   «         S ©N)r<   Ú__repr__)r*   s    r   rs   zCausalBias.__repr__0  s    € Ø× Ò Ñ"Ô"×+Ò+Ñ-Ô-Ð-r   rr   )r=   FNF)r   N)r   r   r   r   r   ra   r)   r0   r,   ÚTensorr5   r9   r<   ÚstaticmethodÚfloatr3   ri   Úclassmethodrk   Ústrrs   Ú__classcell__)r+   s   @r   r   r   V   sÐ  ø€ € € € € ðð ð@ ð ¸#ð È3ð ÐSWð ð ð ð ð ð ð2
 %¤,ð 
°5´<ð 
ð 
ð 
ð 
ð
 5¤<ð 
°E´Lð 
ð 
ð 
ð 
ð-ð - 5¤<°$Ñ#6ð -À%Ä,ð -ð -ð -ð -ð( ð ØØ"Ø ðoð oØŒ|ðoàŒ\ðoð Œ|ðoð  ð	oð
 ðoð ðoð �t‰|ðoð ðoð 
Œðoð oð oñ „\ðoðb ðEð Eð Eð Eð Eñ „[ðEð.˜#ð .ð .ð .ð .ð .ð .ð .ð .r   r   r!   c                  ó†   — t          | ¦  «        dk    rt          d¦  «        ‚| \  }}t          t          j        ||¦  «        S )a&  
    Creates an upper-left triangular causal bias.

    This function generates a upper-left triangular matrix to represent causal attention bias with a
    diagonal offset set so that the inclusive values are aligned to the upper left corner of the matrix.
    This equivalent to the `is_causal=True` argument in `scaled_dot_product_attention`.

    The equivalent pytorch code for constructing this bias is:

    .. code-block:: python

        torch.tril(torch.ones(size, dtype=torch.bool))

    For instance, with `shape=(3,4)`, the materialized bias tensor will be:

    .. code-block:: text

        [[1, 0, 0, 0],
         [1, 1, 0, 0],
         [1, 1, 1, 0]]

    Args:
        size: The size of the bias matrix.

    Returns:
        CausalBias: The UPPER_LEFT triangular causal bias variant.
    r#   z*causal_upper_left only supports 2D tensors)Úlenr&   r   r   r   ©rX   r   r    s      r   r   r   4  sA   € õ8 ˆ4�y„y�A‚~€~ÝÐIÑJÔJÐJØ Ñ€IˆzÝ•mÔ.°	¸:ÑFÔFÐFr   c                  ó†   — t          | ¦  «        dk    rt          d¦  «        ‚| \  }}t          t          j        ||¦  «        S )a:  
    Creates a lower-right triangular causal bias.

    This function generates a lower-right triangular matrix to represent causal attention bias with a
    diagonal offset set so that the inclusive values are aligned to the lower right corner of the matrix.

    The equivalent pytorch code for constructing this bias is:

    .. code-block:: python

        diagonal_offset = size[1] - size[0]
        torch.tril(
            torch.ones(size, dtype=torch.bool),
            diagonal=diagonal_offset,
        )

    For instance, with `shape=(3,4)`, the materialized bias tensor will be:

    .. code-block:: text

        [[1, 1, 0, 0],
         [1, 1, 1, 0],
         [1, 1, 1, 1]]

    Args:
        size: The size of the bias matrix.

    Returns:
        CausalBias: The LOWER_RIGHT triangular causal bias variant.
    r#   z+causal_lower_right only supports 2D tensors)r{   r&   r   r   r   r|   s      r   r   r   V  sA   € õ> ˆ4�y„y�A‚~€~ÝÐJÑKÔKÐKØ Ñ€IˆzÝ•mÔ/°¸JÑGÔGÐGr   )r   Úenumr   r   Úwarningsr   r0   Útorch.nn.functionalrY   rZ   rV   Útorch.backends.cudar   r   r   r	   Útorch.nn.attentionr
   Útorch.nn.attention._utilsr   r   r   r   Ú__all__Ú_dynamoÚallow_in_graphr   rt   r   r   r   r   r   r   ú<module>r‡      sý  ðà IÐ Ià Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à €€€Ø Ð Ð Ð Ð Ð Ð Ð Ð ðð ð ð ð ð ð ð ð ð ð ð ð 6Ð 5Ð 5Ð 5Ð 5Ð 5ðð ð ð ð ð ð ð ð ð ð ð ð UÐ
TÐ
T€ð „× Ò Ð9Ñ :Ô :Ð :Ø „× Ò Ð4Ñ 5Ô 5Ð 5Ø „× Ò Ð8Ñ 9Ô 9Ð 9Ø „× Ò ˜ZÑ (Ô (Ð (ð2ð 2ð 2ð 2ð 2�Gñ 2ô 2ð 2ðj[.ð [.ð [.ð [.ð [.�”ñ [.ô [.ð [.ð|G 
ð Gð Gð Gð GðD"H ð "Hð "Hð "Hð "Hð "Hð "Hr   