CoolFace
Apppublic

Aluode/PerceptionLabPortable

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
modeling_bart.cpython-310.pyc616 linesDownload Raw Back to __pycache__
1o

2.�Yi1_�@s�dZddlZddlZddlmZmZmZddlZddlmZddl	m3Z4mZmZddl
mZddlmZmZmZdd	lmZdd5lmZmZmZddlmZddlmZdd
lmZmZm Z m!Z!m"Z"m#Z#m$Z$ddl%m&Z&m'Z'ddl(m)Z)ddl*m+Z+m,Z,m-Z-m.Z.ddl/m0Z0ddl1m2Z2e,�r�ddl3m4Z4m5Z5e.�6e7�Z8dej9de:de:fdd�Z;Gdd�dej<�Z=Gdd�dej<�Z>			dLdej?d ej9d!ej9d"ej9d#eej9d$ee@d%e@d&eej9fd'd(�ZAGd)d*�d*ej?�ZBGd+d,�d,e�ZCGd-d.�d.e�ZDGd/d0�d0ej?�ZEe+Gd1d2�d2e'��ZFGd3d4�d4eF�ZGGd5d6�d6eF�ZHGd7d8�d8eF�ZIGd9d:�d:eF�ZJe+Gd;d<�d<eF��ZKe+d=d>�Gd?d@�d@eFe��ZLe+dAd>�GdBdC�dCeF��ZMe+GdDdE�dEeF��ZNGdFdG�dGeF�ZOe+dHd>�GdIdJ�dJeFe��ZPgdK�ZQdS)MzPyTorch BART model.�N)�Callable�Optional�Union)�nn)�BCEWithLogitsLoss�CrossEntropyLoss�MSELoss�)�ACT2FN)�Cache�DynamicCache�EncoderDecoderCache)�GenerationMixin)�AttentionMaskConverter�_prepare_4d_attention_mask�#_prepare_4d_attention_mask_for_sdpa)�FlashAttentionKwargs)�GradientCheckpointingLayer)�BaseModelOutput�)BaseModelOutputWithPastAndCrossAttentions�!CausalLMOutputWithCrossAttentions�Seq2SeqLMOutput�Seq2SeqModelOutput�#Seq2SeqQuestionAnsweringModelOutput�Seq2SeqSequenceClassifierOutput)�ALL_ATTENTION_FUNCTIONS�PreTrainedModel)�Unpack)�auto_docstring�is_torch_flex_attn_available�is_torchdynamo_compiling�logging)�deprecate_kwarg�)�6BartConfig)�	BlockMask�make_flex_block_causal_mask�	input_ids�pad_token_id�decoder_start_token_idcCsh|�|j�}|dd�dd�f��|dd�dd�f<||dd�df<|dur*td��|�|dk|�|S)z17    Shift input ids one token to the right.8    N�����r#rz1self.model.config.pad_token_id has to be defined.i����)Z	new_zeros�shape�clone�9ValueErrorZmasked_fill_)r'r(r)Zshifted_input_ids�r.��E:\DocsHouse\542 percep lab latest\PerceptionLab\PerceptionLab_Portable\python_embed\Lib\site-packages\transformers/models/bart/modeling_bart.py�shift_tokens_right?s(r0csPeZdZdZdedef�fdd�Z	d
dejd	ed10eejf�fdd�
Z	�Z11S)�BartLearnedPositionalEmbeddingzN12    This module learns positional embeddings up to a fixed maximum size.13    �num_embeddings�
embedding_dimcsd|_t��||j|�dS�N�)�offset�super�__init__)�selfr2r3��	__class__r.r/r8Tsz'BartLearnedPositionalEmbedding.__init__rNr'�past_key_values_length�position_idscs\|dur |jdd�\}}tj|||tj|jjd��|d�}n|�d�}t��	||j14�S)z3`input_ids' shape is expected to be [bsz x seqlen].Nr5)�dtype�devicer*r)r+�torch�arange�long�weightr?�expandZ	unsqueezer7�forwardr6)r9r'r<r=�bszZseq_lenr:r.r/rEZs��15z&BartLearnedPositionalEmbedding.forward)rN)�__name__�16__module__�__qualname__�__doc__�intr8r@�TensorrrE�
__classcell__r.r.r:r/r1Os����r1c17sLeZdZdZddedededeef�fdd�
Zd	ej	f�fd18d�Z19�ZS)
�BartScaledWordEmbeddingz\20    This module overrides nn.Embeddings' forward by multiplying with embeddings scale.21    ��?r2r3�padding_idx�embed_scalecst��|||�||_dS�N)r7r8rQ)r9r2r3rPrQr:r.r/r8os22z BartScaledWordEmbedding.__init__r'cst��|�|jSrR)r7rErQ)r9r'r:r.r/rEsszBartScaledWordEmbedding.forward)rO)rGrHrIrJrKr�floatr8r@rLrErMr.r.r:r/rNjs$rN��module�query�key�value�attention_mask�scaling�dropout�	head_maskcKs�|dur|�d�d}t�||�dd��|}	|dur|	|}	tjj|	dd�}	|dur5|	|�dddd�}	tjj|	||j	d�}	t�|	|�}23|24�dd��25�}26|27|	fS)Nr*��r5r	��dimr#��p�training)�sizer@�matmul�	transposer�28functionalZsoftmax�viewr[rb�29contiguous)rUrVrWrXrYrZr[r\�kwargs�attn_weights�attn_outputr.r.r/�eager_attention_forwardwsrlcs�eZdZdZ						ddededed	ed30ededeed
eef�fdd�
Z	e31dddd�						ddejdeejdee
deejdeejdedeejdeedeejeejeeejffdd��Z�ZS) �
BartAttentionz=Multi-headed attention from 'Attention Is All You Need' paperrTFTN�	embed_dim�	num_headsr[�32is_decoder�bias�	is_causal�config�	layer_idxc		s�t���||_||_||_|||_||_|j||jkr*td|j�d|�d���|jd|_||_	||_33||_|durK|j	rKt�
d|jj�d��tj|||d�|_tj|||d�|_tj|||d�|_tj|||d�|_dS)Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: z).r]zInstantiating a decoder z� without passing `layer_idx` is not recommended and will lead to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` when creating this class.�rq)r7r8rnror[�head_dimrsr-rZrprrrt�logger�warning_oncer;rGr�Linear�k_proj�v_proj�q_proj�out_proj)	r9rnror[rprqrrrsrtr:r.r/r8�s0343536���zBartAttention.__init__�past_key_value�past_key_values�4.58��new_name�version�
hidden_states�key_value_statesrY�layer_head_mask�output_attentions�cache_positionri�returncKs�|du}	|jdd�\}37}|	r|jdn|}|38|d|jf}
|39|d|jf}|�|�j|
��dd�}d}|durNt|t�rL|j�|j	�}|	rH|j40}n|j}n|}|	rR|n|}|	rk|durk|rk|j|j	j
}|j|j	j}n@|�|�}|�|�}|j|��dd�}|j|��dd�}|dur�|	s�|nd}|�|||j	d|i�\}}|	r�t|t�r�d|j|j	<t}|jjdkr�t|jj}||||||f|js�d	n|j|j||d41�|��\}}|�|42|d���}|�|�}||fS)z#Input shape: Batch x Time x ChannelNr*r#r5Fr�T�eagerrT)r[rZr�r\)r+rvr|rgre�43isinstancer
�44is_updated�getrtZcross_attention_cache�self_attention_cache�layers�keys�valuesrzr{�updaterlrs�_attn_implementationrrbr[rZ�reshaperhr})r9r�r�rrYr�r�r�riZis_cross_attentionrF�tgt_lenZsrc_lenZ
q_input_shapeZkv_input_shapeZquery_statesr�Zcurr_past_key_valueZcurrent_statesZ45key_statesZvalue_statesZattention_interfacerkrjr.r.r/rE�sb464748���49 50�
51zBartAttention.forward)rTFTFNN)NNNNFN)rGrHrIrJrKrS�boolrr$r8r"r@rLrrr�tuplerErMr.r.r:r/rm�sf��������	�'����������rmcsheZdZddedeef�fdd�
Z	ddejdejd	ejd52ee	de53ejeejff54dd
�Z�ZS)�BartEncoderLayerNrsrtcs�t���|j|_t|j|j|j||d�|_t�	|j�|_55|j|_t|j
|_|j|_t�|j|j�|_t�|j|j�|_t�	|j�|_dS)N)rnror[rsrt)r7r8�d_modelrnrmZencoder_attention_heads�attention_dropout�	self_attnr�	LayerNorm�self_attn_layer_normr[r56�activation_function�
activation_fn�activation_dropoutryZencoder_ffn_dim�fc1�fc2�final_layer_norm�r9rsrtr:r.r/r8s 57�zBartEncoderLayer.__init__Fr�rYr�r�r�c	Cs|}|j||||d�\}}tjj||j|jd�}||}|�|�}|}|�|�|��}tjj||j|jd�}|�	|�}tjj||j|jd�}||}|�58|�}|jtj
krut�|���sct�|���rut�|j�jd}tj|||d�}|f}|r||f7}|S)a�59        Args:60            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`61            attention_mask (`torch.FloatTensor`): attention mask of size62                `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.63            layer_head_mask (`torch.FloatTensor`): mask for attention heads in a given layer of size64                `(encoder_attention_heads,)`.65            output_attentions (`bool`, *optional*):66                Whether or not to return the attentions tensors of all attention layers. See `attentions` under67                returned tensors for more detail.68        )r�rYr�r�r`i�)�min�max)r�rrfr[rbr�r�r�r�r�r�r>r@Zfloat16�isinf�any�isnan�finfor��clamp)	r9r�rYr�r��residualrjZclamp_value�outputsr.r.r/rE)s869�707172��73zBartEncoderLayer.forwardrR)F)
rGrHrIr$rrKr8r@�FloatTensorr�r�rErMr.r.r:r/r�s������r�cs�eZdZddedeef�fdd�
Zedddd	�							74		ddej	d
eej	deej	deej	deej	deej	dee75deedeedeej	deej
eeej
ej
fffdd��Z�ZS)�BartDecoderLayerNrsrtc	s�t���|j|_t|j|j|jdd||d�|_|j|_t	|j76|_|j|_t
�|j�|_t|j|j|jd||d�|_t
�|j�|_t
�|j|j�|_t
�|j|j�|_t
�|j�|_dS)NT)rnror[rprrrsrt)r[rprsrt)r7r8r�rnrmZdecoder_attention_headsr�r�r[r77r�r�r�rr�r��encoder_attn�encoder_attn_layer_normryZdecoder_ffn_dimr�r�r�r�r:r.r/r8]s678�	�zBartDecoderLayer.__init__r~rr�r�FTr�rY�encoder_hidden_states�encoder_attention_maskr��cross_attn_layer_head_maskr��	use_cacher�r�c	Cs|}|j||||||79d�\}}tjj||j|jd�}||}|�|�}d}
|durM|}|j|||||||80d�\}}
tjj||j|jd�}||}|�|�}|}|�|�	|��}tjj||j81|jd�}|�|�}tjj||j|jd�}||}|�|�}|f}|r�|||
f7}|S)a&82        Args:83            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`84            attention_mask (`torch.FloatTensor`): attention mask of size85                `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.86            encoder_hidden_states (`torch.FloatTensor`):87                cross attention input to the layer of shape `(batch, seq_len, embed_dim)`88            encoder_attention_mask (`torch.FloatTensor`): encoder attention mask of size89                `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.90            layer_head_mask (`torch.FloatTensor`): mask for attention heads in a given layer of size91                `(encoder_attention_heads,)`.92            cross_attn_layer_head_mask (`torch.FloatTensor`): mask for cross-attention heads in a given layer of93                size `(decoder_attention_heads,)`.94            past_key_values (`Cache`): cached past key and value projection states95            output_attentions (`bool`, *optional*):96                Whether or not to return the attentions tensors of all attention layers. See `attentions` under97                returned tensors for more detail.98            cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):99                Indices depicting the position of the input sequence tokens in the sequence. It is used to update the100                cache in the correct position and to infer the complete sequence length.101        )r�rrYr�r�r�r`N)r�r�rYr�rr�r�)
r�rrfr[rbr�r�r�r�r�r�r�r�)r9r�rYr�r�r�r�rr�r�r�r�Zself_attn_weightsZcross_attn_weightsr�r.r.r/rE|sL#102�103104�	105106107zBartDecoderLayer.forwardrR)	NNNNNNFTN)rGrHrIr$rrKr8r"r@rLrr�r�r�rErMr.r.r:r/r�\sF��������	�108���r�csHeZdZdZdedededef�fdd�Zdejd	ejfd109d�Z	�Z110S)�BartClassificationHeadz-Head for sentence-level classification tasks.�	input_dim�	inner_dim�num_classes�pooler_dropoutcs8t���t�||�|_tj|d�|_t�||�|_dS)N)ra)r7r8rry�denseZDropoutr[r})r9r�r�r�r�r:r.r/r8�s111zBartClassificationHead.__init__r�r�cCs6|�|�}|�|�}t�|�}|�|�}|�|�}|SrR)r[r�r@�tanhr})r9r�r.r.r/rE�s112113114115116zBartClassificationHead.forward)rGrHrIrJrKrSr8r@rLrErMr.r.r:r/r��s����r�c
@s�eZdZUeed<dZdZddgZddgZdZ	dZ117dZdZdZ
d	d118�Zedd��Zd
eejdfdejfdd�Zd
eeejdfdejdejdefdd�Zed
ejdededejdejdefdd��Zdeejdfdeejdfdejdejfd d!�ZdS)"�BartPreTrainedModelrs�modelTzencoder.versionzdecoder.versionr�r�rcCs�|jj}t|tj�r"|jjjd|d�|jdur |jj�	�dSdSt|tj119�rC|jjjd|d�|jdurA|jj|j�	�dSdSt|tj�rX|jj�
d�|jj�	�dSdS)NrT)�mean�stdrO)rsZinit_stdr�rryrC�dataZnormal_rqZzero_�	EmbeddingrPr�Zfill_)r9rUr�r.r.r/�
_init_weights�s120�121��z!BartPreTrainedModel._init_weightscCs>|jj}tjgd�dddd|gg|jd�}|�|�|d�}|S)N)r��122�r5r��r5�r?)rYr')rsr(r@Ztensorr?�ne)r9Z	pad_tokenr'�dummy_inputsr.r.r/r�s"�z BartPreTrainedModel.dummy_inputsrYN�
inputs_embedscCs�|dur>|jjdkrd|vr|}|Sd}|S|jjdkr$t||j�}|S|jjdkr8t|tj�r6t|dd�}|St||j�}|S)N�flash_attention_2r�sdpa�flex_attentionF)rr�	rsr�rr>r�r@rLr&r)r9rYr�r.r.r/�_update_full_masks
�
���z%BartPreTrainedModel._update_full_maskr%�input_tensorr�cCsb|jjdkr*t|tj�rt|�}|S|dur(ttj|jd|jdf|jd��}|S|jjdkr>|dur<|dk�	�r<|SdS|durF|�123�nd}|durO|jnd}|jjdkre|setj
||||jd	�redS|j}|jd}|rt|��}	nt|tj�r|jd124n||d}	|j|||	|||jdd�}125|jjdkr�|dur�|jjdvr�t�|�j}t�|126|�}127|128S)
Nr�rr#)rcr?r�rTFr�)r�r<Zis_trainingr*)�sequence_length�
target_lengthr>r��129batch_size)�cudaZxpuZnpu)rsr�r�r@rLr&�onesr+r?r��get_seq_lengthZis_compileablerZ_ignore_causal_mask_sdparbr>Zget_max_cache_shape�5_prepare_4d_causal_attention_mask_with_cache_position�typer�r�Z_unmask_unattended)r9rYr�r�rZpast_seen_tokensZusing_compilable_cacher>r�r��causal_mask�	min_dtyper.r.r/�_update_causal_mask%s`130����131132133�134��135z'BartPreTrainedModel._update_causal_maskr�r�r>r�cKsD|dur|��dkr|}|St�|�j}tj||f|||jd�}|dkr+tj|dd�}|tj||jd�|�dd�k9}|dddd�dd�f�	|ddd�}|dur�|�136�}|jd}	|dd�dd�dd�d|	�f|dd�dddd�f�|j�}137|138dk}139|dd�dd�dd�d|	�f�
|140|�|dd�dd�dd�d|	�f<|S)	aM141        Creates a causal 4D mask of shape `(batch_size, 1, query_length, key_value_length)` from a 2D mask of shape142        `(batch_size, key_value_length)`, or if the input `attention_mask` is already 4D, do nothing.143 144        Args:145            attention_mask (`torch.Tensor`):146                A 2D attention mask of shape `(batch_size, key_value_length)` or a 4D attention mask of shape147                `(batch_size, 1, query_length, key_value_length)`.148            sequence_length (`int`):149                The sequence length being processed.150            target_length (`int`):151                The target length: when generating with static cache, the mask should be as long as the static cache,152                to account for the 0 padding, the part of the cache that is not filled yet.153            dtype (`torch.dtype`):154                The dtype to use for the 4D attention mask.155            cache_position (`torch.Tensor`):156                Indices depicting the position of the input sequence tokens in the sequence.157            batch_size (`torch.Tensor`):158                Batch size.159        Nr�)Z160fill_valuer>r?r#)Zdiagonalr�r*r)r_r@r�r��fullr?ZtriurAr�rDr,r+�toZmasked_fill)rYr�r�r>r�r�rir�r�Zmask_lengthZpadding_maskr.r.r/r�qs,�� $1616�  �zIBartPreTrainedModel._prepare_4d_causal_attention_mask_with_cache_positionr�r��input_shapecCs�|durM|durM|jjdkrd|vr|}|Sd}|S|jjdkr,t||j|dd�}|S|jjdkrCt|tj�rAt||ddd�}|St||j|dd�}|S)	Nr�rr�r*)r�r�F)Zquery_lengthrrr�)r9r�r�r�r�r.r.r/�_update_cross_attn_mask�s2�������z+BartPreTrainedModel._update_cross_attn_mask)rGrHrIr$�__annotations__�base_model_prefixZsupports_gradient_checkpointingZ"_keys_to_ignore_on_load_unexpectedZ_no_split_modulesZ_skip_keys_device_placementZ_supports_flash_attnZ_supports_sdpaZ_supports_flex_attnZ_can_compile_fullgraphr��propertyr�rr@rLr�rrr��staticmethodrKr>r��Sizer�r.r.r.r/r��sf162163	�164����165�L������6����r�c@�eZdZdd�ZdS)�PretrainedBartModelcC�t�dt�dS�Nz_The class `PretrainedBartModel` has been depreciated, please use `BartPreTrainedModel` instead.��warnings�warn�
FutureWarning�r9r.r.r/�__init_subclass__���z%PretrainedBartModel.__init_subclass__N�rGrHrIr�r.r.r.r/r���r�c@r�)�BartPretrainedModelcCr�r�r�r�r.r.r/r��r�z%BartPretrainedModel.__init_subclass__Nr�r.r.r.r/r��r�r�cs�eZdZdZddedeejf�fdd�
Z							ddee	j166dee	jd	ee	jd167ee	jdee
dee
d
ee
deeeffdd�Z�ZS)�BartEncoderz�168    Transformer encoder consisting of *config.encoder_layers* self attention layers. Each layer is a169    [`BartEncoderLayer`].170 171    Args:172        config: BartConfig173        embed_tokens (nn.Embedding): output embedding174    Nrs�embed_tokenscs�t�����j|_�j|_�j}�j|_�j|_	�j175r!t�|�nd}t
�j||j|d�|_|dur7|j|j_t�j|�|_t��fdd�t�j�D��|_t�|�|_d|_|��dS)NrO�rQc�g|]}t�|d��qS�)rt)r���.0�i�rsr.r/�176<listcomp>��z(BartEncoder.__init__.<locals>.<listcomp>F)r7r8r[Zencoder_layerdrop�	layerdropr�r(rP�max_position_embeddingsZmax_source_positions�scale_embedding�math�sqrtrN�177vocab_sizer�rCr1�embed_positionsr�178ModuleList�rangeZencoder_layersr�r��layernorm_embedding�gradient_checkpointing�	post_init)r9rsr�rnrQr:r�r/r8�s(�179� zBartEncoder.__init__r'rYr\r�r��output_hidden_states�return_dictr�cCs|dur|n|jj}|dur|n|jj}|dur|n|jj}|dur*|dur*td��|dur:|}|�d|jd�}n|durJ|dd�dd�df}ntd��|durW|�|�}|�|�}	|	�	|j180�}	||	}181|�|182�}183tj
j|184|j|jd�}185|�||�}|r�dnd}|r�dnd}|dur�|��dt|j�kr�tdt|j��d	|��d�d186���t|j�D]>\}
}|r�||187f}d}|jr�t�g�}||jkr�d}|r�d
}n||188||dur�||
nd|d�}|d}189|r�||df}q�|r�||190f}|�stdd�|191||fD��St|192||d�S)a~193        Args:194            input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):195                Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you196                provide it.197 198                Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and199                [`PreTrainedTokenizer.__call__`] for details.200 201                [What are input IDs?](../glossary#input-ids)202            attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):203                Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:204 205                - 1 for tokens that are **not masked**,206                - 0 for tokens that are **masked**.207 208                [What are attention masks?](../glossary#attention-mask)209            head_mask (`torch.Tensor` of shape `(encoder_layers, encoder_attention_heads)`, *optional*):210                Mask to nullify selected heads of the attention modules. Mask values selected in `[0, 1]`:211 212                - 1 indicates the head is **not masked**,213                - 0 indicates the head is **masked**.214 215            inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):216                Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.217                This is useful if you want more control over how to convert `input_ids` indices into associated vectors218                than the model's internal embedding lookup matrix.219            output_attentions (`bool`, *optional*):220                Whether or not to return the attentions tensors of all attention layers. See `attentions` under221                returned tensors for more detail.222            output_hidden_states (`bool`, *optional*):223                Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors224                for more detail.225            return_dict (`bool`, *optional*):226                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.227        NzDYou cannot specify both input_ids and inputs_embeds at the same timer*z5You have to specify either input_ids or inputs_embedsr`r.rz&The head_mask should be specified for � layers, but it is for �.FT)NN)r�r�r#cs��|]	}|dur|VqdSrRr.�r��vr.r.r/�	<genexpr>zs�z&BartEncoder.forward.<locals>.<genexpr>��last_hidden_stater��228attentions)rsr�r�use_return_dictr-rgr+r�rr�r?r	rrfr[rbr�rc�lenr��	enumerater@�randrr�r)r9r'rYr\r�r�rr
�inputZ	embed_posr�Zencoder_statesZall_attentions�idxZ
encoder_layerZto_drop�dropout_probability�
layer_outputsr.r.r/rEsv.�229230231�232��233234235��236�zBartEncoder.forwardrR)NNNNNNN)rGrHrIrJr$rrr�r8r@�237LongTensorrLr�r�rr�rrErMr.r.r:r/r��s6	��������238	�r�cs�eZdZdZddedeejf�fdd�
Z													ddee	j239dee	jd	ee	jd240ee	j241dee	jdee	jd
ee
dee	jdeedeedeedeedee	j242deeeffdd�Z�ZS)�BartDecoderz�243    Transformer decoder consisting of *config.decoder_layers* layers. Each layer is a [`BartDecoderLayer`]244 245    Args:246        config: BartConfig247        embed_tokens (nn.Embedding): output embedding248    Nrsr�cs�t�����j|_�j|_�j|_�j|_�j	rt249��j�nd}t
�j�j|j|d�|_|dur6|j|j_t�j�j�|_t��fdd�t�j�D��|_t��j�|_d|_|��dS)NrOr�cr�r�)r�r�r�r.r/r��r�z(BartDecoder.__init__.<locals>.<listcomp>F)r7r8r[Zdecoder_layerdroprr(rPrZmax_target_positionsrrrr�rNrr�rCr1rrrrZdecoder_layersr�r�r	r250r)r9rsr�rQr:r�r/r8�s&�251� zBartDecoder.__init__r'rYr�r�r\�cross_attn_head_maskrr�r�r�rr
r�r�c 
Cs�|252dur|253n|jj}254|dur|n|jj}|	dur|	n|jj}	|dur$|n|jj}|jr7|jr7|	r7t�d�d}	|du|duArCt	d��|durU|}|j255}|�d|d�}n|durm|��dd�}|dd�dd�df}nt	d��|durz|�
|�}|	r�|dur�|dur�tt|jd�t|jd��nt|jd�}|	r�t|t�r�t�d�t�|�}|��dd�\}}|dur�|��nd	}|
dur�tj||||jd256�}
|dur�t�s�||}tj|||jd257�}t|t�r�|jn|}|�|||
|�}|�||||�}|j|||
d�}|�|j�}||}|�|�}tj j!||j!|jd�}|�r d
nd}|258�r'd
nd}|259�r3|du�r3d
nd}t"||gddg�D]+\}}|du�rh|��d	t#|j$�k�rht	d|�dt#|j$��d|��d	�d����q>t%|j$�D]X\}}|�r{||f7}|j�r�t�&g�}||j'k�r��qo||||||du�r�||nd|du�r�||nd||260|	|
d�261}|d	}|262�r�||df7}|du�r�||df7}�qo|�r�||f7}|�s�tdd�|||||fD��St(|||||d�S)an263        Args:264            input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):265                Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you266                provide it.267 268                Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and269                [`PreTrainedTokenizer.__call__`] for details.270 271                [What are input IDs?](../glossary#input-ids)272            attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):273                Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:274 275                - 1 for tokens that are **not masked**,276                - 0 for tokens that are **masked**.277 278                [What are attention masks?](../glossary#attention-mask)279            encoder_hidden_states (`torch.FloatTensor` of shape `(batch_size, encoder_sequence_length, hidden_size)`, *optional*):280                Sequence of hidden-states at the output of the last layer of the encoder. Used in the cross-attention281                of the decoder.282            encoder_attention_mask (`torch.LongTensor` of shape `(batch_size, encoder_sequence_length)`, *optional*):283                Mask to avoid performing cross-attention on padding tokens indices of encoder input_ids. Mask values284                selected in `[0, 1]`:285 286                - 1 for tokens that are **not masked**,287                - 0 for tokens that are **masked**.288 289                [What are attention masks?](../glossary#attention-mask)290            head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):291                Mask to nullify selected heads of the attention modules. Mask values selected in `[0, 1]`:292 293                - 1 indicates the head is **not masked**,294                - 0 indicates the head is **masked**.295 296            cross_attn_head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):297                Mask to nullify selected heads of the cross-attention modules in the decoder to avoid performing298                cross-attention on hidden heads. Mask values selected in `[0, 1]`:299 300                - 1 indicates the head is **not masked**,301                - 0 indicates the head is **masked**.302 303            past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):304                It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).305 306                Contains pre-computed hidden-states (key and values in the self-attention blocks and in the307                cross-attention blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.308 309                If `past_key_values` are used, the user can optionally input only the last `decoder_input_ids` (those310                that don't have their past key value states given to this model) of shape `(batch_size, 1)` instead of311                all `decoder_input_ids` of shape `(batch_size, sequence_length)`.312            inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):313                Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.314                This is useful if you want more control over how to convert `input_ids` indices into associated vectors315                than the model's internal embedding lookup matrix.316            output_attentions (`bool`, *optional*):317                Whether or not to return the attentions tensors of all attention layers. See `attentions` under318                returned tensors for more detail.319            output_hidden_states (`bool`, *optional*):320                Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors321                for more detail.322            return_dict (`bool`, *optional*):323                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.324            cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):325                Indices depicting the position of the input sequence tokens in the sequence. It is used to update the326                cache in the correct position and to infer the complete sequence length.327        NzZ`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...FzTYou cannot specify both decoder_input_ids and decoder_inputs_embeds at the same timer*zEYou have to specify either decoder_input_ids or decoder_inputs_embedsr�z�Passing a tuple of `past_key_values` is deprecated and will be removed in Transformers v4.58.0. You should pass an instance of `EncoderDecoderCache` instead, e.g. `past_key_values=EncoderDecoderCache.from_legacy_cache(past_key_values)`.rr�)r=r`r.r\r!zThe `z` should be specified for rr)r�r�r�rr�r�r�r#r5csrrRr.rr.r.r/rzs���z&BartDecoder.forward.<locals>.<genexpr>)rrr�r�cross_attentions))rsr�rr�rr328rbrwrxr-r+rgrcr�r
rr�r�Zfrom_legacy_cacher�r@rAr?r r�r�r�r�rr�r	rrfr[�ziprr�rrrr) r9r'rYr�r�r\r!rr�r�r�rr
r�rr�r�Z329seq_lengthr<Zmask_seq_lengthZself_attn_cacheZ	positionsr�Zall_hidden_statesZall_self_attnsZall_cross_attentionsZ	attn_maskZ	mask_namerZ
decoder_layerrrr.r.r/rE�s�R��330�331��332�����333334335���336337�338�339��zBartDecoder.forwardrR)
NNNNNNNNNNNNN)rGrHrIrJr$rrr�r8r@rrLr�rr�rr�rrErMr.r.r:r/r �sZ��������	�340���
��341�r c&s eZdZddgZdef�fdd�Zdd�Zdd	�Zd342d�Zdd
�Z	e343																d"deej
deejdeej
deej
deejdeejdeejdeeejdeedeejdeejdeedeedeedeedeej
deeeff"d d!��Z�ZS)#�	BartModel�encoder.embed_tokens.weight�decoder.embed_tokens.weightrscslt��|�|j|j}}|jrt�|j�nd}t||j||d�|_	t344||j	�|_t||j	�|_
|��dS)NrOr�)r7r8r(rrrrr�rN�sharedr��encoderr �decoderr)r9rsrPrrQr:r.r/r8�szBartModel.__init__cCs�|jjrB|jjjt�d�kr.|jjjjt�d�kr.|�|j	j|jj�|�|j|jj�dS|�|j	j|j�|�|jj|j�dSdS)N�meta)345rs�tie_word_embeddingsr'rCr?r@r)r��_tie_or_clone_weightsr(r�r.r.r/�_tie_weights�s��zBartModel._tie_weightscC�|jSrR)r'r�r.r.r/�get_input_embeddings��zBartModel.get_input_embeddingscCs||_|j|j_|j|j_dSrR)r'r(r�r)�r9rXr.r.r/�set_input_embeddings�s346zBartModel.set_input_embeddingscCr.rR)r(r�r.r.r/�get_encoder�r0zBartModel.get_encoderNr'rY�decoder_input_ids�decoder_attention_maskr\�decoder_head_maskr!�encoder_outputsrr��decoder_inputs_embedsr�r�rr
r�r�cCsJ|dur|dur|durtd��t||jj|jj�}|
dur |
n|jj}
|dur*|n|jj}|dur4|n|jj}|dur>|n|jj}|durS|j	||||347|
||d�}n$|rwt348|t�swt|dt|�dkrh|dndt|�dkrs|dndd�}|j
|||d||||	|||
|||d�
}|s�||St|j|j|j|j|j|j|j|jd	�S)349�,350        decoder_input_ids (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):351            Indices of decoder input sequence tokens in the vocabulary.352 353            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and354            [`PreTrainedTokenizer.__call__`] for details.355 356            [What are decoder input IDs?](../glossary#decoder-input-ids)357 358            Bart uses the `eos_token_id` as the starting token for `decoder_input_ids` generation. If `past_key_values`359            is used, optionally only the last `decoder_input_ids` have to be input (see `past_key_values`).360 361            For translation and summarization training, `decoder_input_ids` should be provided. If no362            `decoder_input_ids` is provided, the model will create this tensor by shifting the `input_ids` to the right363            for denoising pre-training following the paper.364        decoder_attention_mask (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):365            Default behavior: generate a tensor that ignores pad tokens in `decoder_input_ids`. Causal mask will also366            be used by default.367 368            If you want to change padding behavior, you should read [`modeling_bart._prepare_decoder_attention_mask`]369            and modify to your needs. See diagram 1 in [the paper](https://huggingface.co/papers/1910.13461) for more370            information on the default strategy.371        cross_attn_head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):372            Mask to nullify selected heads of the cross-attention modules in the decoder. Mask values selected in `[0,373            1]`:374 375            - 1 indicates the head is **not masked**,376            - 0 indicates the head is **masked**.377        Nz�If no `decoder_input_ids` or `decoder_inputs_embeds` are passed, `input_ids` cannot be `None`. Please pass either `input_ids` or `decoder_input_ids` or `decoder_inputs_embeds`.)r'rYr\r�r�rr
rr#r5r�
r'rYr�r�r\r!rr�r�r�rr
r�)rr�decoder_hidden_states�decoder_attentionsr"�encoder_last_hidden_stater��encoder_attentions)r-r0rsr(r)r�rr�rr(r�rrr)rrrr�rr")r9r'rYr4r5r\r6r!r7rr�r8r�r�rr
r�Zdecoder_outputsr.r.r/rE�sp3����378���zBartModel.forward�NNNNNNNNNNNNNNNN)rGrHrI�_tied_weights_keysr$r8r-r/r2r3rrr@rrL�listr�rr�rr�rrErMr.r.r:r/r$�sv
��������	�379���
�����380�r$zV381    The BART Model with a language modeling head. Can be used for summarization.382    )Zcustom_introc(sxeZdZdZgd�ZdgZdef�fdd�Zdd�Zd	d383�Z		d,d
e384dee385dede
jf�fdd�
Zd
e386ddfdd�Zdd�Ze																	d-deejdeejdeejdeejdeejdeejdeejdeeejdeed eejd!eejd"eejd#eed$eed%eed&eed'eejdeeeff$d(d)��Zd"ejfd*d+�Z�ZS).�BartForConditionalGenerationr�)r%r&�lm_head.weight�final_logits_biasrscsXt��|�t|�|_|�dt�d|jjjf��t	j387|j|jjjdd�|_|�
�dS)NrDr#Fru)r7r8r$r��register_bufferr@�zerosr'r2rryr��lm_headr�r9rsr:r.r/r82s388389z%BartForConditionalGeneration.__init__cC�390|j��SrR)r�r3r�r.r.r/r3;�391z(BartForConditionalGeneration.get_encodercCrIrR)r��get_decoderr�r.r.r/rK>rJz(BartForConditionalGeneration.get_decoderNT�new_num_tokens�pad_to_multiple_of�
mean_resizingr�cs&t��|||�}|�|jjd�|S)Nr)r7�resize_token_embeddings�_resize_final_logits_biasrCr+)r9rLrMrNZnew_embeddingsr:r.r/rOAsz4BartForConditionalGeneration.resize_token_embeddingscCsj|jjd}||kr|jdd�d|�f}ntjd||f|jjd�}tj|j|gdd�}|�d|�dS)Nr*r#r�r^rD)rDr+r@rFr?�catrE)r9rLZold_num_tokensZnew_biasZ392extra_biasr.r.r/rPHsz6BartForConditionalGeneration._resize_final_logits_biascCs,|jjr|j��|�|j|jj�dSdSrR)rsr+r�r-r,rGr'r�r.r.r/r-Qs393�z)BartForConditionalGeneration._tie_weightsr'rYr4r5r\r6r!r7rr�r8�labelsr�r�rr
r�cCs.|dur|n|jj}|dur)|
rt�d�d}
|dur)|dur)t||jj|jj�}|j|f||||||||	|394||
||||d��}|�|d�}||j	�395|j�}d}|durm|�396|j�}t�}||�
d|jj�|�
d��}|s�|f|dd�}|dur�|f|S|St|||j|j|j|j|j|j|jd�	S)	ai397        decoder_input_ids (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):398            Indices of decoder input sequence tokens in the vocabulary.399 400            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and401            [`PreTrainedTokenizer.__call__`] for details.402 403            [What are decoder input IDs?](../glossary#decoder-input-ids)404 405            Bart uses the `eos_token_id` as the starting token for `decoder_input_ids` generation. If `past_key_values`406            is used, optionally only the last `decoder_input_ids` have to be input (see `past_key_values`).407 408            For translation and summarization training, `decoder_input_ids` should be provided. If no409            `decoder_input_ids` is provided, the model will create this tensor by shifting the `input_ids` to the right410            for denoising pre-training following the paper.411        decoder_attention_mask (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):412            Default behavior: generate a tensor that ignores pad tokens in `decoder_input_ids`. Causal mask will also413            be used by default.414 415            If you want to change padding behavior, you should read [`modeling_bart._prepare_decoder_attention_mask`]416            and modify to your needs. See diagram 1 in [the paper](https://huggingface.co/papers/1910.13461) for more417            information on the default strategy.418        cross_attn_head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):419            Mask to nullify selected heads of the cross-attention modules in the decoder. Mask values selected in `[0,420            1]`:421 422            - 1 indicates the head is **not masked**,423            - 0 indicates the head is **masked**.424        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):425            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,426            config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored427            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.428 429        Example summarization:430 431        ```python432        >>> from transformers import AutoTokenizer, BartForConditionalGeneration433 434        >>> model = BartForConditionalGeneration.from_pretrained("facebook/bart-large-cnn")435        >>> tokenizer = AutoTokenizer.from_pretrained("facebook/bart-large-cnn")436 437        >>> ARTICLE_TO_SUMMARIZE = (438        ...     "PG&E stated it scheduled the blackouts in response to forecasts for high winds "439        ...     "amid dry conditions. The aim is to reduce the risk of wildfires. Nearly 800 thousand customers were "440        ...     "scheduled to be affected by the shutoffs which were expected to last through at least midday tomorrow."441        ... )442        >>> inputs = tokenizer([ARTICLE_TO_SUMMARIZE], max_length=1024, return_tensors="pt")443 444        >>> # Generate Summary445        >>> summary_ids = model.generate(inputs["input_ids"], num_beams=2, min_length=0, max_length=20)446        >>> tokenizer.batch_decode(summary_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]447        'PG&E scheduled the blackouts in response to forecasts for high winds amid dry conditions'448        ```449 450        Mask filling example:451 452        ```python453        >>> from transformers import AutoTokenizer, BartForConditionalGeneration454 455        >>> tokenizer = AutoTokenizer.from_pretrained("facebook/bart-base")456        >>> model = BartForConditionalGeneration.from_pretrained("facebook/bart-base")457 458        >>> TXT = "My friends are <mask> but they eat too many carbs."459        >>> input_ids = tokenizer([TXT], return_tensors="pt")["input_ids"]460        >>> logits = model(input_ids).logits461 462        >>> masked_index = (input_ids[0] == tokenizer.mask_token_id).nonzero().item()463        >>> probs = logits[0, masked_index].softmax(dim=0)464        >>> values, predictions = probs.topk(5)465 466        >>> tokenizer.decode(predictions).split()467        ['not', 'good', 'healthy', 'great', 'very']468        ```469        NzJThe `use_cache` argument is changed to `False` since `labels` is provided.F)rYr4r7r5r\r6r!rr�r8r�r�rr
r�rr*r#�	�loss�logitsrr;r<r"r=r�r>)rsrrw�warningr0r(r)r�rGrDr�r?rrgrrrr;r<r"r=r�r>)r9r'rYr4r5r\r6r!r7rr�r8rRr�r�rr
r�r�Z	lm_logitsZmasked_lm_loss�loss_fct�outputr.r.r/rEVsb_470����z$BartForConditionalGeneration.forwardcCst||jj|jj�SrR)r0rsr(r))r9rRr.r.r/�%prepare_decoder_input_ids_from_labels�szBBartForConditionalGeneration.prepare_decoder_input_ids_from_labels)NT�NNNNNNNNNNNNNNNNN)rGrHrIr�r@Z_keys_to_ignore_on_load_missingr$r8r3rKrKrr�rr�rOrPr-rr@rrLrAr�rrr�rrErYrMr.r.r:r/rB(s�	�����	��������	�471���
������472�rBz�473    Bart model with a sequence classification/head on top (a linear layer on top of the pooled output) e.g. for GLUE474    tasks.475    c&seZdZddgZdef�fdd�Ze																ddeej	deej476d	eej	d477eej	deej478deej479d
eej480deeejdeejdeejdeej	dee
dee
dee
dee
deej	deeeff"dd��Z�ZS)�BartForSequenceClassificationr%r&rscsBt�j|fi|��t|�|_t|j|j|j|j�|_|�	�dSrR)481r7r8r$r�r�r��482num_labelsZclassifier_dropout�classification_headr)r9rsrir:r.r/r8�s483�z&BartForSequenceClassification.__init__Nr'rYr4r5r\r6r!r7r�r8rRr�r�rr
r�r�cCs<|dur|n|jj}|durd}|dur!|	dur!td|jj����|j|||||||||	|484||
|||d�}|d}|�|jj��|j	�}t485t�|�
d���dkrTtd��||dd�f�|�d�d|�d��dd�ddd�f}|�|�}d}|dur�|�|j	�}|jjdur�|jjdkr�d	|j_n|jjdkr�|jtjks�|jtjkr�d486|j_nd|j_|jjd	kr�t�}|jjdkr�||��|���}n,|||�}n&|jjd487kr�t�}||�d|jj�|�d��}n|jjdkr�t�}|||�}|�s488|f|dd�}|du�r|f|S|St|||j|j|j|j|j |j!|j"d�	S)
aV489        decoder_input_ids (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):490            Indices of decoder input sequence tokens in the vocabulary.491 492            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and493            [`PreTrainedTokenizer.__call__`] for details.494 495            [What are decoder input IDs?](../glossary#decoder-input-ids)496 497            Bart uses the `eos_token_id` as the starting token for `decoder_input_ids` generation. If `past_key_values`498            is used, optionally only the last `decoder_input_ids` have to be input (see `past_key_values`).499 500            For translation and summarization training, `decoder_input_ids` should be provided. If no501            `decoder_input_ids` is provided, the model will create this tensor by shifting the `input_ids` to the right502            for denoising pre-training following the paper.503        decoder_attention_mask (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):504            Default behavior: generate a tensor that ignores pad tokens in `decoder_input_ids`. Causal mask will also505            be used by default.506 507            If you want to change padding behavior, you should read [`modeling_bart._prepare_decoder_attention_mask`]508            and modify to your needs. See diagram 1 in [the paper](https://huggingface.co/papers/1910.13461) for more509            information on the default strategy.510        cross_attn_head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):511            Mask to nullify selected heads of the cross-attention modules in the decoder. Mask values selected in `[0,512            1]`:513 514            - 1 indicates the head is **not masked**,515            - 0 indicates the head is **masked**.516        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):517            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,518            config.num_labels - 1]`. If `config.num_labels > 1` a classification loss is computed (Cross-Entropy).519        NFz8Passing input embeddings is currently not supported for �rYr4r5r\r6r!r7r�r8r�r�rr
r�rr#z7All examples must have the same number of <eos> tokens.r*Z520regressionZsingle_label_classificationZmulti_label_classificationrS)#rsr�NotImplementedErrorr;rGr��eqZeos_token_idr�r?rr@Zunique_consecutive�sumr-rgrcr]Zproblem_typer\r>rBrKr�squeezerrrrr;r<r"r=r�r>)r9r'rYr4r5r\r6r!r7r�r8rRr�r�rr
r�r�r�Zeos_maskZsentence_representationrUrTrWrXr.r.r/rEs�4��$�521522$523524�z%BartForSequenceClassification.forwardr?)rGrHrIr@r$r8rrr@rrLrAr�r�rr�rrErMr.r.r:r/r[�sn
��������	�525���
�����526�r[c(seZdZddgZ�fdd�Ze																	ddeejdeejdeej	d	eej	d527eejdeejdeejd
ee528ejdeej	deej	deejdeejdeedeedeedeedeej	de
eeff$dd��Z�ZS)�BartForQuestionAnsweringr%r&csBt��|�d|_|j|_t|�|_t�|j|j�|_|�	�dSr4)529r7r8r\r$r�rry�hidden_size�530qa_outputsrrHr:r.r/r8�s531z!BartForQuestionAnswering.__init__Nr'rYr4r5r\r6r!r7�start_positions�
end_positionsr�r8r�r�rr
r�r�cCs||dur|n|jj}|	dur|532durd}
|j|||||||||||
||||d�}|d}|�|�}|jddd�\}}|�d���}|�d���}d}|	dur�|533dur�t|	���dkr_|	�d�}	t|534���dkrl|535�d�}536|�d�}|	�	d|�}	|537�	d|�}538t539|d�}|||	�}|||540�}||d	}|s�||f|dd�}|dur�|f|S|St||||j|j
|j|j|j|j|jd541542S)r9NFr^rr#r*r^)Zignore_indexr5)543rT�start_logits�544end_logitsrr;r<r"r=r�r>)rsrr�re�splitrbrhrrcr�rrrr;r<r"r=r�r>)r9r'rYr4r5r\r6r!r7rfrgr�r8r�r�rr
r�r�Zsequence_outputrUrhriZ545total_lossZ
ignored_indexrWZ546start_lossZend_lossrXr.r.r/rE�sr2�547548549550551552553�554��z BartForQuestionAnswering.forwardrZ)rGrHrIr@r8rrr@rLrrAr�r�rr�rrErMr.r.r:r/rc�st��������	�555���
������556�rccs(eZdZdZ�fdd�Zdd�Z�ZS)�BartDecoderWrapperz�557    This wrapper class is a helper class to correctly load pretrained checkpoints when the causal language model is558    used in combination with the [`EncoderDecoderModel`] framework.559    cst��|�t|�|_dSrR)r7r8r r)rHr:r.r/r8szBartDecoderWrapper.__init__cOs|j|i|��SrR)r))r9�argsrir.r.r/rEszBartDecoderWrapper.forward)rGrHrIrJr8rErMr.r.r:r/rksrkzu560    BART decoder with a language modeling head on top (linear layer with weights tied to the input embeddings).561    c"s�eZdZdgZ�fdd�Zdd�Zdd�Zdd	�Zd562d�Ze															dd
e563ejde564ej
de565ejde566ejde567ej
de568ej
de569ede570ejde571ejde572ede573ede574ede575ede576ejdeeeffdd��Z�ZS)�BartForCausalLMrCcsDd|_d|_t��|�t|�|_tj|j|j	dd�|_577|��dS)NTFru)rpZis_encoder_decoderr7r8rkr�rryrdrrGrrHr:r.r/r8"s578zBartForCausalLM.__init__cCs579|jjjSrR�r�r)r�r�r.r.r/r/-rJz$BartForCausalLM.get_input_embeddingscCs||jj_dSrRrnr1r.r.r/r20sz$BartForCausalLM.set_input_embeddingscCs||j_dSrR�r�r))r9r)r.r.r/�set_decoder3szBartForCausalLM.set_decodercCs|jjSrRror�r.r.r/rK6szBartForCausalLM.get_decoderNr'rYr�r�r\r!rr�rRr�r�rr
r�r�cCs�|dur|n|jj}|dur|n|jj}|
dur|
n|jj}
|jj|||||||||580|||
|d�
}|�|d�}d}|	durU|	�|j�}	t	�}||�581d|jj�|	�582d��}|
sk|f|dd�}|duri|f|S|St|||j
|j|j|jd�S)a�583        cross_attn_head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):584            Mask to nullify selected heads of the cross-attention modules. Mask values selected in `[0, 1]`:585 586            - 1 indicates the head is **not masked**,587            - 0 indicates the head is **masked**.588        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):589            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,590            config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored591            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.592 593        Example:594 595        ```python596        >>> from transformers import AutoTokenizer, BartForCausalLM597 598        >>> tokenizer = AutoTokenizer.from_pretrained("facebook/bart-base")599        >>> model = BartForCausalLM.from_pretrained("facebook/bart-base", add_cross_attention=False)600        >>> assert model.config.is_decoder, f"{model.__class__} has to be configured as a decoder."601        >>> inputs = tokenizer("Hello, my dog is cute", return_tensors="pt")602        >>> outputs = model(**inputs)603 604        >>> logits = outputs.logits605        >>> expected_shape = [1, inputs.input_ids.shape[-1], model.config.vocab_size]606        >>> list(logits.shape) == expected_shape607        True608        ```Nr:rr*r#)rTrUrr�rr")rsr�rrr�r)rGr�r?rrgrrrr�rr")r9r'rYr�r�r\r!rr�rRr�r�rr
r�r�rUrTrWrXr.r.r/rE9sH.���zBartForCausalLM.forward)NNNNNNNNNNNNNN)rGrHrIr@r8r/r2rprKrrr@rrLr�rr�rr�rrErMr.r.r:r/rmsj��������	�609���
���610�rm)rmrBrcr[r$r�r�r�)NrTN)RrJrr��typingrrrr@rZtorch.nnrrrZactivationsr611Zcache_utilsrrr
Z612generationrZmodeling_attn_mask_utilsrrrZmodeling_flash_attention_utilsrZmodeling_layersrZmodeling_outputsrrrrrrrZmodeling_utilsrrZprocessing_utilsr�utilsrrr r!Zutils.deprecationr"Zconfiguration_bartr$Zintegrations.flex_attentionr%r&Z613get_loggerrGrwrLrKr0r�r1rN�ModulerSrlrmr�r�r�r�r�r�r�r r$rBr[rcrkrm�__all__r.r.r.r/�<module>s�$	614��������615�Gud$616 �D��u
Aluode/PerceptionLabPortable · CoolFace