(
self,
input_ids,
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.BoolTensor] = None,
full_attention_mask: Optional[torch.BoolTensor] = None,
past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None,
inputs_embeds: Optional[torch.Tensor] = None,
use_cache: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
)
| 784 | return past_key_values |
| 785 | |
| 786 | def forward( |
| 787 | self, |
| 788 | input_ids, |
| 789 | position_ids: Optional[torch.Tensor] = None, |
| 790 | attention_mask: Optional[torch.BoolTensor] = None, |
| 791 | full_attention_mask: Optional[torch.BoolTensor] = None, |
| 792 | past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None, |
| 793 | inputs_embeds: Optional[torch.Tensor] = None, |
| 794 | use_cache: Optional[bool] = None, |
| 795 | output_hidden_states: Optional[bool] = None, |
| 796 | return_dict: Optional[bool] = None, |
| 797 | ): |
| 798 | output_hidden_states = ( |
| 799 | output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states |
| 800 | ) |
| 801 | use_cache = use_cache if use_cache is not None else self.config.use_cache |
| 802 | return_dict = return_dict if return_dict is not None else self.config.use_return_dict |
| 803 | |
| 804 | batch_size, seq_length = input_ids.shape |
| 805 | |
| 806 | if inputs_embeds is None: |
| 807 | inputs_embeds = self.embedding(input_ids) |
| 808 | |
| 809 | if self.pre_seq_len is not None: |
| 810 | if past_key_values is None: |
| 811 | past_key_values = self.get_prompt(batch_size=batch_size, device=input_ids.device, |
| 812 | dtype=inputs_embeds.dtype) |
| 813 | if attention_mask is not None: |
| 814 | attention_mask = torch.cat([attention_mask.new_ones((batch_size, self.pre_seq_len)), |
| 815 | attention_mask], dim=-1) |
| 816 | |
| 817 | if full_attention_mask is None: |
| 818 | if (attention_mask is not None and not attention_mask.all()) or (past_key_values and seq_length != 1): |
| 819 | full_attention_mask = self.get_masks(input_ids, past_key_values, padding_mask=attention_mask) |
| 820 | |
| 821 | # Rotary positional embeddings |
| 822 | rotary_pos_emb = self.rotary_pos_emb(self.seq_length) |
| 823 | if position_ids is not None: |
| 824 | rotary_pos_emb = rotary_pos_emb[position_ids] |
| 825 | else: |
| 826 | rotary_pos_emb = rotary_pos_emb[None, :seq_length] |
| 827 | rotary_pos_emb = rotary_pos_emb.transpose(0, 1).contiguous() |
| 828 | |
| 829 | # Run encoder. |
| 830 | hidden_states, presents, all_hidden_states, all_self_attentions = self.encoder( |
| 831 | inputs_embeds, full_attention_mask, rotary_pos_emb=rotary_pos_emb, |
| 832 | kv_caches=past_key_values, use_cache=use_cache, output_hidden_states=output_hidden_states |
| 833 | ) |
| 834 | |
| 835 | if not return_dict: |
| 836 | return tuple(v for v in [hidden_states, presents, all_hidden_states, all_self_attentions] if v is not None) |
| 837 | |
| 838 | return BaseModelOutputWithPast( |
| 839 | last_hidden_state=hidden_states, |
| 840 | past_key_values=presents, |
| 841 | hidden_states=all_hidden_states, |
| 842 | attentions=all_self_attentions, |
| 843 | ) |
nothing calls this directly
no test coverage detected