bigscience/petals-api
18
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 