(
self,
input_ids: Optional[torch.LongTensor] = None,
past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.LongTensor] = None,
use_cache: Optional[bool] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
**deprecated_arguments,
)
| 634 | self.word_embeddings = new_embeddings |
| 635 | |
| 636 | def forward( |
| 637 | self, |
| 638 | input_ids: Optional[torch.LongTensor] = None, |
| 639 | past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None, |
| 640 | attention_mask: Optional[torch.Tensor] = None, |
| 641 | inputs_embeds: Optional[torch.LongTensor] = None, |
| 642 | use_cache: Optional[bool] = None, |
| 643 | output_attentions: Optional[bool] = None, |
| 644 | output_hidden_states: Optional[bool] = None, |
| 645 | return_dict: Optional[bool] = None, |
| 646 | **deprecated_arguments, |
| 647 | ) -> Union[Tuple[torch.Tensor, ...], BaseModelOutputWithPastAndCrossAttentions]: |
| 648 | |
| 649 | output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions |
| 650 | output_hidden_states = ( |
| 651 | output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states |
| 652 | ) |
| 653 | use_cache = use_cache if use_cache is not None else self.config.use_cache |
| 654 | return_dict = return_dict if return_dict is not None else self.config.use_return_dict |
| 655 | |
| 656 | if input_ids is not None: |
| 657 | batch_size, seq_length = input_ids.shape |
| 658 | elif inputs_embeds is not None: |
| 659 | batch_size, seq_length, _ = inputs_embeds.shape |
| 660 | |
| 661 | if past_key_values is None: |
| 662 | past_key_values = tuple([None] * len(self.h)) |
| 663 | |
| 664 | if inputs_embeds is None: |
| 665 | inputs_embeds = self.word_embeddings(input_ids) |
| 666 | hidden_states = inputs_embeds |
| 667 | |
| 668 | if self.config.embed_layernorm: |
| 669 | hidden_states = self.word_embeddings_layernorm(inputs_embeds) |
| 670 | |
| 671 | presents = () if use_cache else None |
| 672 | all_self_attentions = () if output_attentions else None |
| 673 | all_hidden_states = () if output_hidden_states else None |
| 674 | |
| 675 | if self.gradient_checkpointing and self.training: |
| 676 | if use_cache: |
| 677 | use_cache = False |
| 678 | |
| 679 | seq_length_with_past = seq_length |
| 680 | past_key_values_length = 0 |
| 681 | if past_key_values[0] is not None: |
| 682 | past_key_values_length = past_key_values[0][0].shape[2] |
| 683 | seq_length_with_past = seq_length_with_past + past_key_values_length |
| 684 | if attention_mask is None: |
| 685 | attention_mask = torch.ones((batch_size, seq_length_with_past), device=hidden_states.device) |
| 686 | else: |
| 687 | attention_mask = attention_mask.to(hidden_states.device) |
| 688 | causal_mask = self._prepare_attn_mask( |
| 689 | attention_mask, |
| 690 | input_shape=(batch_size, seq_length), |
| 691 | past_key_values_length=past_key_values_length, |
| 692 | ) |
| 693 |
nothing calls this directly
no test coverage detected