Jingya/tiny-random-bert-remote-code
0157
1# coding=utf-82# Copyright 2018 The Google AI Language Team Authors and The HuggingFace Inc. team.3# Copyright (c) 2018, NVIDIA CORPORATION. All rights reserved.4#5# Licensed under the Apache License, Version 2.0 (the "License");6# you may not use this file except in compliance with the License.7# You may obtain a copy of the License at8#9# http://www.apache.org/licenses/LICENSE-2.010#11# Unless required by applicable law or agreed to in writing, software12# distributed under the License is distributed on an "AS IS" BASIS,13# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.14# See the License for the specific language governing permissions and15# limitations under the License.16"""PyTorch BERT model."""17 18 19import math20import os21import warnings22from dataclasses import dataclass23from typing import List, Optional, Tuple, Union24 25import torch26import torch.utils.checkpoint27from torch import nn28from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss29 30from ...activations import ACT2FN31from ...modeling_outputs import (32 BaseModelOutputWithPastAndCrossAttentions,33 BaseModelOutputWithPoolingAndCrossAttentions,34 CausalLMOutputWithCrossAttentions,35 MaskedLMOutput,36 MultipleChoiceModelOutput,37 NextSentencePredictorOutput,38 QuestionAnsweringModelOutput,39 SequenceClassifierOutput,40 TokenClassifierOutput,41)42from ...modeling_utils import PreTrainedModel43from ...pytorch_utils import apply_chunking_to_forward, find_pruneable_heads_and_indices, prune_linear_layer44from ...utils import (45 ModelOutput,46 add_code_sample_docstrings,47 add_start_docstrings,48 add_start_docstrings_to_model_forward,49 logging,50 replace_return_docstrings,51)52from .configuration_bert import BertConfig53 54 55logger = logging.get_logger(__name__)56 57_CHECKPOINT_FOR_DOC = "bert-base-uncased"58_CONFIG_FOR_DOC = "BertConfig"59 60# TokenClassification docstring61_CHECKPOINT_FOR_TOKEN_CLASSIFICATION = "dbmdz/bert-large-cased-finetuned-conll03-english"62_TOKEN_CLASS_EXPECTED_OUTPUT = (63 "['O', 'I-ORG', 'I-ORG', 'I-ORG', 'O', 'O', 'O', 'O', 'O', 'I-LOC', 'O', 'I-LOC', 'I-LOC'] "64)65_TOKEN_CLASS_EXPECTED_LOSS = 0.0166 67# QuestionAnswering docstring68_CHECKPOINT_FOR_QA = "deepset/bert-base-cased-squad2"69_QA_EXPECTED_OUTPUT = "'a nice puppet'"70_QA_EXPECTED_LOSS = 7.4171_QA_TARGET_START_INDEX = 1472_QA_TARGET_END_INDEX = 1573 74# SequenceClassification docstring75_CHECKPOINT_FOR_SEQUENCE_CLASSIFICATION = "textattack/bert-base-uncased-yelp-polarity"76_SEQ_CLASS_EXPECTED_OUTPUT = "'LABEL_1'"77_SEQ_CLASS_EXPECTED_LOSS = 0.0178 79 80BERT_PRETRAINED_MODEL_ARCHIVE_LIST = [81 "bert-base-uncased",82 "bert-large-uncased",83 "bert-base-cased",84 "bert-large-cased",85 "bert-base-multilingual-uncased",86 "bert-base-multilingual-cased",87 "bert-base-chinese",88 "bert-base-german-cased",89 "bert-large-uncased-whole-word-masking",90 "bert-large-cased-whole-word-masking",91 "bert-large-uncased-whole-word-masking-finetuned-squad",92 "bert-large-cased-whole-word-masking-finetuned-squad",93 "bert-base-cased-finetuned-mrpc",94 "bert-base-german-dbmdz-cased",95 "bert-base-german-dbmdz-uncased",96 "cl-tohoku/bert-base-japanese",97 "cl-tohoku/bert-base-japanese-whole-word-masking",98 "cl-tohoku/bert-base-japanese-char",99 "cl-tohoku/bert-base-japanese-char-whole-word-masking",100 "TurkuNLP/bert-base-finnish-cased-v1",101 "TurkuNLP/bert-base-finnish-uncased-v1",102 "wietsedv/bert-base-dutch-cased",103 # See all BERT models at https://huggingface.co/models?filter=bert104]105 106 107def load_tf_weights_in_bert(model, config, tf_checkpoint_path):108 """Load tf checkpoints in a pytorch model."""109 try:110 import re111 112 import numpy as np113 import tensorflow as tf114 except ImportError:115 logger.error(116 "Loading a TensorFlow model in PyTorch, requires TensorFlow to be installed. Please see "117 "https://www.tensorflow.org/install/ for installation instructions."118 )119 raise120 tf_path = os.path.abspath(tf_checkpoint_path)121 logger.info(f"Converting TensorFlow checkpoint from {tf_path}")122 # Load weights from TF model123 init_vars = tf.train.list_variables(tf_path)124 names = []125 arrays = []126 for name, shape in init_vars:127 logger.info(f"Loading TF weight {name} with shape {shape}")128 array = tf.train.load_variable(tf_path, name)129 names.append(name)130 arrays.append(array)131 132 for name, array in zip(names, arrays):133 name = name.split("/")134 # adam_v and adam_m are variables used in AdamWeightDecayOptimizer to calculated m and v135 # which are not required for using pretrained model136 if any(137 n in ["adam_v", "adam_m", "AdamWeightDecayOptimizer", "AdamWeightDecayOptimizer_1", "global_step"]138 for n in name139 ):140 logger.info(f"Skipping {'/'.join(name)}")141 continue142 pointer = model143 for m_name in name:144 if re.fullmatch(r"[A-Za-z]+_\d+", m_name):145 scope_names = re.split(r"_(\d+)", m_name)146 else:147 scope_names = [m_name]148 if scope_names[0] == "kernel" or scope_names[0] == "gamma":149 pointer = getattr(pointer, "weight")150 elif scope_names[0] == "output_bias" or scope_names[0] == "beta":151 pointer = getattr(pointer, "bias")152 elif scope_names[0] == "output_weights":153 pointer = getattr(pointer, "weight")154 elif scope_names[0] == "squad":155 pointer = getattr(pointer, "classifier")156 else:157 try:158 pointer = getattr(pointer, scope_names[0])159 except AttributeError:160 logger.info(f"Skipping {'/'.join(name)}")161 continue162 if len(scope_names) >= 2:163 num = int(scope_names[1])164 pointer = pointer[num]165 if m_name[-11:] == "_embeddings":166 pointer = getattr(pointer, "weight")167 elif m_name == "kernel":168 array = np.transpose(array)169 try:170 if pointer.shape != array.shape:171 raise ValueError(f"Pointer shape {pointer.shape} and array shape {array.shape} mismatched")172 except AssertionError as e:173 e.args += (pointer.shape, array.shape)174 raise175 logger.info(f"Initialize PyTorch weight {name}")176 pointer.data = torch.from_numpy(array)177 return model178 179 180class BertEmbeddings(nn.Module):181 """Construct the embeddings from word, position and token_type embeddings."""182 183 def __init__(self, config):184 super().__init__()185 self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id)186 self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size)187 self.token_type_embeddings = nn.Embedding(config.type_vocab_size, config.hidden_size)188 189 # self.LayerNorm is not snake-cased to stick with TensorFlow model variable name and be able to load190 # any TensorFlow checkpoint file191 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)192 self.dropout = nn.Dropout(config.hidden_dropout_prob)193 # position_ids (1, len position emb) is contiguous in memory and exported when serialized194 self.position_embedding_type = getattr(config, "position_embedding_type", "absolute")195 self.register_buffer("position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)))196 self.register_buffer(197 "token_type_ids", torch.zeros(self.position_ids.size(), dtype=torch.long), persistent=False198 )199 200 def forward(201 self,202 input_ids: Optional[torch.LongTensor] = None,203 token_type_ids: Optional[torch.LongTensor] = None,204 position_ids: Optional[torch.LongTensor] = None,205 inputs_embeds: Optional[torch.FloatTensor] = None,206 past_key_values_length: int = 0,207 ) -> torch.Tensor:208 if input_ids is not None:209 input_shape = input_ids.size()210 else:211 input_shape = inputs_embeds.size()[:-1]212 213 seq_length = input_shape[1]214 215 if position_ids is None:216 position_ids = self.position_ids[:, past_key_values_length : seq_length + past_key_values_length]217 218 # Setting the token_type_ids to the registered buffer in constructor where it is all zeros, which usually occurs219 # when its auto-generated, registered buffer helps users when tracing the model without passing token_type_ids, solves220 # issue #5664221 if token_type_ids is None:222 if hasattr(self, "token_type_ids"):223 buffered_token_type_ids = self.token_type_ids[:, :seq_length]224 buffered_token_type_ids_expanded = buffered_token_type_ids.expand(input_shape[0], seq_length)225 token_type_ids = buffered_token_type_ids_expanded226 else:227 token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=self.position_ids.device)228 229 if inputs_embeds is None:230 inputs_embeds = self.word_embeddings(input_ids)231 token_type_embeddings = self.token_type_embeddings(token_type_ids)232 233 embeddings = inputs_embeds + token_type_embeddings234 if self.position_embedding_type == "absolute":235 position_embeddings = self.position_embeddings(position_ids)236 embeddings += position_embeddings237 embeddings = self.LayerNorm(embeddings)238 embeddings = self.dropout(embeddings)239 return embeddings240 241 242class BertSelfAttention(nn.Module):243 def __init__(self, config, position_embedding_type=None):244 super().__init__()245 if config.hidden_size % config.num_attention_heads != 0 and not hasattr(config, "embedding_size"):246 raise ValueError(247 f"The hidden size ({config.hidden_size}) is not a multiple of the number of attention "248 f"heads ({config.num_attention_heads})"249 )250 251 self.num_attention_heads = config.num_attention_heads252 self.attention_head_size = int(config.hidden_size / config.num_attention_heads)253 self.all_head_size = self.num_attention_heads * self.attention_head_size254 255 self.query = nn.Linear(config.hidden_size, self.all_head_size)256 self.key = nn.Linear(config.hidden_size, self.all_head_size)257 self.value = nn.Linear(config.hidden_size, self.all_head_size)258 259 self.dropout = nn.Dropout(config.attention_probs_dropout_prob)260 self.position_embedding_type = position_embedding_type or getattr(261 config, "position_embedding_type", "absolute"262 )263 if self.position_embedding_type == "relative_key" or self.position_embedding_type == "relative_key_query":264 self.max_position_embeddings = config.max_position_embeddings265 self.distance_embedding = nn.Embedding(2 * config.max_position_embeddings - 1, self.attention_head_size)266 267 self.is_decoder = config.is_decoder268 269 def transpose_for_scores(self, x: torch.Tensor) -> torch.Tensor:270 new_x_shape = x.size()[:-1] + (self.num_attention_heads, self.attention_head_size)271 x = x.view(new_x_shape)272 return x.permute(0, 2, 1, 3)273 274 def forward(275 self,276 hidden_states: torch.Tensor,277 attention_mask: Optional[torch.FloatTensor] = None,278 head_mask: Optional[torch.FloatTensor] = None,279 encoder_hidden_states: Optional[torch.FloatTensor] = None,280 encoder_attention_mask: Optional[torch.FloatTensor] = None,281 past_key_value: Optional[Tuple[Tuple[torch.FloatTensor]]] = None,282 output_attentions: Optional[bool] = False,283 ) -> Tuple[torch.Tensor]:284 mixed_query_layer = self.query(hidden_states)285 286 # If this is instantiated as a cross-attention module, the keys287 # and values come from an encoder; the attention mask needs to be288 # such that the encoder's padding tokens are not attended to.289 is_cross_attention = encoder_hidden_states is not None290 291 if is_cross_attention and past_key_value is not None:292 # reuse k,v, cross_attentions293 key_layer = past_key_value[0]294 value_layer = past_key_value[1]295 attention_mask = encoder_attention_mask296 elif is_cross_attention:297 key_layer = self.transpose_for_scores(self.key(encoder_hidden_states))298 value_layer = self.transpose_for_scores(self.value(encoder_hidden_states))299 attention_mask = encoder_attention_mask300 elif past_key_value is not None:301 key_layer = self.transpose_for_scores(self.key(hidden_states))302 value_layer = self.transpose_for_scores(self.value(hidden_states))303 key_layer = torch.cat([past_key_value[0], key_layer], dim=2)304 value_layer = torch.cat([past_key_value[1], value_layer], dim=2)305 else:306 key_layer = self.transpose_for_scores(self.key(hidden_states))307 value_layer = self.transpose_for_scores(self.value(hidden_states))308 309 query_layer = self.transpose_for_scores(mixed_query_layer)310 311 use_cache = past_key_value is not None312 if self.is_decoder:313 # if cross_attention save Tuple(torch.Tensor, torch.Tensor) of all cross attention key/value_states.314 # Further calls to cross_attention layer can then reuse all cross-attention315 # key/value_states (first "if" case)316 # if uni-directional self-attention (decoder) save Tuple(torch.Tensor, torch.Tensor) of317 # all previous decoder key/value_states. Further calls to uni-directional self-attention318 # can concat previous decoder key/value_states to current projected key/value_states (third "elif" case)319 # if encoder bi-directional self-attention `past_key_value` is always `None`320 past_key_value = (key_layer, value_layer)321 322 # Take the dot product between "query" and "key" to get the raw attention scores.323 attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))324 325 if self.position_embedding_type == "relative_key" or self.position_embedding_type == "relative_key_query":326 query_length, key_length = query_layer.shape[2], key_layer.shape[2]327 if use_cache:328 position_ids_l = torch.tensor(key_length - 1, dtype=torch.long, device=hidden_states.device).view(329 -1, 1330 )331 else:332 position_ids_l = torch.arange(query_length, dtype=torch.long, device=hidden_states.device).view(-1, 1)333 position_ids_r = torch.arange(key_length, dtype=torch.long, device=hidden_states.device).view(1, -1)334 distance = position_ids_l - position_ids_r335 336 positional_embedding = self.distance_embedding(distance + self.max_position_embeddings - 1)337 positional_embedding = positional_embedding.to(dtype=query_layer.dtype) # fp16 compatibility338 339 if self.position_embedding_type == "relative_key":340 relative_position_scores = torch.einsum("bhld,lrd->bhlr", query_layer, positional_embedding)341 attention_scores = attention_scores + relative_position_scores342 elif self.position_embedding_type == "relative_key_query":343 relative_position_scores_query = torch.einsum("bhld,lrd->bhlr", query_layer, positional_embedding)344 relative_position_scores_key = torch.einsum("bhrd,lrd->bhlr", key_layer, positional_embedding)345 attention_scores = attention_scores + relative_position_scores_query + relative_position_scores_key346 347 attention_scores = attention_scores / math.sqrt(self.attention_head_size)348 if attention_mask is not None:349 # Apply the attention mask is (precomputed for all layers in BertModel forward() function)350 attention_scores = attention_scores + attention_mask351 352 # Normalize the attention scores to probabilities.353 attention_probs = nn.functional.softmax(attention_scores, dim=-1)354 355 # This is actually dropping out entire tokens to attend to, which might356 # seem a bit unusual, but is taken from the original Transformer paper.357 attention_probs = self.dropout(attention_probs)358 359 # Mask heads if we want to360 if head_mask is not None:361 attention_probs = attention_probs * head_mask362 363 context_layer = torch.matmul(attention_probs, value_layer)364 365 context_layer = context_layer.permute(0, 2, 1, 3).contiguous()366 new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)367 context_layer = context_layer.view(new_context_layer_shape)368 369 outputs = (context_layer, attention_probs) if output_attentions else (context_layer,)370 371 if self.is_decoder:372 outputs = outputs + (past_key_value,)373 return outputs374 375 376class BertSelfOutput(nn.Module):377 def __init__(self, config):378 super().__init__()379 self.dense = nn.Linear(config.hidden_size, config.hidden_size)380 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)381 self.dropout = nn.Dropout(config.hidden_dropout_prob)382 383 def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:384 hidden_states = self.dense(hidden_states)385 hidden_states = self.dropout(hidden_states)386 hidden_states = self.LayerNorm(hidden_states + input_tensor)387 return hidden_states388 389 390class BertAttention(nn.Module):391 def __init__(self, config, position_embedding_type=None):392 super().__init__()393 self.self = BertSelfAttention(config, position_embedding_type=position_embedding_type)394 self.output = BertSelfOutput(config)395 self.pruned_heads = set()396 397 def prune_heads(self, heads):398 if len(heads) == 0:399 return400 heads, index = find_pruneable_heads_and_indices(401 heads, self.self.num_attention_heads, self.self.attention_head_size, self.pruned_heads402 )403 404 # Prune linear layers405 self.self.query = prune_linear_layer(self.self.query, index)406 self.self.key = prune_linear_layer(self.self.key, index)407 self.self.value = prune_linear_layer(self.self.value, index)408 self.output.dense = prune_linear_layer(self.output.dense, index, dim=1)409 410 # Update hyper params and store pruned heads411 self.self.num_attention_heads = self.self.num_attention_heads - len(heads)412 self.self.all_head_size = self.self.attention_head_size * self.self.num_attention_heads413 self.pruned_heads = self.pruned_heads.union(heads)414 415 def forward(416 self,417 hidden_states: torch.Tensor,418 attention_mask: Optional[torch.FloatTensor] = None,419 head_mask: Optional[torch.FloatTensor] = None,420 encoder_hidden_states: Optional[torch.FloatTensor] = None,421 encoder_attention_mask: Optional[torch.FloatTensor] = None,422 past_key_value: Optional[Tuple[Tuple[torch.FloatTensor]]] = None,423 output_attentions: Optional[bool] = False,424 ) -> Tuple[torch.Tensor]:425 self_outputs = self.self(426 hidden_states,427 attention_mask,428 head_mask,429 encoder_hidden_states,430 encoder_attention_mask,431 past_key_value,432 output_attentions,433 )434 attention_output = self.output(self_outputs[0], hidden_states)435 outputs = (attention_output,) + self_outputs[1:] # add attentions if we output them436 return outputs437 438 439class BertIntermediate(nn.Module):440 def __init__(self, config):441 super().__init__()442 self.dense = nn.Linear(config.hidden_size, config.intermediate_size)443 if isinstance(config.hidden_act, str):444 self.intermediate_act_fn = ACT2FN[config.hidden_act]445 else:446 self.intermediate_act_fn = config.hidden_act447 448 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:449 hidden_states = self.dense(hidden_states)450 hidden_states = self.intermediate_act_fn(hidden_states)451 return hidden_states452 453 454class BertOutput(nn.Module):455 def __init__(self, config):456 super().__init__()457 self.dense = nn.Linear(config.intermediate_size, config.hidden_size)458 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)459 self.dropout = nn.Dropout(config.hidden_dropout_prob)460 461 def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:462 hidden_states = self.dense(hidden_states)463 hidden_states = self.dropout(hidden_states)464 hidden_states = self.LayerNorm(hidden_states + input_tensor)465 return hidden_states466 467 468class BertLayer(nn.Module):469 def __init__(self, config):470 super().__init__()471 self.chunk_size_feed_forward = config.chunk_size_feed_forward472 self.seq_len_dim = 1473 self.attention = BertAttention(config)474 self.is_decoder = config.is_decoder475 self.add_cross_attention = config.add_cross_attention476 if self.add_cross_attention:477 if not self.is_decoder:478 raise ValueError(f"{self} should be used as a decoder model if cross attention is added")479 self.crossattention = BertAttention(config, position_embedding_type="absolute")480 self.intermediate = BertIntermediate(config)481 self.output = BertOutput(config)482 483 def forward(484 self,485 hidden_states: torch.Tensor,486 attention_mask: Optional[torch.FloatTensor] = None,487 head_mask: Optional[torch.FloatTensor] = None,488 encoder_hidden_states: Optional[torch.FloatTensor] = None,489 encoder_attention_mask: Optional[torch.FloatTensor] = None,490 past_key_value: Optional[Tuple[Tuple[torch.FloatTensor]]] = None,491 output_attentions: Optional[bool] = False,492 ) -> Tuple[torch.Tensor]:493 # decoder uni-directional self-attention cached key/values tuple is at positions 1,2494 self_attn_past_key_value = past_key_value[:2] if past_key_value is not None else None495 self_attention_outputs = self.attention(496 hidden_states,497 attention_mask,498 head_mask,499 output_attentions=output_attentions,500 past_key_value=self_attn_past_key_value,501 )502 attention_output = self_attention_outputs[0]503 504 # if decoder, the last output is tuple of self-attn cache505 if self.is_decoder:506 outputs = self_attention_outputs[1:-1]507 present_key_value = self_attention_outputs[-1]508 else:509 outputs = self_attention_outputs[1:] # add self attentions if we output attention weights510 511 cross_attn_present_key_value = None512 if self.is_decoder and encoder_hidden_states is not None:513 if not hasattr(self, "crossattention"):514 raise ValueError(515 f"If `encoder_hidden_states` are passed, {self} has to be instantiated with cross-attention layers"516 " by setting `config.add_cross_attention=True`"517 )518 519 # cross_attn cached key/values tuple is at positions 3,4 of past_key_value tuple520 cross_attn_past_key_value = past_key_value[-2:] if past_key_value is not None else None521 cross_attention_outputs = self.crossattention(522 attention_output,523 attention_mask,524 head_mask,525 encoder_hidden_states,526 encoder_attention_mask,527 cross_attn_past_key_value,528 output_attentions,529 )530 attention_output = cross_attention_outputs[0]531 outputs = outputs + cross_attention_outputs[1:-1] # add cross attentions if we output attention weights532 533 # add cross-attn cache to positions 3,4 of present_key_value tuple534 cross_attn_present_key_value = cross_attention_outputs[-1]535 present_key_value = present_key_value + cross_attn_present_key_value536 537 layer_output = apply_chunking_to_forward(538 self.feed_forward_chunk, self.chunk_size_feed_forward, self.seq_len_dim, attention_output539 )540 outputs = (layer_output,) + outputs541 542 # if decoder, return the attn key/values as the last output543 if self.is_decoder:544 outputs = outputs + (present_key_value,)545 546 return outputs547 548 def feed_forward_chunk(self, attention_output):549 intermediate_output = self.intermediate(attention_output)550 layer_output = self.output(intermediate_output, attention_output)551 return layer_output552 553 554class BertEncoder(nn.Module):555 def __init__(self, config):556 super().__init__()557 self.config = config558 self.layer = nn.ModuleList([BertLayer(config) for _ in range(config.num_hidden_layers)])559 self.gradient_checkpointing = False560 561 def forward(562 self,563 hidden_states: torch.Tensor,564 attention_mask: Optional[torch.FloatTensor] = None,565 head_mask: Optional[torch.FloatTensor] = None,566 encoder_hidden_states: Optional[torch.FloatTensor] = None,567 encoder_attention_mask: Optional[torch.FloatTensor] = None,568 past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None,569 use_cache: Optional[bool] = None,570 output_attentions: Optional[bool] = False,571 output_hidden_states: Optional[bool] = False,572 return_dict: Optional[bool] = True,573 ) -> Union[Tuple[torch.Tensor], BaseModelOutputWithPastAndCrossAttentions]:574 all_hidden_states = () if output_hidden_states else None575 all_self_attentions = () if output_attentions else None576 all_cross_attentions = () if output_attentions and self.config.add_cross_attention else None577 578 if self.gradient_checkpointing and self.training:579 if use_cache:580 logger.warning_once(581 "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."582 )583 use_cache = False584 585 next_decoder_cache = () if use_cache else None586 for i, layer_module in enumerate(self.layer):587 if output_hidden_states:588 all_hidden_states = all_hidden_states + (hidden_states,)589 590 layer_head_mask = head_mask[i] if head_mask is not None else None591 past_key_value = past_key_values[i] if past_key_values is not None else None592 593 if self.gradient_checkpointing and self.training:594 595 def create_custom_forward(module):596 def custom_forward(*inputs):597 return module(*inputs, past_key_value, output_attentions)598 599 return custom_forward600 601 layer_outputs = torch.utils.checkpoint.checkpoint(602 create_custom_forward(layer_module),603 hidden_states,604 attention_mask,605 layer_head_mask,606 encoder_hidden_states,607 encoder_attention_mask,608 )609 else:610 layer_outputs = layer_module(611 hidden_states,612 attention_mask,613 layer_head_mask,614 encoder_hidden_states,615 encoder_attention_mask,616 past_key_value,617 output_attentions,618 )619 620 hidden_states = layer_outputs[0]621 if use_cache:622 next_decoder_cache += (layer_outputs[-1],)623 if output_attentions:624 all_self_attentions = all_self_attentions + (layer_outputs[1],)625 if self.config.add_cross_attention:626 all_cross_attentions = all_cross_attentions + (layer_outputs[2],)627 628 if output_hidden_states:629 all_hidden_states = all_hidden_states + (hidden_states,)630 631 if not return_dict:632 return tuple(633 v634 for v in [635 hidden_states,636 next_decoder_cache,637 all_hidden_states,638 all_self_attentions,639 all_cross_attentions,640 ]641 if v is not None642 )643 return BaseModelOutputWithPastAndCrossAttentions(644 last_hidden_state=hidden_states,645 past_key_values=next_decoder_cache,646 hidden_states=all_hidden_states,647 attentions=all_self_attentions,648 cross_attentions=all_cross_attentions,649 )650 651 652class BertPooler(nn.Module):653 def __init__(self, config):654 super().__init__()655 self.dense = nn.Linear(config.hidden_size, config.hidden_size)656 self.activation = nn.Tanh()657 658 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:659 # We "pool" the model by simply taking the hidden state corresponding660 # to the first token.661 first_token_tensor = hidden_states[:, 0]662 pooled_output = self.dense(first_token_tensor)663 pooled_output = self.activation(pooled_output)664 return pooled_output665 666 667class BertPredictionHeadTransform(nn.Module):668 def __init__(self, config):669 super().__init__()670 self.dense = nn.Linear(config.hidden_size, config.hidden_size)671 if isinstance(config.hidden_act, str):672 self.transform_act_fn = ACT2FN[config.hidden_act]673 else:674 self.transform_act_fn = config.hidden_act675 self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)676 677 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:678 hidden_states = self.dense(hidden_states)679 hidden_states = self.transform_act_fn(hidden_states)680 hidden_states = self.LayerNorm(hidden_states)681 return hidden_states682 683 684class BertLMPredictionHead(nn.Module):685 def __init__(self, config):686 super().__init__()687 self.transform = BertPredictionHeadTransform(config)688 689 # The output weights are the same as the input embeddings, but there is690 # an output-only bias for each token.691 self.decoder = nn.Linear(config.hidden_size, config.vocab_size, bias=False)692 693 self.bias = nn.Parameter(torch.zeros(config.vocab_size))694 695 # Need a link between the two variables so that the bias is correctly resized with `resize_token_embeddings`696 self.decoder.bias = self.bias697 698 def forward(self, hidden_states):699 hidden_states = self.transform(hidden_states)700 hidden_states = self.decoder(hidden_states)701 return hidden_states702 703 704class BertOnlyMLMHead(nn.Module):705 def __init__(self, config):706 super().__init__()707 self.predictions = BertLMPredictionHead(config)708 709 def forward(self, sequence_output: torch.Tensor) -> torch.Tensor:710 prediction_scores = self.predictions(sequence_output)711 return prediction_scores712 713 714class BertOnlyNSPHead(nn.Module):715 def __init__(self, config):716 super().__init__()717 self.seq_relationship = nn.Linear(config.hidden_size, 2)718 719 def forward(self, pooled_output):720 seq_relationship_score = self.seq_relationship(pooled_output)721 return seq_relationship_score722 723 724class BertPreTrainingHeads(nn.Module):725 def __init__(self, config):726 super().__init__()727 self.predictions = BertLMPredictionHead(config)728 self.seq_relationship = nn.Linear(config.hidden_size, 2)729 730 def forward(self, sequence_output, pooled_output):731 prediction_scores = self.predictions(sequence_output)732 seq_relationship_score = self.seq_relationship(pooled_output)733 return prediction_scores, seq_relationship_score734 735 736class BertPreTrainedModel(PreTrainedModel):737 """738 An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained739 models.740 """741 742 config_class = BertConfig743 load_tf_weights = load_tf_weights_in_bert744 base_model_prefix = "bert"745 supports_gradient_checkpointing = True746 _keys_to_ignore_on_load_missing = [r"position_ids"]747 748 def _init_weights(self, module):749 """Initialize the weights"""750 if isinstance(module, nn.Linear):751 # Slightly different from the TF version which uses truncated_normal for initialization752 # cf https://github.com/pytorch/pytorch/pull/5617753 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)754 if module.bias is not None:755 module.bias.data.zero_()756 elif isinstance(module, nn.Embedding):757 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)758 if module.padding_idx is not None:759 module.weight.data[module.padding_idx].zero_()760 elif isinstance(module, nn.LayerNorm):761 module.bias.data.zero_()762 module.weight.data.fill_(1.0)763 764 def _set_gradient_checkpointing(self, module, value=False):765 if isinstance(module, BertEncoder):766 module.gradient_checkpointing = value767 768 769@dataclass770class BertForPreTrainingOutput(ModelOutput):771 """772 Output type of [`BertForPreTraining`].773 774 Args:775 loss (*optional*, returned when `labels` is provided, `torch.FloatTensor` of shape `(1,)`):776 Total loss as the sum of the masked language modeling loss and the next sequence prediction777 (classification) loss.778 prediction_logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):779 Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).780 seq_relationship_logits (`torch.FloatTensor` of shape `(batch_size, 2)`):781 Prediction scores of the next sequence prediction (classification) head (scores of True/False continuation782 before SoftMax).783 hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):784 Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of785 shape `(batch_size, sequence_length, hidden_size)`.786 787 Hidden-states of the model at the output of each layer plus the initial embedding outputs.788 attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):789 Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,790 sequence_length)`.791 792 Attentions weights after the attention softmax, used to compute the weighted average in the self-attention793 heads.794 """795 796 loss: Optional[torch.FloatTensor] = None797 prediction_logits: torch.FloatTensor = None798 seq_relationship_logits: torch.FloatTensor = None799 hidden_states: Optional[Tuple[torch.FloatTensor]] = None800 attentions: Optional[Tuple[torch.FloatTensor]] = None801 802 803BERT_START_DOCSTRING = r"""804 805 This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the806 library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads807 etc.)808 809 This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.810 Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage811 and behavior.812 813 Parameters:814 config ([`BertConfig`]): Model configuration class with all the parameters of the model.815 Initializing with a config file does not load the weights associated with the model, only the816 configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.817"""818 819BERT_INPUTS_DOCSTRING = r"""820 Args:821 input_ids (`torch.LongTensor` of shape `({0})`):822 Indices of input sequence tokens in the vocabulary.823 824 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and825 [`PreTrainedTokenizer.__call__`] for details.826 827 [What are input IDs?](../glossary#input-ids)828 attention_mask (`torch.FloatTensor` of shape `({0})`, *optional*):829 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:830 831 - 1 for tokens that are **not masked**,832 - 0 for tokens that are **masked**.833 834 [What are attention masks?](../glossary#attention-mask)835 token_type_ids (`torch.LongTensor` of shape `({0})`, *optional*):836 Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,837 1]`:838 839 - 0 corresponds to a *sentence A* token,840 - 1 corresponds to a *sentence B* token.841 842 [What are token type IDs?](../glossary#token-type-ids)843 position_ids (`torch.LongTensor` of shape `({0})`, *optional*):844 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,845 config.max_position_embeddings - 1]`.846 847 [What are position IDs?](../glossary#position-ids)848 head_mask (`torch.FloatTensor` of shape `(num_heads,)` or `(num_layers, num_heads)`, *optional*):849 Mask to nullify selected heads of the self-attention modules. Mask values selected in `[0, 1]`:850 851 - 1 indicates the head is **not masked**,852 - 0 indicates the head is **masked**.853 854 inputs_embeds (`torch.FloatTensor` of shape `({0}, hidden_size)`, *optional*):855 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This856 is useful if you want more control over how to convert `input_ids` indices into associated vectors than the857 model's internal embedding lookup matrix.858 output_attentions (`bool`, *optional*):859 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned860 tensors for more detail.861 output_hidden_states (`bool`, *optional*):862 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for863 more detail.864 return_dict (`bool`, *optional*):865 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.866"""867 868 869@add_start_docstrings(870 "The bare Bert Model transformer outputting raw hidden-states without any specific head on top.",871 BERT_START_DOCSTRING,872)873class BertModel(BertPreTrainedModel):874 """875 876 The model can behave as an encoder (with only self-attention) as well as a decoder, in which case a layer of877 cross-attention is added between the self-attention layers, following the architecture described in [Attention is878 all you need](https://arxiv.org/abs/1706.03762) by Ashish Vaswani, Noam Shazeer, Niki Parmar, Jakob Uszkoreit,879 Llion Jones, Aidan N. Gomez, Lukasz Kaiser and Illia Polosukhin.880 881 To behave as an decoder the model needs to be initialized with the `is_decoder` argument of the configuration set882 to `True`. To be used in a Seq2Seq model, the model needs to initialized with both `is_decoder` argument and883 `add_cross_attention` set to `True`; an `encoder_hidden_states` is then expected as an input to the forward pass.884 """885 886 def __init__(self, config, add_pooling_layer=True):887 super().__init__(config)888 self.config = config889 890 self.embeddings = BertEmbeddings(config)891 self.encoder = BertEncoder(config)892 893 self.pooler = BertPooler(config) if add_pooling_layer else None894 895 # Initialize weights and apply final processing896 self.post_init()897 898 def get_input_embeddings(self):899 return self.embeddings.word_embeddings900 901 def set_input_embeddings(self, value):902 self.embeddings.word_embeddings = value903 904 def _prune_heads(self, heads_to_prune):905 """906 Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base907 class PreTrainedModel908 """909 for layer, heads in heads_to_prune.items():910 self.encoder.layer[layer].attention.prune_heads(heads)911 912 @add_start_docstrings_to_model_forward(BERT_INPUTS_DOCSTRING.format("batch_size, sequence_length"))913 @add_code_sample_docstrings(914 checkpoint=_CHECKPOINT_FOR_DOC,915 output_type=BaseModelOutputWithPoolingAndCrossAttentions,916 config_class=_CONFIG_FOR_DOC,917 )918 def forward(919 self,920 input_ids: Optional[torch.Tensor] = None,921 attention_mask: Optional[torch.Tensor] = None,922 token_type_ids: Optional[torch.Tensor] = None,923 position_ids: Optional[torch.Tensor] = None,924 head_mask: Optional[torch.Tensor] = None,925 inputs_embeds: Optional[torch.Tensor] = None,926 encoder_hidden_states: Optional[torch.Tensor] = None,927 encoder_attention_mask: Optional[torch.Tensor] = None,928 past_key_values: Optional[List[torch.FloatTensor]] = None,929 use_cache: Optional[bool] = None,930 output_attentions: Optional[bool] = None,931 output_hidden_states: Optional[bool] = None,932 return_dict: Optional[bool] = None,933 ) -> Union[Tuple[torch.Tensor], BaseModelOutputWithPoolingAndCrossAttentions]:934 r"""935 encoder_hidden_states (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):936 Sequence of hidden-states at the output of the last layer of the encoder. Used in the cross-attention if937 the model is configured as a decoder.938 encoder_attention_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length)`, *optional*):939 Mask to avoid performing attention on the padding token indices of the encoder input. This mask is used in940 the cross-attention if the model is configured as a decoder. Mask values selected in `[0, 1]`:941 942 - 1 for tokens that are **not masked**,943 - 0 for tokens that are **masked**.944 past_key_values (`tuple(tuple(torch.FloatTensor))` of length `config.n_layers` with each tuple having 4 tensors of shape `(batch_size, num_heads, sequence_length - 1, embed_size_per_head)`):945 Contains precomputed key and value hidden states of the attention blocks. Can be used to speed up decoding.946 947 If `past_key_values` are used, the user can optionally input only the last `decoder_input_ids` (those that948 don't have their past key value states given to this model) of shape `(batch_size, 1)` instead of all949 `decoder_input_ids` of shape `(batch_size, sequence_length)`.950 use_cache (`bool`, *optional*):951 If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see952 `past_key_values`).953 """954 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions955 output_hidden_states = (956 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states957 )958 return_dict = return_dict if return_dict is not None else self.config.use_return_dict959 960 if self.config.is_decoder:961 use_cache = use_cache if use_cache is not None else self.config.use_cache962 else:963 use_cache = False964 965 if input_ids is not None and inputs_embeds is not None:966 raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")967 elif input_ids is not None:968 input_shape = input_ids.size()969 elif inputs_embeds is not None:970 input_shape = inputs_embeds.size()[:-1]971 else:972 raise ValueError("You have to specify either input_ids or inputs_embeds")973 974 batch_size, seq_length = input_shape975 device = input_ids.device if input_ids is not None else inputs_embeds.device976 977 # past_key_values_length978 past_key_values_length = past_key_values[0][0].shape[2] if past_key_values is not None else 0979 980 if attention_mask is None:981 attention_mask = torch.ones(((batch_size, seq_length + past_key_values_length)), device=device)982 983 if token_type_ids is None:984 if hasattr(self.embeddings, "token_type_ids"):985 buffered_token_type_ids = self.embeddings.token_type_ids[:, :seq_length]986 buffered_token_type_ids_expanded = buffered_token_type_ids.expand(batch_size, seq_length)987 token_type_ids = buffered_token_type_ids_expanded988 else:989 token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=device)990 991 # We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]992 # ourselves in which case we just need to make it broadcastable to all heads.993 extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(attention_mask, input_shape)994 995 # If a 2D or 3D attention mask is provided for the cross-attention996 # we need to make broadcastable to [batch_size, num_heads, seq_length, seq_length]997 if self.config.is_decoder and encoder_hidden_states is not None:998 encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()999 encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)1000 if encoder_attention_mask is None:1001 encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)1002 encoder_extended_attention_mask = self.invert_attention_mask(encoder_attention_mask)1003 else:1004 encoder_extended_attention_mask = None1005 1006 # Prepare head mask if needed1007 # 1.0 in head_mask indicate we keep the head1008 # attention_probs has shape bsz x n_heads x N x N1009 # input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]1010 # and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]1011 head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)1012 1013 embedding_output = self.embeddings(1014 input_ids=input_ids,1015 position_ids=position_ids,1016 token_type_ids=token_type_ids,1017 inputs_embeds=inputs_embeds,1018 past_key_values_length=past_key_values_length,1019 )1020 encoder_outputs = self.encoder(1021 embedding_output,1022 attention_mask=extended_attention_mask,1023 head_mask=head_mask,1024 encoder_hidden_states=encoder_hidden_states,1025 encoder_attention_mask=encoder_extended_attention_mask,1026 past_key_values=past_key_values,1027 use_cache=use_cache,1028 output_attentions=output_attentions,1029 output_hidden_states=output_hidden_states,1030 return_dict=return_dict,1031 )1032 sequence_output = encoder_outputs[0]1033 pooled_output = self.pooler(sequence_output) if self.pooler is not None else None1034 1035 if not return_dict:1036 return (sequence_output, pooled_output) + encoder_outputs[1:]1037 1038 return BaseModelOutputWithPoolingAndCrossAttentions(1039 last_hidden_state=sequence_output,1040 pooler_output=pooled_output,1041 past_key_values=encoder_outputs.past_key_values,1042 hidden_states=encoder_outputs.hidden_states,1043 attentions=encoder_outputs.attentions,1044 cross_attentions=encoder_outputs.cross_attentions,1045 )1046 1047 1048@add_start_docstrings(1049 """1050 Bert Model with two heads on top as done during the pretraining: a `masked language modeling` head and a `next1051 sentence prediction (classification)` head.1052 """,1053 BERT_START_DOCSTRING,1054)1055class BertForPreTraining(BertPreTrainedModel):1056 _keys_to_ignore_on_load_missing = [r"position_ids", r"predictions.decoder.bias", r"cls.predictions.decoder.weight"]1057 1058 def __init__(self, config):1059 super().__init__(config)1060 1061 self.bert = BertModel(config)1062 self.cls = BertPreTrainingHeads(config)1063 1064 # Initialize weights and apply final processing1065 self.post_init()1066 1067 def get_output_embeddings(self):1068 return self.cls.predictions.decoder1069 1070 def set_output_embeddings(self, new_embeddings):1071 self.cls.predictions.decoder = new_embeddings1072 1073 @add_start_docstrings_to_model_forward(BERT_INPUTS_DOCSTRING.format("batch_size, sequence_length"))1074 @replace_return_docstrings(output_type=BertForPreTrainingOutput, config_class=_CONFIG_FOR_DOC)1075 def forward(1076 self,1077 input_ids: Optional[torch.Tensor] = None,1078 attention_mask: Optional[torch.Tensor] = None,1079 token_type_ids: Optional[torch.Tensor] = None,1080 position_ids: Optional[torch.Tensor] = None,1081 head_mask: Optional[torch.Tensor] = None,1082 inputs_embeds: Optional[torch.Tensor] = None,1083 labels: Optional[torch.Tensor] = None,1084 next_sentence_label: Optional[torch.Tensor] = None,1085 output_attentions: Optional[bool] = None,1086 output_hidden_states: Optional[bool] = None,1087 return_dict: Optional[bool] = None,1088 ) -> Union[Tuple[torch.Tensor], BertForPreTrainingOutput]:1089 r"""1090 labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1091 Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,1092 config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are ignored (masked),1093 the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`1094 next_sentence_label (`torch.LongTensor` of shape `(batch_size,)`, *optional*):1095 Labels for computing the next sequence prediction (classification) loss. Input should be a sequence1096 pair (see `input_ids` docstring) Indices should be in `[0, 1]`:1097 1098 - 0 indicates sequence B is a continuation of sequence A,1099 - 1 indicates sequence B is a random sequence.1100 kwargs (`Dict[str, any]`, optional, defaults to *{}*):1101 Used to hide legacy arguments that have been deprecated.1102 1103 Returns:1104 1105 Example:1106 1107 ```python1108 >>> from transformers import AutoTokenizer, BertForPreTraining1109 >>> import torch1110 1111 >>> tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")1112 >>> model = BertForPreTraining.from_pretrained("bert-base-uncased")1113 1114 >>> inputs = tokenizer("Hello, my dog is cute", return_tensors="pt")1115 >>> outputs = model(**inputs)1116 1117 >>> prediction_logits = outputs.prediction_logits1118 >>> seq_relationship_logits = outputs.seq_relationship_logits1119 ```1120 """1121 return_dict = return_dict if return_dict is not None else self.config.use_return_dict1122 1123 outputs = self.bert(1124 input_ids,1125 attention_mask=attention_mask,1126 token_type_ids=token_type_ids,1127 position_ids=position_ids,1128 head_mask=head_mask,1129 inputs_embeds=inputs_embeds,1130 output_attentions=output_attentions,1131 output_hidden_states=output_hidden_states,1132 return_dict=return_dict,1133 )1134 1135 sequence_output, pooled_output = outputs[:2]1136 prediction_scores, seq_relationship_score = self.cls(sequence_output, pooled_output)1137 1138 total_loss = None1139 if labels is not None and next_sentence_label is not None:1140 loss_fct = CrossEntropyLoss()1141 masked_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), labels.view(-1))1142 next_sentence_loss = loss_fct(seq_relationship_score.view(-1, 2), next_sentence_label.view(-1))1143 total_loss = masked_lm_loss + next_sentence_loss1144 1145 if not return_dict:1146 output = (prediction_scores, seq_relationship_score) + outputs[2:]1147 return ((total_loss,) + output) if total_loss is not None else output1148 1149 return BertForPreTrainingOutput(1150 loss=total_loss,1151 prediction_logits=prediction_scores,1152 seq_relationship_logits=seq_relationship_score,1153 hidden_states=outputs.hidden_states,1154 attentions=outputs.attentions,1155 )1156 1157 1158@add_start_docstrings(1159 """Bert Model with a `language modeling` head on top for CLM fine-tuning.""", BERT_START_DOCSTRING1160)1161class BertCustomLMHeadModel(BertPreTrainedModel):1162 _keys_to_ignore_on_load_unexpected = [r"pooler"]1163 _keys_to_ignore_on_load_missing = [r"position_ids", r"predictions.decoder.bias", r"cls.predictions.decoder.weight"]1164 1165 def __init__(self, config):1166 super().__init__(config)1167 1168 if not config.is_decoder:1169 logger.warning("If you want to use `BertLMHeadModel` as a standalone, add `is_decoder=True.`")1170 1171 self.bert = BertModel(config, add_pooling_layer=False)1172 self.cls = BertOnlyMLMHead(config)1173 1174 # Initialize weights and apply final processing1175 self.post_init()1176 1177 def get_output_embeddings(self):1178 return self.cls.predictions.decoder1179 1180 def set_output_embeddings(self, new_embeddings):1181 self.cls.predictions.decoder = new_embeddings1182 1183 @add_start_docstrings_to_model_forward(BERT_INPUTS_DOCSTRING.format("batch_size, sequence_length"))1184 @add_code_sample_docstrings(1185 checkpoint=_CHECKPOINT_FOR_DOC,1186 output_type=CausalLMOutputWithCrossAttentions,1187 config_class=_CONFIG_FOR_DOC,1188 )1189 def forward(1190 self,1191 input_ids: Optional[torch.Tensor] = None,1192 attention_mask: Optional[torch.Tensor] = None,1193 token_type_ids: Optional[torch.Tensor] = None,1194 position_ids: Optional[torch.Tensor] = None,1195 head_mask: Optional[torch.Tensor] = None,1196 inputs_embeds: Optional[torch.Tensor] = None,1197 encoder_hidden_states: Optional[torch.Tensor] = None,1198 encoder_attention_mask: Optional[torch.Tensor] = None,1199 labels: Optional[torch.Tensor] = None,1200 past_key_values: Optional[List[torch.Tensor]] = None,