
    ^jo                     J   d Z ddlZddlmZ ddlZddlmZ ddlmZmZm	Z	 ddl
mZ ddlmZmZ dd	lmZ dd
lmZ ddlmZmZmZmZmZmZ ddlmZ ddlmZ ddlm Z  ddl!m"Z"m#Z#m$Z$m%Z% ddl&m'Z' ddl(m)Z) ddl*m+Z+  e%jX                  e-      Z. G d dej^                        Z0 G d dej^                        Z1 G d dej^                        Z2 G d dej^                        Z3 G d dej^                        Z4 G d dej^                        Z5 G d  d!ej^                        Z6 G d" d#ej^                        Z7 G d$ d%e      Z8e# G d& d'e             Z9 G d( d)ej^                        Z: G d* d+ej^                        Z; G d, d-ej^                        Z<e# G d. d/e9             Z= G d0 d1ej^                        Z>e# G d2 d3e9             Z? G d4 d5ej^                        Z@ e#d67       G d8 d9e9             ZAe# G d: d;e9             ZBe# G d< d=e9             ZCe# G d> d?e9             ZDg d@ZEy)AzPyTorch ConvBERT model.    N)Callable)nn)BCEWithLogitsLossCrossEntropyLossMSELoss   )initialization)ACT2FNget_activation)create_bidirectional_mask)GradientCheckpointingLayer)"BaseModelOutputWithCrossAttentionsMaskedLMOutputMultipleChoiceModelOutputQuestionAnsweringModelOutputSequenceClassifierOutputTokenClassifierOutput)PreTrainedModel)Unpack)apply_chunking_to_forward)TransformersKwargsauto_docstringcan_return_tuplelogging)merge_with_config_defaults)capture_outputs   )ConvBertConfigc                        e Zd ZdZ fdZ	 	 	 	 d
dej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  f
d	Z xZ	S )ConvBertEmbeddingszGConstruct the embeddings from word, position and token_type embeddings.c                    t         |           t        j                  |j                  |j
                  |j                        | _        t        j                  |j                  |j
                        | _	        t        j                  |j                  |j
                        | _        t        j                  |j
                  |j                        | _        t        j                  |j                        | _        | j#                  dt%        j&                  |j                        j)                  d      d       | j#                  dt%        j*                  | j,                  j/                         t$        j0                        d       y )	N)padding_idxepsposition_idsr   F)
persistenttoken_type_idsdtype)super__init__r   	Embedding
vocab_sizeembedding_sizepad_token_idword_embeddingsmax_position_embeddingsposition_embeddingstype_vocab_sizetoken_type_embeddings	LayerNormlayer_norm_epsDropouthidden_dropout_probdropoutregister_buffertorcharangeexpandzerosr%   sizelongselfconfig	__class__s     y/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/convbert/modeling_convbert.pyr-   zConvBertEmbeddings.__init__7   s   !||F,=,=v?T?Tbhbubuv#%<<0N0NPVPePe#f %'\\&2H2H&J_J_%`"f&;&;AVAVWzz&"<"<=ELL)G)GHOOPWXej 	 	
 	ekk$*;*;*@*@*B%**Ubg 	 	
    N	input_idsr)   r%   inputs_embedsreturnc                 2   ||j                         }n|j                         d d }|d   }|| j                  d d d |f   }|st        | d      r-| j                  d d d |f   }|j	                  |d   |      }|}n:t        j                  |t
        j                  | j                  j                        }|| j                  |      }| j                  |      }	| j                  |      }
||	z   |
z   }| j                  |      }| j                  |      }|S )Nr'   r   r)   r   r+   device)rA   r%   hasattrr)   r?   r=   r@   rB   rN   r2   r4   r6   r7   r;   )rD   rI   r)   r%   rJ   input_shape
seq_lengthbuffered_token_type_ids buffered_token_type_ids_expandedr4   r6   
embeddingss               rG   forwardzConvBertEmbeddings.forwardG   s,     #..*K',,.s3K ^
,,Q^<L
 !t-.*.*=*=a*n*M'3J3Q3QR]^_R`bl3m0!A!&[

SWSdSdSkSk!l  00;M"66|D $ : :> J"%88;PP
^^J/
\\*-
rH   )NNNN)
__name__
__module____qualname____doc__r-   r=   
LongTensorFloatTensorrU   __classcell__rF   s   @rG   r    r    4   s    Q
$ .2260426$##d*$ ((4/$ &&-	$
 ((4/$ 
		$rH   r    c                   Z     e Zd ZdZ fdZdej                  dej                  fdZ xZS )SeparableConv1DzSThis class implements separable convolution, i.e. a depthwise and a pointwise layerc                    t         |           t        j                  |||||dz  d      | _        t        j                  ||dd      | _        t        j                  t        j                  |d            | _	        | j                  j                  j                  j                  d|j                         | j
                  j                  j                  j                  d|j                         y )N   F)kernel_sizegroupspaddingbiasr   )rb   re           meanstd)r,   r-   r   Conv1d	depthwise	pointwise	Parameterr=   r@   re   weightdatanormal_initializer_range)rD   rE   input_filtersoutput_filtersrb   kwargsrF   s         rG   r-   zSeparableConv1D.__init__q   s    # 1$
 =.aV[\LL^Q!?@	""**9Q9Q*R""**9Q9Q*RrH   hidden_statesrK   c                 h    | j                  |      }| j                  |      }|| j                  z  }|S N)rk   rl   re   )rD   ru   xs      rG   rU   zSeparableConv1D.forward   s0    NN=)NN1	TYYrH   	rV   rW   rX   rY   r-   r=   TensorrU   r\   r]   s   @rG   r_   r_   n   s'    ]S U\\ ell rH   r_   c                        e Zd Z fdZ	 	 d	dej
                  dej                  dz  dej
                  dz  dee   de	ej
                  ej
                  f   f
