(
self,
input_ids: torch.LongTensor = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_values: Optional[HybridCache] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
use_cache: Optional[bool] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
cache_position: Optional[torch.LongTensor] = None,
)
| 678 | |
| 679 | @add_start_docstrings_to_model_forward(GEMMA2_INPUTS_DOCSTRING) |
| 680 | def forward( |
| 681 | self, |
| 682 | input_ids: torch.LongTensor = None, |
| 683 | attention_mask: Optional[torch.Tensor] = None, |
| 684 | position_ids: Optional[torch.LongTensor] = None, |
| 685 | past_key_values: Optional[HybridCache] = None, |
| 686 | inputs_embeds: Optional[torch.FloatTensor] = None, |
| 687 | use_cache: Optional[bool] = None, |
| 688 | output_attentions: Optional[bool] = None, |
| 689 | output_hidden_states: Optional[bool] = None, |
| 690 | return_dict: Optional[bool] = None, |
| 691 | cache_position: Optional[torch.LongTensor] = None, |
| 692 | ) -> Union[Tuple, BaseModelOutputWithPast]: |
| 693 | output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions |
| 694 | output_hidden_states = ( |
| 695 | output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states |
| 696 | ) |
| 697 | use_cache = use_cache if use_cache is not None else self.config.use_cache |
| 698 | return_dict = return_dict if return_dict is not None else self.config.use_return_dict |
| 699 | |
| 700 | if (input_ids is None) ^ (inputs_embeds is not None): |
| 701 | raise ValueError("You must specify exactly one of input_ids or inputs_embeds") |
| 702 | |
| 703 | if self.gradient_checkpointing and self.training and use_cache: |
| 704 | logger.warning_once( |
| 705 | "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`." |
| 706 | ) |
| 707 | use_cache = False |
| 708 | |
| 709 | if inputs_embeds is None: |
| 710 | inputs_embeds = self.embed_tokens(input_ids) |
| 711 | |
| 712 | if use_cache and past_key_values is None and not self.training: |
| 713 | batch_size, seq_len, _ = inputs_embeds.shape |
| 714 | past_key_values = HybridCache( |
| 715 | self.config, |
| 716 | batch_size=batch_size, |
| 717 | max_cache_len=seq_len, |
| 718 | device=self.device, |
| 719 | dtype=inputs_embeds.dtype, |
| 720 | ) |
| 721 | |
| 722 | if cache_position is None: |
| 723 | past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 |
| 724 | cache_position = torch.arange( |
| 725 | past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device |
| 726 | ) |
| 727 | |
| 728 | if position_ids is None: |
| 729 | position_ids = cache_position.unsqueeze(0) |
| 730 | |
| 731 | causal_mask = self._update_causal_mask( |
| 732 | attention_mask, inputs_embeds, cache_position, past_key_values, output_attentions |
| 733 | ) |
| 734 | |
| 735 | # embed positions |
| 736 | hidden_states = inputs_embeds |
| 737 |
nothing calls this directly
no test coverage detected