Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
attention.cpython-38.pyc204 linesDownload Raw Back to __pycache__
1U

2E$�c\Y�@sddlZddlmZmZddlZddlmZddlmZddlm	Z	m3Z4mZddlm
Z
ddlmZmZddlmZmZmZmZdeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeeefd5�dd�ZGd
d�dej�ZGdd�dej�ZdS)�N)�Optional�Tuple)�Tensor)�	constant_�xavier_normal_�xavier_uniform_)�	Parameter)�has_torch_function�handle_torch_function)�pad�linear�softmax�dropoutTF)�query�key�value�embed_dim_to_check�	num_heads�in_proj_weight�in_proj_bias�bias_k�bias_v�
add_zero_attn�	dropout_p�out_proj_weight�
out_proj_bias�training�key_padding_mask�need_weights�	attn_mask�use_separate_proj_weight�
q_proj_weight�
k_proj_weight�
v_proj_weight�static_k�static_v�returnc,Cs�|||||||||f	}t|�rXtt|||||||||||	|6|||
|||||||||d�S|��\}}}||ksrt�|�d�|�d�kr�|�d�|�d�ks�t�||}|||ks�td��t|�d}|�s�||ks�t�||��r||ks�t�||��rt|||�j	ddd�\}}}�q�||k�s0t�||��r�|} d}!|}"||!|"�d	d	�f}#| d	k	�rf| |!|"�} t||#| �}|d	k�r�|d	k�s�t�d	}d	}nP|} |}!d	}"||!d	�d	d	�f}#| d	k	�r�| |!d	�} t||#| �j	d7dd�\}}n�|} d}!|}"||!|"�d	d	�f}#| d	k	�r| |!|"�} t||#| �}|} |}!|d8}"||!|"�d	d	�f}#| d	k	�rb| |!|"�} t||#| �}|} |d9}!d	}"||!d	�d	d	�f}#| d	k	�r�| |!d	�} t||#| �}�ntj10�|�}$|$��\}%}&|%|k�r�|&|�d�k�s�t�tj11�|�}'|'��\}%}&|%|k�r |&|�d�k�s$t�tj12�|�}(|(��\}%}&|%|k�rV|&|�d�k�sZt�|d	k	�r�t||$|d|��}t||'|||d13��}t||(||d14d	��}n$t||$|�}t||'|�}t||(|�}||}|d	k	�r�|jtj
k�s6|jtjk�s6|jtjk�s6|jtjk�s6|jtjk�s6td�|j���|jtjk�rZt�d�|�tj�}|��d15k�r�|�d�}t|���d|�d�|�d�gk�r�td
��nR|��dk�r�t|���|||�d�|�d�gk�r�td��ntd�|�����|d	k	�r |jtjk�r t�d�|�tj�}|d	k	�r�|d	k	�r�|d	k�r�|d	k�r�t�||�d|d�g�}t�||�d|d�g�}|d	k	�r�t|d�}|d	k	�r�t|d�}n$|d	k�s�td��|d	k�s�td��n|d	k�s�t�|d	k�s�t�|���||||��dd�}|d	k	�r*|���d|||��dd�}|d	k	�rR|���d|||��dd�}|d	k	�r�|�d�||k�stt�|�d16�|k�s�t�|}|d	k	�r�|�d�||k�s�t�|�d17�|k�s�t�|}|�d�})|d	k	�r�|�d�|)k�s�t�|	�r�|)d7})tj|tj |�d�df|��d18d	�|j|j!d�gdd�}tj|tj |�d�df|��d19d	�|j|j!d�gdd�}|d	k	�r�t|d�}|d	k	�r�t|d�}t�"||�dd20��}*t|*���||||)gk�s�t�|d	k	�r|jtjk�r�|*�#|td��n|*|7}*|d	k	�rD|*�||||)�}*|*�$|�d�td��}*|*�||||)�}*t%|*dd�}*t&|*|21|
d�}*t�"|*|�}+t|+���||||gk�s�t�|+�dd����|||�}+t|+||�}+|�r�|*�||||)�}*|+|*j'dd�|fS|+d	fSd	S)a?22    Args:23        query, key, value: map a query and a set of key-value pairs to an output.24            See "Attention Is All You Need" for more details.25        embed_dim_to_check: total dimension of the model.26        num_heads: parallel attention heads.27        in_proj_weight, in_proj_bias: input projection weight and bias.28        bias_k, bias_v: bias of the key and value sequences to be added at dim=0.29        add_zero_attn: add a new batch of zeros to the key and30                       value sequences at dim=1.31        dropout_p: probability of an element to be zeroed.32        out_proj_weight, out_proj_bias: the output projection weight and bias.33        training: apply dropout if is ``True``.34        key_padding_mask: if provided, specified padding elements in the key will35            be ignored by the attention. This is an binary mask. When the value is True,36            the corresponding value on the attention layer will be filled with -inf.37        need_weights: output attn_output_weights.38        attn_mask: 2D or 3D mask that prevents attention to certain positions. A 2D mask will be broadcasted for all39            the batches while a 3D mask allows to specify a different mask for the entries of each batch.40        use_separate_proj_weight: the function accept the proj. weights for query, key,41            and value in different forms. If false, in_proj_weight will be used, which is42            a combination of q_proj_weight, k_proj_weight, v_proj_weight.43        q_proj_weight, k_proj_weight, v_proj_weight, in_proj_bias: input projection weight and bias.44        static_k, static_v: static key and value used for attention operators.45 46 47    Shape:48        Inputs:49        - query: :math:`(L, N, E)` where L is the target sequence length, N is the batch size, E is50          the embedding dimension.51        - key: :math:`(S, N, E)`, where S is the source sequence length, N is the batch size, E is52          the embedding dimension.53        - value: :math:`(S, N, E)` where S is the source sequence length, N is the batch size, E is54          the embedding dimension.55        - key_padding_mask: :math:`(N, S)` where N is the batch size, S is the source sequence length.56          If a ByteTensor is provided, the non-zero positions will be ignored while the zero positions57          will be unchanged. If a BoolTensor is provided, the positions with the58          value of ``True`` will be ignored while the position with the value of ``False`` will be unchanged.59        - attn_mask: 2D mask :math:`(L, S)` where L is the target sequence length, S is the source sequence length.60          3D mask :math:`(N*num_heads, L, S)` where N is the batch size, L is the target sequence length,61          S is the source sequence length. attn_mask ensures that position i is allowed to attend the unmasked62          positions. If a ByteTensor is provided, the non-zero positions are not allowed to attend63          while the zero positions will be unchanged. If a BoolTensor is provided, positions with ``True``64          are not allowed to attend while ``False`` values will be unchanged. If a FloatTensor65          is provided, it will be added to the attention weight.66        - static_k: :math:`(N*num_heads, S, E/num_heads)`, where S is the source sequence length,67          N is the batch size, E is the embedding dimension. E/num_heads is the head dimension.68        - static_v: :math:`(N*num_heads, S, E/num_heads)`, where S is the source sequence length,69          N is the batch size, E is the embedding dimension. E/num_heads is the head dimension.70 71        Outputs:72        - attn_output: :math:`(L, N, E)` where L is the target sequence length, N is the batch size,73          E is the embedding dimension.74        - attn_output_weights: :math:`(N, L, S)` where N is the batch size,75          L is the target sequence length, S is the source sequence length.76    )77rrrrr r!r"r#r$r%r��(embed_dim must be divisible by num_headsg�������)�dimN�zDOnly float, byte, and bool types are supported for attn_mask, not {}zZByte tensor for attn_mask in nn.MultiheadAttention is deprecated. Use bool tensor instead.z,The size of the 2D attn_mask is not correct.z,The size of the 3D attn_mask is not correct.z)attn_mask's dimension {} is not supportedzaByte tensor for key_padding_mask in nn.MultiheadAttention is deprecated. Use bool tensor instead.)rr'z#bias cannot be added to static key.z%bias cannot be added to static value.)�dtype�devicez-inf)�pr)(r	r78�multi_head_attention_forward�size�AssertionError�float�torch�equalr�chunk�jit�_unwrap_optionalr-�float32�float64�float16�uint8�bool�format�warnings�warn�tor+�	unsqueeze�list�RuntimeError�cat�repeatr�79contiguous�view�	transpose�zerosr.�bmm�masked_fill_�masked_fillr
r�sum),rrrrrrrrrrrrrrrrrr r!r"r#r$r%�tens_ops�tgt_len�bsz�	embed_dim�head_dim�scaling�q�k�v�_b�_start�_end�_w�q_proj_weight_non_opt�len1�len2�k_proj_weight_non_opt�v_proj_weight_non_opt�src_len�attn_output_weights�attn_output�rd�3/data/arXiv/demo/Demo/xdecoder/modules/attention.pyr0snQ�,, 808182838485868788�89�90�91�92�93�9495$96(97�9899100101102103104105106<<107108109110 111112� r0cs0eZdZUeed<eedd��fdd�Z�ZS)�_LinearWithBias�biasN)�in_features�out_featuresr&cst�j||dd�dS)NT)rg)�super�__init__)�selfrhri��	__class__rdrerkIsz_LinearWithBias.__init__)�__name__�113__module__�__qualname__r�__annotations__�intrk�
__classcell__rdrdrmrerfFs114rfcs�eZdZUdZeejed<eejed<d�fdd	�	Zd115d�Z	�fdd
�Z116deeeeeeeeeeeefd�dd�Z
�ZS)�MultiheadAttentiona�Allows the model to jointly attend to information117    from different representation subspaces.118    See `Attention Is All You Need <https://arxiv.org/abs/1706.03762>`_119 120    .. math::121        \text{MultiHead}(Q, K, V) = \text{Concat}(head_1,\dots,head_h)W^O122 123    where :math:`head_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)`.124 125    Args:126        embed_dim: total dimension of the model.127        num_heads: parallel attention heads.128        dropout: a Dropout layer on attn_output_weights. Default: 0.0.129        bias: add bias as module parameter. Default: True.130        add_bias_kv: add bias to the key and value sequences at dim=0.131        add_zero_attn: add a new batch of zeros to the key and132                       value sequences at dim=1.133        kdim: total number of features in key. Default: None.134        vdim: total number of features in value. Default: None.135 136    Note that if :attr:`kdim` and :attr:`vdim` are None, they will be set137    to :attr:`embed_dim` such that query, key, and value have the same138    number of features.139 140    Examples::141 142        >>> multihead_attn = nn.MultiheadAttention(embed_dim, num_heads)143        >>> attn_output, attn_output_weights = multihead_attn(query, key, value)144    rr�TFNc		s�tt|���||_|dk	r |n||_|dk	r2|n||_|j|koJ|j|k|_||_||_|||_	|j	||jks|t145d��|jdkr�tt�
||��|_tt�
||j��|_tt�
||j��|_|�dd�n:tt�d||��|_|�dd�|�dd�|�dd�|�r$tt�d|��|_n|�dd�t||�|_|�rltt�d	d	|��|_tt�d	d	|��|_nd|_|_||_|��dS)146Nr(Frr)r!r"r#rr')rjrurkrR�kdim�vdim�_qkv_same_embed_dimrrrSr2rr4rr!r"r#�register_parameter�emptyrrrf�out_projrrr�_reset_parameters)	rlrRrrrg�add_bias_kvrrwrxrmrdrerkns8147148zMultiheadAttention.__init__cCs�|jrt|j�nt|j�t|j�t|j�|jdk	rTt|jd�t|jj	d�|j149dk	rht|j150�|jdk	r|t|j�dS)Nrv)
ryrrr!r"r#rrr|rgrrr)rlrdrdrer}�s151152153154155156157z$MultiheadAttention._reset_parameterscs$d|krd|d<tt|��|�dS)NryT)rjru�__setstate__)rl�statermrdrer�szMultiheadAttention.__setstate__)rrrrrrr&cCs�|jsXt||||j|j|j|j|j|j|j|j	|j158j|j159j|j
|||d|j|j|jd�St||||j|j|j|j|j|j|j|j	|j160j|j161j|j
|||d�SdS)a�162163    Args:164        query, key, value: map a query and a set of key-value pairs to an output.165            See "Attention Is All You Need" for more details.166        key_padding_mask: if provided, specified padding elements in the key will167            be ignored by the attention. When given a binary mask and a value is True,168            the corresponding value on the attention layer will be ignored. When given169            a byte mask and a value is non-zero, the corresponding value on the attention170            layer will be ignored171        need_weights: output attn_output_weights.172        attn_mask: 2D or 3D mask that prevents attention to certain positions. A 2D mask will be broadcasted for all173            the batches while a 3D mask allows to specify a different mask for the entries of each batch.174 175    Shapes for inputs:176        - query: :math:`(L, N, E)` where L is the target sequence length, N is the batch size, E is177          the embedding dimension.178        - key: :math:`(S, N, E)`, where S is the source sequence length, N is the batch size, E is179          the embedding dimension.180        - value: :math:`(S, N, E)` where S is the source sequence length, N is the batch size, E is181          the embedding dimension.182        - key_padding_mask: :math:`(N, S)` where N is the batch size, S is the source sequence length.183          If a ByteTensor is provided, the non-zero positions will be ignored while the position184          with the zero positions will be unchanged. If a BoolTensor is provided, the positions with the185          value of ``True`` will be ignored while the position with the value of ``False`` will be unchanged.186        - attn_mask: if a 2D mask: :math:`(L, S)` where L is the target sequence length, S is the187          source sequence length.188 189          If a 3D mask: :math:`(N\cdot\text{num\_heads}, L, S)` where N is the batch size, L is the target sequence190          length, S is the source sequence length. ``attn_mask`` ensure that position i is allowed to attend191          the unmasked positions. If a ByteTensor is provided, the non-zero positions are not allowed to attend192          while the zero positions will be unchanged. If a BoolTensor is provided, positions with ``True``193          is not allowed to attend while ``False`` values will be unchanged. If a FloatTensor194          is provided, it will be added to the attention weight.195 196    Shapes for outputs:197        - attn_output: :math:`(L, N, E)` where L is the target sequence length, N is the batch size,198          E is the embedding dimension.199        - attn_output_weights: :math:`(N, L, S)` where N is the batch size,200          L is the target sequence length, S is the source sequence length.201        T)rrrrr r!r"r#)rrrrN)ryr0rRrrrrrrrr|�weightrgrr!r"r#)rlrrrrrrrdrdre�forward�sV*��zMultiheadAttention.forward)rvTFFNN)NTN)rorprq�__doc__rr4rrrrkr}rr=rr�rtrdrdrmreruMs202'��ru)203TNTNFNNNNN)r?�typingrrr4�torch.nn�nnr�
torch.nn.initrrrZtorch.nn.parameterr�torch.overridesr	r204Ztorch.nn.functionalrrr
rrsr=r3r0�Linearrf�Modulerurdrdrdre�<module>s`��9