dZ
 xZS )
ConvBertSelfAttentionc                 j   t         |           |j                  |j                  z  dk7  r2t	        |d      s&t        d|j                   d|j                   d      |j                  |j                  z  }|dk  r|j                  | _        d| _        n|| _        |j                  | _        |j                  | _        |j                  | j                  z  dk7  rt        d      |j                  | j                  z  dz  | _        | j                  | j                  z  | _	        t        j                  |j                  | j                        | _        t        j                  |j                  | j                        | _        t        j                  |j                  | j                        | _        t        ||j                  | j                  | j                        | _        t        j                  | j                  | j                  | j                  z        | _        t        j                  |j                  | j                        | _        t        j&                  | j                  dgt)        | j                  dz
  dz        dg	      | _        t        j,                  |j.                        | _        y )
Nr   r0   zThe hidden size (z6) is not a multiple of the number of attention heads ()r   z6hidden_size should be divisible by num_attention_headsra   )rb   rd   )r,   r-   hidden_sizenum_attention_headsrO   
ValueError
head_ratioconv_kernel_sizeattention_head_sizeall_head_sizer   Linearquerykeyvaluer_   key_conv_attn_layerconv_kernel_layerconv_out_layerUnfoldintunfoldr9   attention_probs_dropout_probr;   )rD   rE   new_num_attention_headsrF   s      rG   r-   zConvBertSelfAttention.__init__   s>    : ::a?PVXhHi#F$6$6#7 8 445Q8 
 #)"<"<@Q@Q"Q"Q&$88DO'(D$'>D$$//DO & 7 7 8 88A=UVV$*$6$6$:R:R$RWX#X !558P8PPYYv1143E3EF
99V//1C1CDYYv1143E3EF
#2F&&(:(:D<Q<Q$
  "$4+=+=t?W?WZ^ZoZo?o!p ii(:(:D<N<NOii..2S$BWBWZ[B[_`A`=acd<e
 zz&"E"EFrH   Nru   attention_maskencoder_hidden_statesrt   rK   c                     |j                   d d }g |d| j                  }|#| j                  |      }| j                  |      }n"| j                  |      }| j                  |      }| j	                  |j                  dd            }	|	j                  dd      }	| j                  |      }
|
j                  |      j                  dd      }|j                  |      j                  dd      }|j                  |      j                  dd      }t        j                  |	|
      }| j                  |      }t        j                  |d| j                  dg      }t        j                  |d      }| j                  |      }t        j                  ||d   d| j                  g      }|j                  dd      j!                         j#                  d      }t$        j&                  j)                  || j                  dgd| j                  dz
  dz  dgd      }|j                  dd      j                  |d   d| j                  | j                        }t        j                  |d| j                  | j                  g      }t        j*                  ||      }t        j                  |d| j                  g      }t        j*                  ||j                  dd            }|t-        j.                  | j                        z  }|||z   }t$        j&                  j                  |d      }| j1                  |      }t        j*                  ||      }|j3                  dddd      j!                         }t        j                  ||d   d| j4                  | j                  g      }t        j6                  ||gd      }|j9                         d d | j4                  | j                  z  dz  fz   } |j                  | }||fS )	Nr'   r   ra   dimr   )rb   dilationrd   strider   )shaper   r   r   r   	transposer   viewr=   multiplyr   reshaper   softmaxr   r   
contiguous	unsqueezer   
functionalr   matmulmathsqrtr;   permuter   catrA   )rD   ru   r   r   rt   rP   hidden_shapemixed_key_layermixed_value_layermixed_key_conv_attn_layermixed_query_layerquery_layer	key_layervalue_layerconv_attn_layerr   r   attention_scoresattention_probscontext_layerconv_outnew_context_layer_shapes                         rG   rU   zConvBertSelfAttention.forward   s    $))#2.CCbC$*B*BC !,"hh'<=O $

+@ A"hh}5O $

= 9$($<$<]=T=TUVXY=Z$[!$=$G$G1$M! JJ}5',,\:DDQJ#((6@@AF	',,\:DDQJ..)BDUV 22?C!MM*;b$BWBWYZ=[\!MM*;C,,];~ADL^L^7_`'11!Q7BBDNNrR--..2++a/A5q9 . 
 (11!Q7??NB 2 2D4I4I
 ~D<T<TVZVkVk7lmn6GH~D<N<N7OP !<<Y5H5HR5PQ+dii8P8P.QQ%/.@ --//0@b/I ,,7_kB%--aAq9DDF==[^R1I1I4KcKcd
 		=(";Q? #0"4"4"6s";$$t'?'??!C?
 #
 +**,CDo--rH   NN)rV   rW   rX   r-   r=   rz   r[   r   r   tuplerU   r\   r]   s   @rG   r|   r|      s}    %GT 4859	O.||O. ))D0O.  %||d2	O.
 +,O. 
u||U\\)	*O.rH   r|   c                   n     e Zd Z fdZdej
                  dej
                  dej
                  fdZ xZS )ConvBertSelfOutputc                 (   t         |           t        j                  |j                  |j                        | _        t        j                  |j                  |j                        | _        t        j                  |j                        | _
        y Nr#   )r,   r-   r   r   r   denser7   r8   r9   r:   r;   rC   s     rG   r-   zConvBertSelfOutput.__init__  s`    YYv1163E3EF
