CoolFace
Modelpublic

apps1/draft_overall_student

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes6downloads
modeling_bert_hash.py520 linesDownload Raw Back to root
1from typing import Optional, Union2 3import torch4from torch import nn5from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss6 7from transformers.cache_utils import Cache8from transformers.models.bert.modeling_bert import BertEncoder, BertPooler, BertPreTrainedModel, BertOnlyMLMHead9from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask_for_sdpa, _prepare_4d_causal_attention_mask_for_sdpa10from transformers.modeling_outputs import (11    BaseModelOutputWithPoolingAndCrossAttentions,12    MaskedLMOutput,13    SequenceClassifierOutput,14)15from transformers.utils import auto_docstring, logging16 17from .configuration_bert_hash import BertHashConfig18 19logger = logging.get_logger(__name__)20 21 22class BertHashTokens(nn.Module):23    """24    Module that embeds token vocabulary to an intermediate embeddings layer then projects those embeddings to the25    hidden size.26 27    The number of projections is like a hash. Setting the projections parameter to 5 is like generating a28    160-bit hash (5 x float32) for each token. That hash is then projected to the hidden size.29 30    This significantly reduces the number of parameters necessary for token embeddings.31 32    For example:33        Standard token embeddings:34          30,522 (vocab size) x 768 (hidden size) = 23,440,896 parameters35          23,440,896 x 4 (float32) = 93,763,584 bytes36 37        Hash token embeddings:38          30,522 (vocab size) x 5 (hash buckets) + 5 x 768 (projection matrix)= 156,450 parameters39          156,450 x 4 (float32) = 625,800 bytes    40    """41 42    def __init__(self, config):43        super().__init__()44        self.config = config45 46        # Token embeddings47        self.embeddings = nn.Embedding(config.vocab_size, config.projections, padding_idx=config.pad_token_id)48 49        # Token embeddings projections50        self.projections = nn.Linear(config.projections, config.hidden_size)51 52    def forward(self, input_ids):53        # Project embeddings to hidden size54        return self.projections(self.embeddings(input_ids))55 56 57class BertHashEmbeddings(nn.Module):58    """Construct the embeddings from word, position and token_type embeddings."""59 60    def __init__(self, config):61        super().__init__()62        self.word_embeddings = BertHashTokens(config)63        self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size)64        self.token_type_embeddings = nn.Embedding(config.type_vocab_size, config.hidden_size)65 66        # self.LayerNorm is not snake-cased to stick with TensorFlow model variable name and be able to load67        # any TensorFlow checkpoint file68        self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)69        self.dropout = nn.Dropout(config.hidden_dropout_prob)70        # position_ids (1, len position emb) is contiguous in memory and exported when serialized71        self.position_embedding_type = getattr(config, "position_embedding_type", "absolute")72        self.register_buffer(73            "position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False74        )75        self.register_buffer(76            "token_type_ids", torch.zeros(self.position_ids.size(), dtype=torch.long), persistent=False77        )78 79    def forward(80        self,81        input_ids: Optional[torch.LongTensor] = None,82        token_type_ids: Optional[torch.LongTensor] = None,83        position_ids: Optional[torch.LongTensor] = None,84        inputs_embeds: Optional[torch.FloatTensor] = None,85        past_key_values_length: int = 0,86    ) -> torch.Tensor:87        if input_ids is not None:88            input_shape = input_ids.size()89        else:90            input_shape = inputs_embeds.size()[:-1]91 92        seq_length = input_shape[1]93 94        if position_ids is None:95            position_ids = self.position_ids[:, past_key_values_length : seq_length + past_key_values_length]96 97        # Setting the token_type_ids to the registered buffer in constructor where it is all zeros, which usually occurs98        # when its auto-generated, registered buffer helps users when tracing the model without passing token_type_ids, solves99        # issue #5664100        if token_type_ids is None:101            if hasattr(self, "token_type_ids"):102                buffered_token_type_ids = self.token_type_ids[:, :seq_length]103                buffered_token_type_ids_expanded = buffered_token_type_ids.expand(input_shape[0], seq_length)104                token_type_ids = buffered_token_type_ids_expanded105            else:106                token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=self.position_ids.device)107 108        if inputs_embeds is None:109            inputs_embeds = self.word_embeddings(input_ids)110        token_type_embeddings = self.token_type_embeddings(token_type_ids)111 112        embeddings = inputs_embeds + token_type_embeddings113        if self.position_embedding_type == "absolute":114            position_embeddings = self.position_embeddings(position_ids)115            embeddings += position_embeddings116        embeddings = self.LayerNorm(embeddings)117        embeddings = self.dropout(embeddings)118        return embeddings119 120 121@auto_docstring(122    custom_intro="""123    The model can behave as an encoder (with only self-attention) as well as a decoder, in which case a layer of124    cross-attention is added between the self-attention layers, following the architecture described in [Attention is125    all you need](https://huggingface.co/papers/1706.03762) by Ashish Vaswani, Noam Shazeer, Niki Parmar, Jakob Uszkoreit,126    Llion Jones, Aidan N. Gomez, Lukasz Kaiser and Illia Polosukhin.127 128    To behave as an decoder the model needs to be initialized with the `is_decoder` argument of the configuration set129    to `True`. To be used in a Seq2Seq model, the model needs to initialized with both `is_decoder` argument and130    `add_cross_attention` set to `True`; an `encoder_hidden_states` is then expected as an input to the forward pass.131    """132)133class BertHashModel(BertPreTrainedModel):134    config_class = BertHashConfig135 136    _no_split_modules = ["BertEmbeddings", "BertLayer"]137 138    def __init__(self, config, add_pooling_layer=True):139        r"""140        add_pooling_layer (bool, *optional*, defaults to `True`):141            Whether to add a pooling layer142        """143        super().__init__(config)144        self.config = config145 146        self.embeddings = BertHashEmbeddings(config)147        self.encoder = BertEncoder(config)148 149        self.pooler = BertPooler(config) if add_pooling_layer else None150 151        self.attn_implementation = config._attn_implementation152        self.position_embedding_type = config.position_embedding_type153 154        # Initialize weights and apply final processing155        self.post_init()156 157    def get_input_embeddings(self):158        return self.embeddings.word_embeddings.embeddings159 160    def set_input_embeddings(self, value):161        self.embeddings.word_embeddings.embeddings = value162 163    def _prune_heads(self, heads_to_prune):164        """165        Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base166        class PreTrainedModel167        """168        for layer, heads in heads_to_prune.items():169            self.encoder.layer[layer].attention.prune_heads(heads)170 171    @auto_docstring172    def forward(173        self,174        input_ids: Optional[torch.Tensor] = None,175        attention_mask: Optional[torch.Tensor] = None,176        token_type_ids: Optional[torch.Tensor] = None,177        position_ids: Optional[torch.Tensor] = None,178        head_mask: Optional[torch.Tensor] = None,179        inputs_embeds: Optional[torch.Tensor] = None,180        encoder_hidden_states: Optional[torch.Tensor] = None,181        encoder_attention_mask: Optional[torch.Tensor] = None,182        past_key_values: Optional[list[torch.FloatTensor]] = None,183        use_cache: Optional[bool] = None,184        output_attentions: Optional[bool] = None,185        output_hidden_states: Optional[bool] = None,186        return_dict: Optional[bool] = None,187        cache_position: Optional[torch.Tensor] = None,188    ) -> Union[tuple[torch.Tensor], BaseModelOutputWithPoolingAndCrossAttentions]:189        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions190        output_hidden_states = (191            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states192        )193        return_dict = return_dict if return_dict is not None else self.config.use_return_dict194 195        if self.config.is_decoder:196            use_cache = use_cache if use_cache is not None else self.config.use_cache197        else:198            use_cache = False199 200        if input_ids is not None and inputs_embeds is not None:201            raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")202        elif input_ids is not None:203            self.warn_if_padding_and_no_attention_mask(input_ids, attention_mask)204            input_shape = input_ids.size()205        elif inputs_embeds is not None:206            input_shape = inputs_embeds.size()[:-1]207        else:208            raise ValueError("You have to specify either input_ids or inputs_embeds")209 210        batch_size, seq_length = input_shape211        device = input_ids.device if input_ids is not None else inputs_embeds.device212 213        past_key_values_length = 0214        if past_key_values is not None:215            past_key_values_length = (216                past_key_values[0][0].shape[-2]217                if not isinstance(past_key_values, Cache)218                else past_key_values.get_seq_length()219            )220 221        if token_type_ids is None:222            if hasattr(self.embeddings, "token_type_ids"):223                buffered_token_type_ids = self.embeddings.token_type_ids[:, :seq_length]224                buffered_token_type_ids_expanded = buffered_token_type_ids.expand(batch_size, seq_length)225                token_type_ids = buffered_token_type_ids_expanded226            else:227                token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=device)228 229        embedding_output = self.embeddings(230            input_ids=input_ids,231            position_ids=position_ids,232            token_type_ids=token_type_ids,233            inputs_embeds=inputs_embeds,234            past_key_values_length=past_key_values_length,235        )236 237        if attention_mask is None:238            attention_mask = torch.ones((batch_size, seq_length + past_key_values_length), device=device)239 240        use_sdpa_attention_masks = (241            self.attn_implementation == "sdpa"242            and self.position_embedding_type == "absolute"243            and head_mask is None244            and not output_attentions245        )246 247        # Expand the attention mask248        if use_sdpa_attention_masks and attention_mask.dim() == 2:249            # Expand the attention mask for SDPA.250            # [bsz, seq_len] -> [bsz, 1, seq_len, seq_len]251            if self.config.is_decoder:252                extended_attention_mask = _prepare_4d_causal_attention_mask_for_sdpa(253                    attention_mask,254                    input_shape,255                    embedding_output,256                    past_key_values_length,257                )258            else:259                extended_attention_mask = _prepare_4d_attention_mask_for_sdpa(260                    attention_mask, embedding_output.dtype, tgt_len=seq_length261                )262        else:263            # We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]264            # ourselves in which case we just need to make it broadcastable to all heads.265            extended_attention_mask = self.get_extended_attention_mask(attention_mask, input_shape)266 267        # If a 2D or 3D attention mask is provided for the cross-attention268        # we need to make broadcastable to [batch_size, num_heads, seq_length, seq_length]269        if self.config.is_decoder and encoder_hidden_states is not None:270            encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()271            encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)272            if encoder_attention_mask is None:273                encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)274 275            if use_sdpa_attention_masks and encoder_attention_mask.dim() == 2:276                # Expand the attention mask for SDPA.277                # [bsz, seq_len] -> [bsz, 1, seq_len, seq_len]278                encoder_extended_attention_mask = _prepare_4d_attention_mask_for_sdpa(279                    encoder_attention_mask, embedding_output.dtype, tgt_len=seq_length280                )281            else:282                encoder_extended_attention_mask = self.invert_attention_mask(encoder_attention_mask)283        else:284            encoder_extended_attention_mask = None285 286        # Prepare head mask if needed287        # 1.0 in head_mask indicate we keep the head288        # attention_probs has shape bsz x n_heads x N x N289        # input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]290        # and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]291        head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)292 293        encoder_outputs = self.encoder(294            embedding_output,295            attention_mask=extended_attention_mask,296            head_mask=head_mask,297            encoder_hidden_states=encoder_hidden_states,298            encoder_attention_mask=encoder_extended_attention_mask,299            past_key_values=past_key_values,300            use_cache=use_cache,301            output_attentions=output_attentions,302            output_hidden_states=output_hidden_states,303            return_dict=return_dict,304            cache_position=cache_position,305        )306        sequence_output = encoder_outputs[0]307        pooled_output = self.pooler(sequence_output) if self.pooler is not None else None308 309        if not return_dict:310            return (sequence_output, pooled_output) + encoder_outputs[1:]311 312        return BaseModelOutputWithPoolingAndCrossAttentions(313            last_hidden_state=sequence_output,314            pooler_output=pooled_output,315            past_key_values=encoder_outputs.past_key_values,316            hidden_states=encoder_outputs.hidden_states,317            attentions=encoder_outputs.attentions,318            cross_attentions=encoder_outputs.cross_attentions,319        )320 321 322@auto_docstring323class BertHashForMaskedLM(BertPreTrainedModel):324    _tied_weights_keys = ["predictions.decoder.bias", "cls.predictions.decoder.weight"]325    config_class = BertHashConfig326 327    def __init__(self, config):328        super().__init__(config)329 330        if config.is_decoder:331            logger.warning(332                "If you want to use `BertForMaskedLM` make sure `config.is_decoder=False` for "333                "bi-directional self-attention."334            )335 336        self.bert = BertHashModel(config, add_pooling_layer=False)337        self.cls = BertOnlyMLMHead(config)338 339        # Initialize weights and apply final processing340        self.post_init()341 342    @auto_docstring343    def forward(344        self,345        input_ids: Optional[torch.Tensor] = None,346        attention_mask: Optional[torch.Tensor] = None,347        token_type_ids: Optional[torch.Tensor] = None,348        position_ids: Optional[torch.Tensor] = None,349        head_mask: Optional[torch.Tensor] = None,350        inputs_embeds: Optional[torch.Tensor] = None,351        encoder_hidden_states: Optional[torch.Tensor] = None,352        encoder_attention_mask: Optional[torch.Tensor] = None,353        labels: Optional[torch.Tensor] = None,354        output_attentions: Optional[bool] = None,355        output_hidden_states: Optional[bool] = None,356        return_dict: Optional[bool] = None,357    ) -> Union[tuple[torch.Tensor], MaskedLMOutput]:358        r"""359        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):360            Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,361            config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are ignored (masked), the362            loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`363        """364 365        return_dict = return_dict if return_dict is not None else self.config.use_return_dict366 367        outputs = self.bert(368            input_ids,369            attention_mask=attention_mask,370            token_type_ids=token_type_ids,371            position_ids=position_ids,372            head_mask=head_mask,373            inputs_embeds=inputs_embeds,374            encoder_hidden_states=encoder_hidden_states,375            encoder_attention_mask=encoder_attention_mask,376            output_attentions=output_attentions,377            output_hidden_states=output_hidden_states,378            return_dict=return_dict,379        )380 381        sequence_output = outputs[0]382        prediction_scores = self.cls(sequence_output)383 384        masked_lm_loss = None385        if labels is not None:386            loss_fct = CrossEntropyLoss()  # -100 index = padding token387            masked_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), labels.view(-1))388 389        if not return_dict:390            output = (prediction_scores,) + outputs[2:]391            return ((masked_lm_loss,) + output) if masked_lm_loss is not None else output392 393        return MaskedLMOutput(394            loss=masked_lm_loss,395            logits=prediction_scores,396            hidden_states=outputs.hidden_states,397            attentions=outputs.attentions,398        )399 400    def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **model_kwargs):401        input_shape = input_ids.shape402        effective_batch_size = input_shape[0]403 404        #  add a dummy token405        if self.config.pad_token_id is None:406            raise ValueError("The PAD token should be defined for generation")407 408        attention_mask = torch.cat([attention_mask, attention_mask.new_zeros((attention_mask.shape[0], 1))], dim=-1)409        dummy_token = torch.full(410            (effective_batch_size, 1), self.config.pad_token_id, dtype=torch.long, device=input_ids.device411        )412        input_ids = torch.cat([input_ids, dummy_token], dim=1)413 414        return {"input_ids": input_ids, "attention_mask": attention_mask}415 416    @classmethod417    def can_generate(cls) -> bool:418        """419        Legacy correction: BertForMaskedLM can't call `generate()` from `GenerationMixin`, even though it has a420        `prepare_inputs_for_generation` method.421        """422        return False423 424 425@auto_docstring(426    custom_intro="""427    Bert Model transformer with a sequence classification/regression head on top (a linear layer on top of the pooled428    output) e.g. for GLUE tasks.429    """430)431class BertHashForSequenceClassification(BertPreTrainedModel):432    config_class = BertHashConfig433 434    def __init__(self, config):435        super().__init__(config)436        self.num_labels = config.num_labels437        self.config = config438 439        self.bert = BertHashModel(config)440        classifier_dropout = (441            config.classifier_dropout if config.classifier_dropout is not None else config.hidden_dropout_prob442        )443        self.dropout = nn.Dropout(classifier_dropout)444        self.classifier = nn.Linear(config.hidden_size, config.num_labels)445 446        # Initialize weights and apply final processing447        self.post_init()448 449    @auto_docstring450    def forward(451        self,452        input_ids: Optional[torch.Tensor] = None,453        attention_mask: Optional[torch.Tensor] = None,454        token_type_ids: Optional[torch.Tensor] = None,455        position_ids: Optional[torch.Tensor] = None,456        head_mask: Optional[torch.Tensor] = None,457        inputs_embeds: Optional[torch.Tensor] = None,458        labels: Optional[torch.Tensor] = None,459        output_attentions: Optional[bool] = None,460        output_hidden_states: Optional[bool] = None,461        return_dict: Optional[bool] = None,462    ) -> Union[tuple[torch.Tensor], SequenceClassifierOutput]:463        r"""464        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):465            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,466            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If467            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).468        """469        return_dict = return_dict if return_dict is not None else self.config.use_return_dict470 471        outputs = self.bert(472            input_ids,473            attention_mask=attention_mask,474            token_type_ids=token_type_ids,475            position_ids=position_ids,476            head_mask=head_mask,477            inputs_embeds=inputs_embeds,478            output_attentions=output_attentions,479            output_hidden_states=output_hidden_states,480            return_dict=return_dict,481        )482 483        pooled_output = outputs[1]484 485        pooled_output = self.dropout(pooled_output)486        logits = self.classifier(pooled_output)487 488        loss = None489        if labels is not None:490            if self.config.problem_type is None:491                if self.num_labels == 1:492                    self.config.problem_type = "regression"493                elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int):494                    self.config.problem_type = "single_label_classification"495                else:496                    self.config.problem_type = "multi_label_classification"497 498            if self.config.problem_type == "regression":499                loss_fct = MSELoss()500                if self.num_labels == 1:501                    loss = loss_fct(logits.squeeze(), labels.squeeze())502                else:503                    loss = loss_fct(logits, labels)504            elif self.config.problem_type == "single_label_classification":505                loss_fct = CrossEntropyLoss()506                loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))507            elif self.config.problem_type == "multi_label_classification":508                loss_fct = BCEWithLogitsLoss()509                loss = loss_fct(logits, labels)510        if not return_dict:511            output = (logits,) + outputs[2:]512            return ((loss,) + output) if loss is not None else output513 514        return SequenceClassifierOutput(515            loss=loss,516            logits=logits,517            hidden_states=outputs.hidden_states,518            attentions=outputs.attentions,519        )520