Team Ai
Modelpublic

Jingya/tiny-random-bert-remote-code

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes157downloads
modeling_bert.py1894 linesDownload Raw Back to root
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,

Showing the first 1,200 of 1894 lines. Download the file for the rest.