f&8&8f>S>STzz&"<"<=rH   ru   input_tensorrK   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S rw   r   r;   r7   rD   ru   r   s      rG   rU   zConvBertSelfOutput.forward	  7    

=1]3}|'CDrH   rV   rW   rX   r-   r=   rz   rU   r\   r]   s   @rG   r   r     s1    >U\\  RWR^R^ rH   r   c                        e Zd Z fdZ	 	 d	dej
                  dej                  dz  dej
                  dz  dee   dej
                  f
dZ	 xZ
S )
ConvBertAttentionc                 b    t         |           t        |      | _        t	        |      | _        y rw   )r,   r-   r|   rD   r   outputrC   s     rG   r-   zConvBertAttention.__init__  s&    )&1	(0rH   Nru   r   r   rt   rK   c                 \     | j                   ||fd|i|\  }}| j                  ||      }|S )Nr   )rD   r   )rD   ru   r   r   rt   r   _attention_outputs           rG   rU   zConvBertAttention.forward  sL     %499
 #8
 	
q  ;;}mDrH   r   )rV   rW   rX   r-   r=   rz   r[   r   r   rU   r\   r]   s   @rG   r   r     sg    1 4859	 ||  ))D0   %||d2	 
 +,  
 rH   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )GroupedLinearLayerc                    t         |           || _        || _        || _        | j                  | j                  z  | _        | j                  | j                  z  | _        t        j                  t        j                  | j                  | j
                  | j                              | _        t        j                  t        j                  |            | _        y rw   )r,   r-   
input_sizeoutput_size
num_groupsgroup_in_dimgroup_out_dimr   rm   r=   emptyrn   re   )rD   r   r   r   rF   s       rG   r-   zGroupedLinearLayer.__init__(  s    $&$ OOt>!--@ll5;;t@Q@QSWSeSe#fgLL[!9:	rH   ru   rK   c                    t        |j                               d   }t        j                  |d| j                  | j
                  g      }|j                  ddd      }t        j                  || j                        }|j                  ddd      }t        j                  ||d| j                  g      }|| j                  z   }|S )Nr   r'   r   ra   )listrA   r=   r   r   r   r   r   rn   r   re   )rD   ru   
batch_sizerx   s       rG   rU   zGroupedLinearLayer.forward2  s    -,,./2
MM-"doot?P?P)QRIIaALLDKK(IIaAMM!j"d.>.>?@		MrH   r   r]   s   @rG   r   r   '  s#    ;U\\ ell rH   r   c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )ConvBertIntermediatec                    t         |           |j                  dk(  r0t        j                  |j
                  |j                        | _        n1t        |j
                  |j                  |j                        | _        t        |j                  t              rt        |j                     | _        y |j                  | _        y )Nr   r   r   r   )r,   r-   r   r   r   r   intermediate_sizer   r   
isinstance
hidden_actstrr
   intermediate_act_fnrC   s     rG   r-   zConvBertIntermediate.__init__>  s    !6#5#5v7O7OPDJ+!--6;S;S`f`q`qDJ f''-'-f.?.?'@D$'-'8'8D$rH   ru   rK   c                 J    | j                  |      }| j                  |      }|S rw   )r   r   rD   ru   s     rG   rU   zConvBertIntermediate.forwardK  s&    

=100?rH   r   r]   s   @rG   r   r   =  s#    9U\\ ell rH   r   c                   n     e Zd Z fdZdej
                  dej
                  dej
                  fdZ xZS )ConvBertOutputc                    t         |           |j                  dk(  r0t        j                  |j
                  |j                        | _        n1t        |j
                  |j                  |j                        | _        t        j                  |j                  |j                        | _	        t        j                  |j                        | _        y )Nr   r   r#   )r,   r-   r   r   r   r   r   r   r   r7   r8   r9   r:   r;   rC   s     rG   r-   zConvBertOutput.__init__R  s    !6#;#;V=O=OPDJ+!33ASAS`f`q`qDJ f&8&8f>S>STzz&"<"<=rH   ru   r   rK   c                 r    | j                  |      }| j                  |      }| j                  ||z         }|S rw   r   r   s      rG   rU   zConvBertOutput.forward]  r   rH   r   r]   s   @rG   r   r   Q  s1    	>U\\  RWR^R^ rH   r   c                        e Zd Z fdZ	 	 	 ddej
                  dej                  dz  dej
                  dz  dej
                  dz  dee   dej
                  fd	Z	d
 Z
 xZS )ConvBertLayerc                 b   t         |           |j                  | _        d| _        t	        |      | _        |j                  | _        |j                  | _        | j                  r*| j                  st        |  d      t	        |      | _	        t        |      | _        t        |      | _        y )Nr   z> should be used as a decoder model if cross attention is added)r,   r-   chunk_size_feed_forwardseq_len_dimr   	attention
