Team Ai
Apppublic

bigscience/petals-api

sourceHugging Faceupdated 4y agoView on Hugging Face
18likes
model.py409 linesDownload Raw Back to bloom
1"""2PyTorch BLOOM model that implements several memory-efficient modes.3Based on https://github.com/huggingface/transformers/commit/ca2a55e9dfb245527b5e1c954fec6ffbb7aef07b4See commit history for authorship.5"""6from typing import Tuple7 8import torch9import torch.nn.functional as F10import torch.utils.checkpoint11from hivemind import use_hivemind_log_handler12from torch import nn13from torch.nn import CrossEntropyLoss, LayerNorm14from transformers.file_utils import (add_code_sample_docstrings, add_start_docstrings,15                                     add_start_docstrings_to_model_forward)16from transformers.modeling_outputs import BaseModelOutputWithPastAndCrossAttentions, CausalLMOutputWithCrossAttentions17from transformers.modeling_utils import PreTrainedModel18from transformers.models.bloom.configuration_bloom import BloomConfig19from transformers.utils import logging20 21from src.bloom.block import BloomBlock22 23use_hivemind_log_handler("in_root_logger")24logger = logging.get_logger(__file__)25 26_CHECKPOINT_FOR_DOC = "bigscience/Bloom"27_CONFIG_FOR_DOC = "BloomConfig"28_TOKENIZER_FOR_DOC = "BloomTokenizer"29 30 31class BloomPreTrainedModel(PreTrainedModel):32    _keys_to_ignore_on_load_missing = [r"h.*.self_attention.scale_mask_softmax.causal_mask", r"lm_head.weight"]33    """34    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained35    models.36    """37 38    config_class = BloomConfig39    base_model_prefix = "transformer"40    supports_gradient_checkpointing = True41    _no_split_modules = ["BloomBlock"]42 43    def __init__(self, *inputs, **kwargs):44        super().__init__(*inputs, **kwargs)45 46    def _init_weights(self, module):47        """Initialize the weights."""48        if isinstance(module, (nn.Linear)):49            # Slightly different from the TF version which uses truncated_normal for initialization50            # cf https://github.com/pytorch/pytorch/pull/561751            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)52            if module.bias is not None:53                module.bias.data.zero_()54        elif isinstance(module, nn.Embedding):55            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)56            if module.padding_idx is not None:57                module.weight.data[module.padding_idx].zero_()58        elif isinstance(module, LayerNorm):59            module.bias.data.zero_()60            module.weight.data.fill_(1.0)61 62    def _set_gradient_checkpointing(self, module, value=False):63        if isinstance(module, BloomModel):64            module.gradient_checkpointing = value65 66 67BLOOM_START_DOCSTRING = r"""68 69    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the70    library implements for all its model (such as downloading or saving, resizing the input embeddings etc.)71 72    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.73    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage74    and behavior.75 76    Parameters:77        config ([`MemoryEfficientBloomConfig`]): Model configuration class with all the parameters of the model.78            Initializing with a config file does not load the weights associated with the model, only the79            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.80"""81 82BLOOM_INPUTS_DOCSTRING = r"""83    Args:84        input_ids (`torch.LongTensor` of shape `(batch_size, input_ids_length)`):85            `input_ids_length` = `sequence_length` if `past_key_values` is `None` else86            `past_key_values[0][0].shape[-2]` (`sequence_length` of input past key value states). Indices of input87            sequence tokens in the vocabulary.88 89            If `past_key_values` is used, only `input_ids` that do not have their past calculated should be passed as90            `input_ids`.91 92            Indices can be obtained using [`BloomTokenizer`]. See [`PreTrainedTokenizer.encode`] and93            [`PreTrainedTokenizer.__call__`] for details.94 95            [What are input IDs?](../glossary#input-ids)96        past_key_values (`Tuple[Tuple[torch.Tensor]]` of length `config.n_layers`):97            Contains precomputed hidden-states (key and values in the attention blocks) as computed by the model (see98            `past_key_values` output below). Can be used to speed up sequential decoding. The `input_ids` which have99            their past given to this model should not be passed as `input_ids` as they have already been computed.100        attention_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length)`, *optional*):101            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:102 103            - 1 for tokens that are **not masked**,104            - 0 for tokens that are **masked**.105 106            [What are attention masks?](../glossary#attention-mask)107        position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):108            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,109            config.max_position_embeddings - 1]`.110 111            [What are position IDs?](../glossary#position-ids)112        head_mask (`torch.FloatTensor` of shape `(num_heads,)` or `(num_layers, num_heads)`, *optional*):113            Mask to nullify selected heads of the self-attention modules. Mask values selected in `[0, 1]`:114 115            - 1 indicates the head is **not masked**,116            - 0 indicates the head is **masked**.117 118        inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):119            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This120            is useful if you want more control over how to convert `input_ids` indices into associated vectors than the121            model's internal embedding lookup matrix.122 123            If `past_key_values` is used, optionally only the last `inputs_embeds` have to be input (see124            `past_key_values`).125        use_cache (`bool`, *optional*):126            If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see127            `past_key_values`).128        output_attentions (`bool`, *optional*):129            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned130            tensors for more detail.131        output_hidden_states (`bool`, *optional*):132            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for133            more detail.134        return_dict (`bool`, *optional*):135            Whether or not to return a [`~file_utils.ModelOutput`] instead of a plain tuple.136"""137 138 139@add_start_docstrings(140    "The bare Bloom Model transformer outputting raw hidden-states without any specific head on top.",141    BLOOM_START_DOCSTRING,142)143class BloomModel(BloomPreTrainedModel):144    def __init__(self, config):145        super().__init__(config)146        assert not config.slow_but_exact, "slow_but_exact mode was removed for code simplicity"147 148        self.embed_dim = config.hidden_size149        self.n_head = config.n_head150 151        # Embedding + LN Embedding152 153        # TODO: @dbaranchuk make efficient fp16 on cpu (convert only word_embeddings!)154        self.word_embeddings = nn.Embedding(config.vocab_size, self.embed_dim)  # dtype=config.torch_dtype155        self.word_embeddings_layernorm = LayerNorm(self.embed_dim, eps=config.layer_norm_epsilon)156 157        # Transformer blocks158        self.h = nn.ModuleList([BloomBlock(config, layer_number=i) for i in range(config.num_hidden_layers)])159 160        # Final Layer Norm161        self.ln_f = LayerNorm(self.embed_dim, eps=config.layer_norm_epsilon)162 163        self.gradient_checkpointing = False164 165        # Initialize weights and apply final processing166        self.post_init()167 168        # Forbid accumulate grads for embeddings and layernorm169        self.set_requires_grad(False)170 171    def get_input_embeddings(self):172        return self.word_embeddings173 174    def set_input_embeddings(self, new_embeddings):175        self.word_embeddings = new_embeddings176 177    def set_requires_grad(self, value):178        for p in self.parameters():179            p.requires_grad = value180 181    @add_start_docstrings_to_model_forward(BLOOM_INPUTS_DOCSTRING)182    @add_code_sample_docstrings(183        processor_class=_TOKENIZER_FOR_DOC,184        checkpoint=_CHECKPOINT_FOR_DOC,185        output_type=BaseModelOutputWithPastAndCrossAttentions,186        config_class=_CONFIG_FOR_DOC,187    )188    def forward(189        self,190        input_ids=None,191        past_key_values=None,192        attention_mask=None,193        position_ids=None,194        head_mask=None,195        inputs_embeds=None,196        use_cache=None,197        output_attentions=None,198        output_hidden_states=None,199        return_dict=None,200    ):201        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions202        output_hidden_states = (203            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states204        )205        use_cache = use_cache if use_cache is not None else self.config.use_cache206        return_dict = return_dict if return_dict is not None else self.config.use_return_dict207 208        if input_ids is not None and inputs_embeds is not None:209            raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")210        if position_ids is not None:211            logger.warning("position_ids are ignored in this bloom implementation")212        elif input_ids is not None:213            input_shape = input_ids.size()214            input_ids = input_ids.view(-1, input_shape[-1])215        elif inputs_embeds is not None:216            input_shape = inputs_embeds.size()[:-1]217        else:218            raise ValueError("You have to specify either input_ids or inputs_embeds")219 220        if past_key_values is None:221            past_key_values = tuple([None] * len(self.h))222 223        # Prepare head mask if needed224        # 1.0 in head_mask indicate we keep the head225        # attention_probs has shape bsz x n_head x N x N226        # head_mask has shape n_layer x batch x n_head x N x N227        head_mask = self.get_head_mask(head_mask, self.config.n_layer)228 229        if inputs_embeds is None:230            inputs_embeds = self.word_embeddings(input_ids)231 232        hidden_states = self.word_embeddings_layernorm(inputs_embeds.float())233 234        output_shape = input_shape + (hidden_states.size(-1),)235 236        presents = () if use_cache else None237        all_self_attentions = () if output_attentions else None238        all_hidden_states = () if output_hidden_states else None239 240        # Compute alibi tensor: check build_alibi_tensor documentation241        current_sequence_length = hidden_states.shape[1]242        if past_key_values and past_key_values[0]:243            current_sequence_length += past_key_values[0][0].shape[1]244 245        for i, (block, layer_past) in enumerate(zip(self.h, past_key_values)):246 247            if output_hidden_states:248                all_hidden_states = all_hidden_states + (hidden_states,)249 250            if self.gradient_checkpointing and self.training:251 252                if use_cache:253                    logger.warning(254                        "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."255                    )256                    use_cache = False257 258                def create_custom_forward(module):259                    def custom_forward(*inputs):260                        # None for past_key_value261                        return module(*inputs, use_cache, output_attentions, alibi=None)262 263                    return custom_forward264 265                outputs = torch.utils.checkpoint.checkpoint(266                    create_custom_forward(block),267                    hidden_states,268                    None,269                    attention_mask,270                    head_mask[i],271                )272            else:273                outputs = block(274                    hidden_states,275                    layer_past=layer_past,276                    attention_mask=attention_mask,277                    head_mask=head_mask[i],278                    use_cache=use_cache,279                    output_attentions=output_attentions,280                    alibi=None,281                )282 283            hidden_states = outputs[0]284            if use_cache is True:285                presents = presents + (outputs[1],)286 287            if output_attentions:288                all_self_attentions = all_self_attentions + (outputs[2 if use_cache else 1],)289 290        # Add last hidden state291        hidden_states = self.ln_f(hidden_states)292 293        if output_hidden_states:294            all_hidden_states = all_hidden_states + (hidden_states,)295 296        hidden_states = hidden_states.view(output_shape)297 298        if not return_dict:299            return tuple(v for v in [hidden_states, presents, all_hidden_states, all_self_attentions] if v is not None)300 301        return BaseModelOutputWithPastAndCrossAttentions(302            last_hidden_state=hidden_states,303            past_key_values=presents,304            hidden_states=all_hidden_states,305            attentions=all_self_attentions,306        )307 308 309@add_start_docstrings(310    """311    The Bloom Model transformer with a language modeling head on top (linear layer with weights tied to the input312    embeddings).313    """,314    BLOOM_START_DOCSTRING,315)316class BloomForCausalLM(BloomPreTrainedModel):317    _keys_to_ignore_on_load_missing = [r"h.*.self_attention.scale_mask_softmax.causal_mask", r"lm_head.weight"]318 319    def __init__(self, config):320        super().__init__(config)321        self.transformer = BloomModel(config)322        # Initialize weights and apply final processing323        self.post_init()324 325    def get_output_embeddings(self):326        return self.transformer.word_embeddings327 328    def set_output_embeddings(self, new_embeddings):329        self.transformer.word_embeddings.weight = new_embeddings.weight330 331    def prepare_inputs_for_generation(self, input_ids, past=None, **kwargs):332        # only last token for inputs_ids if past is defined in kwargs333        if past:334            input_ids = input_ids[:, -1].unsqueeze(-1)335 336        attention_mask = kwargs.get("attention_mask", None)337        position_ids = kwargs.get("position_ids", None)338 339        if attention_mask is not None and position_ids is None:340            # create position_ids on the fly for batch generation341            position_ids = attention_mask.long().cumsum(-1) - 1342            position_ids.masked_fill_(attention_mask == 0, 1)343            if past:344                position_ids = position_ids[:, -1].unsqueeze(-1)345        else:346            position_ids = None347        return {348            "input_ids": input_ids,349            "past_key_values": past,350            "use_cache": kwargs.get("use_cache"),351            "position_ids": position_ids,352            "attention_mask": attention_mask,353        }354 355    @add_start_docstrings_to_model_forward(BLOOM_INPUTS_DOCSTRING)356    @add_code_sample_docstrings(357        processor_class=_TOKENIZER_FOR_DOC,358        checkpoint=_CHECKPOINT_FOR_DOC,359        output_type=CausalLMOutputWithCrossAttentions,360        config_class=_CONFIG_FOR_DOC,361    )362    def forward(self, input_ids=None, labels=None, return_dict=None, **kwargs):363        r"""364        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):365            Labels for language modeling. Note that the labels **are shifted** inside the model, i.e. you can set366            `labels = input_ids` Indices are selected in `[-100, 0, ..., config.vocab_size]` All labels set to `-100`367            are ignored (masked), the loss is only computed for labels in `[0, ..., config.vocab_size]`368        """369        return_dict = return_dict if return_dict is not None else self.config.use_return_dict370        transformer_outputs = self.transformer.forward(input_ids=input_ids, return_dict=return_dict, **kwargs)371        word_embeddings = self.transformer.word_embeddings.weight372 373        # Switch dtype in case word_embeddings are fp16/bf16374        hidden_states = transformer_outputs[0].to(word_embeddings.dtype)375        lm_logits = F.linear(hidden_states, word_embeddings).float()376 377        loss = None378        if labels is not None:379            # Shift so that tokens < n predict n380            shift_logits = lm_logits[..., :-1, :].contiguous()381            shift_labels = labels[..., 1:].contiguous()382            # Flatten the tokens383            loss_fct = CrossEntropyLoss()384            loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))385 386        if not return_dict:387            output = (lm_logits,) + transformer_outputs[1:]388            return ((loss,) + output) if loss is not None else output389 390        return CausalLMOutputWithCrossAttentions(391            loss=loss,392            logits=lm_logits,393            past_key_values=transformer_outputs.past_key_values,394            hidden_states=transformer_outputs.hidden_states,395            attentions=transformer_outputs.attentions,396        )397 398    @staticmethod399    def _reorder_cache(past: Tuple[Tuple[torch.Tensor]], beam_idx: torch.Tensor) -> Tuple[Tuple[torch.Tensor]]:400        """401        This function is used to re-order the `past_key_values` cache if [`~PreTrainedModel.beam_search`] or402        [`~PreTrainedModel.beam_sample`] is called. This is required to match `past_key_values` with the correct403        beam_idx at every generation step.404        """405        return tuple(406            tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past)407            for layer_past in past408        )409