is_decoderadd_cross_attention	TypeErrorcrossattentionr   intermediater   r   rC   s     rG   r-   zConvBertLayer.__init__e  s    '-'E'E$*62 ++#)#=#= ##??4&(f ghh"3F";D08$V,rH   Nru   r   r   encoder_attention_maskrt   rK   c                     | j                   ||fi |}| j                  r3|1t        | d      st        d|  d       | j                  ||fd|i|}t        | j                  | j                  | j                  |      }|S )Nr   z'If `encoder_hidden_states` are passed, z` has to be instantiated with cross-attention layers by setting `config.add_cross_attention=True`r   )	r   r   rO   AttributeErrorr   r   feed_forward_chunkr   r   )rD   ru   r   r   r   rt   r   layer_outputs           rG   rU   zConvBertLayer.forwards  s     *4>>
 
 ??4@4!12$=dV DD D   3t22 &  '<  	  1##T%A%A4CSCSUe
 rH   c                 L    | j                  |      }| j                  ||      }|S rw   )r   r   )rD   r   intermediate_outputr   s       rG   r   z ConvBertLayer.feed_forward_chunk  s,    "//0@A{{#68HIrH   NNN)rV   rW   rX   r-   r=   rz   r[   r   r   rU   r   r\   r]   s   @rG   r   r   d  s    -" 48596:|| ))D0  %||d2	
 !&t 3 +, 
@rH   r   c                   d     e Zd ZU eed<   dZdZeedZ	 e
j                          fd       Z xZS )ConvBertPreTrainedModelrE   convbertT)ru   
attentionsc                 b   t         |   |       t        |t              r t	        j
                  |j                         yt        |t              rVt	        j                  |j                  d| j                  j                         t	        j
                  |j                         yt        |t              ryt	        j                  |j                  t        j                   |j                  j"                  d         j%                  d             t	        j
                  |j&                         yy)zInitialize the weightsrf   rg   r'   r&   N)r,   _init_weightsr   r_   initzeros_re   r   rp   rn   rE   rq   r    copy_r%   r=   r>   r   r?   r)   )rD   modulerF   s     rG   r   z%ConvBertPreTrainedModel._init_weights  s     	f%fo.KK$ 23LLSdkk6S6STKK$ 23JJv**ELL9L9L9R9RSU9V,W,^,^_f,ghKK--. 4rH   )rV   rW   rX   r   __annotations__base_model_prefixsupports_gradient_checkpointingr   r|   _can_record_outputsr=   no_gradr   r\   r]   s   @rG   r   r     s?    "&*#&+
 U]]_
/ 
/rH   r   c                        e Zd Z fdZ	 	 	 d	dej
                  dej                  dz  dej
                  dz  dej
                  dz  def
dZ xZ	S )
ConvBertEncoderc                     t         |           || _        t        j                  t        |j                        D cg c]  }t        |       c}      | _        d| _	        y c c}w )NF)
r,   r-   rE   r   
ModuleListrangenum_hidden_layersr   layergradient_checkpointing)rD   rE   r   rF   s      rG   r-   zConvBertEncoder.__init__  sN    ]]5IaIaCb#caM&$9#cd
&+# $ds   A#Nru   r   r   r   rK   c                 V    | j                   D ]  } |||f||d|} t        |      S )N)r   r   )last_hidden_state)r  r   )rD   ru   r   r   r   rt   layer_modules          rG   rU   zConvBertEncoder.forward  sP     !JJ 	L( '<'=	
 M	 2+
 	
rH   r   )
rV   rW   rX   r-   r=   rz   r[   r   rU   r\   r]   s   @rG   r  r    si    , 48596:
||
 ))D0
  %||d2	

 !&t 3
 
,
rH   r  c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )ConvBertPredictionHeadTransformc                 h   t         |           t        j                  |j                  |j                        | _        t        |j                  t              rt        |j                     | _
        n|j                  | _
        t        j                  |j                  |j                        | _        y r   )r,   r-   r   r   r   r   r   r   r   r
   transform_act_fnr7   r8   rC   s     rG   r-   z(ConvBertPredictionHeadTransform.__init__  s{    YYv1163E3EF
f''-$*6+<+<$=D!$*$5$5D!f&8&8f>S>STrH   ru   rK   c                 l    | j                  |      }| j                  |      }| j                  |      }|S rw   )r   r  r7   r   s     rG   rU   z'ConvBertPredictionHeadTransform.forward  s4    

=1--m<}5rH   r   r]   s   @rG   r  r    s$    UU\\ ell rH   r  c                        e Zd ZdZdef fdZ	 d	dej                  dej                  dz  dej                  fdZ	 xZ
S )
ConvBertSequenceSummarya  
    Compute a single vector summary of a sequence hidden states.

    Args:
        config ([`ConvBertConfig`]):
            The config used by the model. Relevant arguments in the config class of the model are (refer to the actual
            config class of your model for the default values it uses):

            - **summary_type** (`str`) -- The method to use to make this summary. Accepted values are:

                - `"last"` -- Take the last token hidden state (like XLNet)
                - `"first"` -- Take the first token hidden state (like Bert)
                - `"mean"` -- Take the mean of all tokens hidden states
                - `"cls_index"` -- Supply a Tensor of classification token position (GPT/GPT-2)
                - `"attn"` -- Not implemented now, use multi-head attention

            - **summary_use_proj** (`bool`) -- Add a projection after the vector extraction.
            - **summary_proj_to_labels** (`bool`) -- If `True`, the projection outputs to `config.num_labels` classes
              (otherwise to `config.hidden_size`).
            - **summary_activation** (`Optional[str]`) -- Set to `"tanh"` to add a tanh activation to the output,
              another string or `None` will add no activation.
            - **summary_first_dropout** (`float`) -- Optional dropout probability before the projection and activation.
            - **summary_last_dropout** (`float`)-- Optional dropout probability after the projection and activation.
    rE   c                 f   t         |           t        |dd      | _        | j                  dk(  rt        t        j                         | _        t        |d      rq|j                  ret        |d      r(|j                  r|j                  dkD  r|j                  }n|j                  }t        j                  |j                  |      | _        t        |dd       }|rt        |      nt        j                         | _        t        j                         | _        t        |d      r3|j"                  dkD  r$t        j$                  |j"                        | _        t        j                         | _        t        |d	      r5|j(                  dkD  r%t        j$                  |j(                        | _        y y y )
Nsummary_typelastattnsummary_use_projsummary_proj_to_labelsr   summary_activationsummary_first_dropoutsummary_last_dropout)r,   r-   getattrr  NotImplementedErrorr   IdentitysummaryrO   r  r  
num_labelsr   r   r   
activationfirst_dropoutr   r9   last_dropoutr!  )rD   rE   num_classesactivation_stringrF   s       rG   r-   z ConvBertSequenceSummary.__init__  sU   #FNFC& &%{{}6-.63J3Jv78V=Z=Z_e_p_pst_t$//$0099V%7%7EDL#F,@$GIZN3D$E`b`k`k`m[[]6238T8TWX8X!#F,H,H!IDKKM612v7R7RUV7V "

6+F+F GD 8W2rH   Nru   	cls_indexrK   c                    | j                   dk(  r|dddf   }n| j                   dk(  r|dddf   }n| j                   dk(  r|j                  d      }n| j                   d	k(  r|At        j                  |d
ddddf   |j                  d   dz
  t        j
                        }nX|j                  d      j                  d      }|j                  d|j                         dz
  z  |j                  d      fz         }|j                  d|      j                  d      }n| j                   dk(  rt        | j                        }| j                  |      }| j                  |      }| j!                  |      }|S )ak  
        Compute a single vector summary of a sequence hidden states.

        Args:
            hidden_states (`torch.FloatTensor` of shape `[batch_size, seq_len, hidden_size]`):
                The hidden states of the last layer.
            cls_index (`torch.LongTensor` of shape `[batch_size]` or `[batch_size, ...]` where ... are optional leading dimensions of `hidden_states`, *optional*):
                Used if `summary_type == "cls_index"` and takes the last token of the sequence as classification token.

        Returns:
            `torch.FloatTensor`: The summary of the sequence hidden states.
        r  Nr'   firstr   rh   r   r   r,  .r   r*   )r'   r  )r  rh   r=   	full_liker   rB   r   r?   r   rA   gathersqueezer#  r(  r%  r'  r)  )rD   ru   r,  r   s       rG   rU   zConvBertSequenceSummary.forward  sn    &"1b5)F')"1a4(F&("''A'.F+- !OO!#rr1*-!''+a/**	 &//3==bA	%,,Uimmo6I-JmN`N`acNdMf-fg	"))"i8@@DF&(%%##F+f%(""6*rH   rw   )rV   rW   rX   rY   r   r-   r=   r[   rZ   rU   r\   r]   s   @rG   r  r    sQ    2H~ H< VZ)"..);@;K;Kd;R)			)rH   r  c                        e Zd Z fdZd Zd Zeee	 	 	 	 	 dde	j                  dz  de	j                  dz  de	j                  dz  de	j                  dz  d	e	j                  dz  d
ee   defd                     Z xZS )ConvBertModelc                 "   t         |   |       t        |      | _        |j                  |j
                  k7  r/t        j                  |j                  |j
                        | _        t        |      | _
        || _        | j                          y rw   )r,   r-   r    rT   r0   r   r   r   embeddings_projectr  encoderrE   	post_initrC   s     rG   r-   zConvBertModel.__init__E  sl     ,V4  F$6$66&(ii0E0EvGYGY&ZD#&v.rH   c                 .    | j                   j                  S rw   rT   r2   rD   s    rG   get_input_embeddingsz"ConvBertModel.get_input_embeddingsQ  s    ...rH   c                 &    || j                   _        y rw   r9  )rD   r   s     rG   set_input_embeddingsz"ConvBertModel.set_input_embeddingsT  s    */'rH   NrI   r   r)   r%   rJ   rt   rK   c                    ||t        d      |#| j                  ||       |j                         }n!||j                         d d }nt        d      |\  }}	||j                  n|j                  }
|t	        j
                  ||
      }|pt        | j                  d      r4| j                  j                  d d d |	f   }|j                  ||	      }|}n&t	        j                  |t        j                  |
      }| j                  ||||      }t        | d      r| j                  |      }t        | j                  ||	      } | j                  |fd
|i|}|S )NzDYou cannot specify both input_ids and inputs_embeds at the same timer'   z5You have to specify either input_ids or inputs_embeds)rN   r)   rM   )rI   r%   r)   rJ   r5  )rE   rJ   r   r   )r   %warn_if_padding_and_no_attention_maskrA   rN   r=   onesrO   rT   r)   r?   r@   rB   r5  r   rE   r6  )rD   rI   r   r)   r%   rJ   rt   rP   r   rQ   rN   rR   rS   ru   encoder_outputss                  rG   rU   zConvBertModel.forwardW  s     ]%>cdd"66y.Q#..*K&',,.s3KTUU!,
J%.%:!!@T@T!"ZZFCN!t(89*.//*H*HKZK*X'3J3Q3QR\^h3i0!A!&[

SY!Zl>iv ( 
 4-. 33MBM2;;')
 ?Kdll?
)?
 ?
 rH   )NNNNN)rV   rW   rX   r-   r;  r=  r   r   r   r=   rZ   r[   r   r   r   rU   r\   r]   s   @rG   r3  r3  C  s    
/0   .2372604263##d*3 ))D03 ((4/	3
 &&-3 ((4/3 +,3 
,3    3rH   r3  c                   Z     e Zd ZdZ fdZdej                  dej                  fdZ xZS )ConvBertGeneratorPredictionszAPrediction module for the generator, made up of two dense layers.c                     t         |           t        d      | _        t	        j
                  |j                  |j                        | _        t	        j                  |j                  |j                        | _
        y )Ngelur#   )r,   r-   r   r'  r   r7   r0   r8   r   r   r   rC   s     rG   r-   z%ConvBertGeneratorPredictions.__init__  sV    (0f&;&;AVAVWYYv1163H3HI
rH   generator_hidden_statesrK   c                 l    | j                  |      }| j                  |      }| j                  |      }|S rw   )r   r'  r7   )rD   rF  ru   s      rG   rU   z$ConvBertGeneratorPredictions.forward  s3    

#:;6}5rH   )	rV   rW   rX   rY   r-   r=   r[   rU   r\   r]   s   @rG   rC  rC    s+    KJu/@/@ UEVEV rH   rC  c                   $    e Zd ZddiZ fdZd Zd Zee	 	 	 	 	 	 dde	j                  dz  de	j                  dz  d	e	j                  dz  d
e	j                  dz  de	j                  dz  de	j                  dz  dee   deez  fd              Z xZS )ConvBertForMaskedLMzgenerator_lm_head.weightz*convbert.embeddings.word_embeddings.weightc                     t         |   |       t        |      | _        t	        |      | _        t        j                  |j                  |j                        | _
        | j                          y rw   )r,   r-   r3  r   rC  generator_predictionsr   r   r0   r/   generator_lm_headr7  rC   s     rG   r-   zConvBertForMaskedLM.__init__  sR     %f-%A&%I"!#6+@+@&BSBS!TrH   c                     | j                   S rw   rL  r:  s    rG   get_output_embeddingsz)ConvBertForMaskedLM.get_output_embeddings  s    %%%rH   c                     || _         y rw   rN  )rD   r2   s     rG   set_output_embeddingsz)ConvBertForMaskedLM.set_output_embeddings  s
    !0rH   NrI   r   r)   r%   rJ   labelsrt   rK   c                 n    | j                   |f||||d|}|d   }	| j                  |	      }
| j                  |
      }
d}|Pt        j                         } ||
j                  d| j                  j                        |j                  d            }t        ||
|j                  |j                        S )a  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,
            config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are ignored (masked), the
            loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`
        r   r)   r%   rJ   r   Nr'   losslogitsru   r   )r   rK  rL  r   r   r   rE   r/   r   ru   r   )rD   rI   r   r)   r%   rJ   rR  rt   rF  generator_sequence_outputprediction_scoresrV  loss_fcts                rG   rU   zConvBertForMaskedLM.forward  s    $ GTdmmG
))%'G
 G
 %<A$>! 667PQ 223DE**,H-222t{{7M7MNPVP[P[\^P_`D$1??.99	
 	
rH   NNNNNN)rV   rW   rX   _tied_weights_keysr-   rO  rQ  r   r   r=   rZ   r[   r   r   r   r   rU   r\   r]   s   @rG   rI  rI    s    46bc&1  .237260426*.(
##d*(
 ))D0(
 ((4/	(

 &&-(
 ((4/(
   4'(
 +,(
 
	(
  (
rH   rI  c                   Z     e Zd ZdZ fdZdej                  dej                  fdZ xZS )ConvBertClassificationHeadz-Head for sentence-level classification tasks.c                 h   t         |           t        j                  |j                  |j                        | _        |j                  |j                  n|j                  }t        j                  |      | _	        t        j                  |j                  |j                        | _        || _        y rw   )r,   r-   r   r   r   r   classifier_dropoutr:   r9   r;   r&  out_projrE   rD   rE   r`  rF   s      rG   r-   z#ConvBertClassificationHead.__init__  s    YYv1163E3EF
)/)B)B)NF%%TZTnTn 	 zz"45		&"4"4f6G6GHrH   ru   rK   c                     |d d dd d f   }| j                  |      }| j                  |      }t        | j                  j                     |      }| j                  |      }| j                  |      }|S )Nr   )r;   r   r
   rE   r   ra  )rD   ru   rt   rx   s       rG   rU   z"ConvBertClassificationHead.forward  se    !Q'"LLOJJqM4;;))*1-LLOMM!rH   ry   r]   s   @rG   r^  r^    s&    7	U\\  rH   r^  z
    ConvBERT Model transformer with a sequence classification/regression head on top (a linear layer on top of the
    pooled output) e.g. for GLUE tasks.
    )custom_introc                       e Zd Z fdZee	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  d	e	e
   d
eez  fd              Z xZS )!ConvBertForSequenceClassificationc                     t         |   |       |j                  | _        || _        t	        |      | _        t        |      | _        | j                          y rw   )	r,   r-   r&  rE   r3  r   r^  
classifierr7  rC   s     rG   r-   z*ConvBertForSequenceClassification.__init__  sH      ++%f-4V< 	rH   NrI   r   r)   r%   rJ   rR  rt   rK   c                     | j                   |f||||d|}|d   }	| j                  |	      }
d}|| j                  j                  | j                  dk(  rd| j                  _        nl| j                  dkD  rL|j
                  t        j                  k(  s|j
                  t        j                  k(  rd| j                  _        nd| j                  _        | j                  j                  dk(  rIt               }| j                  dk(  r& ||
j                         |j                               }n ||
|      }n| j                  j                  dk(  r=t               } ||
j                  d| j                        |j                  d            }n,| j                  j                  dk(  rt               } ||
|      }t        ||
|j                  |j                   	      S )
a  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
        rT  r   Nr   
regressionsingle_label_classificationmulti_label_classificationr'   rU  )r   rh  rE   problem_typer&  r+   r=   rB   r   r   r1  r   r   r   r   ru   r   rD   rI   r   r)   r%   rJ   rR  rt   outputssequence_outputrW  rV  rZ  s                rG   rU   z)ConvBertForSequenceClassification.forward  s   $ 7Ddmm7
))%'7
 7
 "!*1{{''/??a'/;DKK,__q(fllejj.HFLL\a\e\eLe/LDKK,/KDKK,{{''<7"9??a'#FNN$4fnn6FGD#FF3D))-JJ+-B @&++b/R))-II,./'!//))	
 	
rH   r[  )rV   rW   rX   r-   r   r   r=   rZ   r[   r   r   r   r   rU   r\   r]   s   @rG   rf  rf    s      .237260426*.8
##d*8
 ))D08
 ((4/	8

 &&-8
 ((4/8
   4'8
 +,8
 
)	)8
  8
rH   rf  c                       e Zd Z fdZee	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  d	e	e
   d
eez  fd              Z xZS )ConvBertForMultipleChoicec                     t         |   |       t        |      | _        t	        |      | _        t        j                  |j                  d      | _	        | j                          y )Nr   )r,   r-   r3  r   r  sequence_summaryr   r   r   rh  r7  rC   s     rG   r-   z"ConvBertForMultipleChoice.__init__K  sM     %f- 7 ?))F$6$6: 	rH   NrI   r   r)   r%   rJ   rR  rt   rK   c                    ||j                   d   n|j                   d   }|!|j                  d|j                  d            nd}|!|j                  d|j                  d            nd}|!|j                  d|j                  d            nd}|!|j                  d|j                  d            nd}|1|j                  d|j                  d      |j                  d            nd} | j                  |f||||d|}	|	d   }
| j	                  |
      }| j                  |      }|j                  d|      }d}|t               } |||      }t        |||	j                  |	j                        S )a\  
        input_ids (`torch.LongTensor` of shape `(batch_size, num_choices, sequence_length)`):
            Indices of input sequence tokens in the vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            [What are input IDs?](../glossary#input-ids)
        token_type_ids (`torch.LongTensor` of shape `(batch_size, num_choices, sequence_length)`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:


            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.

            [What are token type IDs?](../glossary#token-type-ids)
        position_ids (`torch.LongTensor` of shape `(batch_size, num_choices, sequence_length)`, *optional*):
            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.

            [What are position IDs?](../glossary#position-ids)
        inputs_embeds (`torch.FloatTensor` of shape `(batch_size, num_choices, sequence_length, hidden_size)`, *optional*):
            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
            is useful if you want more control over how to convert *input_ids* indices into associated vectors than the
            model's internal embedding lookup matrix.
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the multiple choice classification loss. Indices should be in `[0, ...,
            num_choices-1]` where `num_choices` is the size of the second dimension of the input tensors. (See
            `input_ids` above)
        Nr   r'   r   rT  r   rU  )
r   r   rA   r   rt  rh  r   r   ru   r   )rD   rI   r   r)   r%   rJ   rR  rt   num_choicesro  rp  pooled_outputrW  reshaped_logitsrV  rZ  s                   rG   rU   z!ConvBertForMultipleChoice.forwardU  s   V -6,Aiooa(}GZGZ[\G]>G>SINN2y~~b'9:Y]	M[Mg,,R1D1DR1HImqM[Mg,,R1D1DR1HImqGSG_|((\->->r-BCei ( r=#5#5b#9=;M;Mb;QR 	 7Ddmm7
))%'7
 7
 "!*--o>/ ++b+6')HOV4D("!//))	
 	
rH   r[  )rV   rW   rX   r-   r   r   r=   rZ   r[   r   r   r   r   rU   r\   r]   s   @rG   rr  rr  I  s      .237260426*.N
##d*N
 ))D0N
 ((4/	N

 &&-N
 ((4/N
   4'N
 +,N
 
*	*N
  N
rH   rr  c                       e Zd Z fdZee	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  d	e	e
   d
eez  fd              Z xZS )ConvBertForTokenClassificationc                 `   t         |   |       |j                  | _        t        |      | _        |j
                  |j
                  n|j                  }t        j                  |      | _	        t        j                  |j                  |j                        | _        | j                          y rw   )r,   r-   r&  r3  r   r`  r:   r   r9   r;   r   r   rh  r7  rb  s      rG   r-   z'ConvBertForTokenClassification.__init__  s      ++%f-)/)B)B)NF%%TZTnTn 	 zz"45))F$6$68I8IJ 	rH   NrI   r   r)   r%   rJ   rR  rt   rK   c                 F    | j                   |f||||d|}|d   }	| j                  |	      }	| j                  |	      }
d}|<t               } ||
j	                  d| j
                        |j	                  d            }t        ||
|j                  |j                        S )z
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the token classification loss. Indices should be in `[0, ..., config.num_labels - 1]`.
        rT  r   Nr'   rU  )	r   r;   rh  r   r   r&  r   ru   r   rn  s                rG   rU   z&ConvBertForTokenClassification.forward  s      7Ddmm7
))%'7
 7
 "!*,,71')HFKKDOO<fkk"oND$!//))	
 	
rH   r[  )rV   rW   rX   r-   r   r   r=   rZ   r[   r   r   r   r   rU   r\   r]   s   @rG   rz  rz    s      .237260426*.&
##d*&
 ))D0&
 ((4/	&

 &&-&
 ((4/&
   4'&
 +,&
 
&	&&
  &
rH   rz  c                   *    e Zd Z fdZee	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  dej                  dz  d	ej                  dz  d
e	e
   defd              Z xZS )ConvBertForQuestionAnsweringc                     t         |   |       |j                  | _        t        |      | _        t        j                  |j                  |j                        | _        | j                          y rw   )
r,   r-   r&  r3  r   r   r   r   
qa_outputsr7  rC   s     rG   r-   z%ConvBertForQuestionAnswering.__init__  sS      ++%f-))F$6$68I8IJ 	rH   NrI   r   r)   r%   rJ   start_positionsend_positionsrt   rK   c                     | j                   |f||||d|}	|	d   }
| j                  |
      }|j                  dd      \  }}|j                  d      j	                         }|j                  d      j	                         }d }||t        |j                               dkD  r|j                  d      }t        |j                               dkD  r|j                  d      }|j                  d      }|j                  d|      }|j                  d|      }t        |      } |||      } |||      }||z   dz  }t        ||||	j                  |	j                        S )	NrT  r   r   r'   r   )ignore_indexra   )rV  start_logits
end_logitsru   r   )r   r  splitr1  r   lenrA   clampr   r   ru   r   )rD   rI   r   r)   r%   rJ   r  r  rt   ro  rp  rW  r  r  
total_lossignored_indexrZ  
start_lossend_losss                      rG   rU   z$ConvBertForQuestionAnswering.forward  s~    7Ddmm7
))%'7
 7
 "!*1#)<<r<#: j#++B/::<''+668

&=+D?'')*Q."1"9"9""==%%'(1, - 5 5b 9(--a0M-33A}EO)//=AM']CH!,@J
M:H$x/14J+%!!//))
 	
rH   )NNNNNNN)rV   rW   rX   r-   r   r   r=   rZ   r[   r   r   r   rU   r\   r]   s   @rG   r~  r~    s      .23726042637152
##d*2
 ))D02
 ((4/	2

 &&-2
 ((4/2
 ))D02
 ''$.2
 +,2
 
&2
  2
rH   r~  )rI  rr  r~  rf  rz  r   r3  r   )FrY   r   collections.abcr   r=   r   torch.nnr   r   r    r	   r   activationsr
   r   masking_utilsr   modeling_layersr   modeling_outputsr   r   r   r   r   r   modeling_utilsr   processing_utilsr   pytorch_utilsr   utilsr   r   r   r   utils.genericr   utils.output_capturingr   configuration_convbertr   
get_loggerrV   loggerModuler    r_   r|   r   r   r   r   r   r   r   r  r  r  r3  rC  rI  r^  rf  rr  rz  r~  __all__ rH   rG   <module>r     s.     $   A A & 1 6 9  . & 6  8 5 2 
		H	%7 7tbii 4w.BII w.t  		  . ,299 (RYY &2. 2j /o / /.
bii 
:bii $`bii `F I+ I IX299 $ =
1 =
 =
@ 0 E
(? E
E
P [
 7 [
 [
| 7
%< 7
 7
t ?
#: ?
 ?
D	rH   