(
self,
input_ids: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
past_key_values: Optional[Tuple[torch.FloatTensor]] = None,
inputs_embeds: Optional[torch.Tensor] = None,
labels: Optional[Tuple[torch.Tensor]] = None,
use_cache: Optional[bool] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
return_last_logit: Optional[bool] = False,
)
| 784 | } |
| 785 | |
| 786 | def forward( |
| 787 | self, |
| 788 | input_ids: Optional[torch.Tensor] = None, |
| 789 | position_ids: Optional[torch.Tensor] = None, |
| 790 | attention_mask: Optional[torch.Tensor] = None, |
| 791 | past_key_values: Optional[Tuple[torch.FloatTensor]] = None, |
| 792 | inputs_embeds: Optional[torch.Tensor] = None, |
| 793 | labels: Optional[Tuple[torch.Tensor]] = None, |
| 794 | use_cache: Optional[bool] = None, |
| 795 | output_attentions: Optional[bool] = None, |
| 796 | output_hidden_states: Optional[bool] = None, |
| 797 | return_dict: Optional[bool] = None, |
| 798 | return_last_logit: Optional[bool] = False, |
| 799 | ): |
| 800 | use_cache = use_cache if use_cache is not None else self.config.use_cache |
| 801 | return_dict = return_dict if return_dict is not None else self.config.use_return_dict |
| 802 | |
| 803 | transformer_outputs = self.transformer( |
| 804 | input_ids=input_ids, |
| 805 | position_ids=position_ids, |
| 806 | attention_mask=attention_mask, |
| 807 | past_key_values=past_key_values, |
| 808 | inputs_embeds=inputs_embeds, |
| 809 | use_cache=use_cache, |
| 810 | output_hidden_states=output_hidden_states, |
| 811 | return_dict=return_dict, |
| 812 | ) |
| 813 | |
| 814 | hidden_states = transformer_outputs[0] |
| 815 | if return_last_logit: |
| 816 | hidden_states = hidden_states[-1:] |
| 817 | lm_logits = self.transformer.output_layer(hidden_states) |
| 818 | lm_logits = lm_logits.transpose(0, 1).contiguous() |
| 819 | |
| 820 | loss = None |
| 821 | if labels is not None: |
| 822 | lm_logits = lm_logits.to(torch.float32) |
| 823 | # Shift so that tokens < n predict n |
| 824 | shift_logits = lm_logits[..., :-1, :].contiguous() |
| 825 | if isinstance(labels, tuple) or isinstance(labels, list): |
| 826 | labels, weights = labels |
| 827 | shift_labels = labels[..., 1:].contiguous() |
| 828 | if self.pack_loss: |
| 829 | loss_fct = CrossEntropyLoss(ignore_index=-100)#, reduction='none') |
| 830 | loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)) |
| 831 | loss *= weights |
| 832 | else: |
| 833 | loss_fct = CrossEntropyLoss(ignore_index=-100) |
| 834 | loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)) |
| 835 | |
| 836 | lm_logits = lm_logits.to(hidden_states.dtype) |
| 837 | loss = loss.to(hidden_states.dtype) |
| 838 | |
| 839 | if not return_dict: |
| 840 | output = (lm_logits,) + transformer_outputs[1:] |
| 841 | return ((loss,) + output) if loss is not None else output |
| 842 | |
| 843 | return CausalLMOutputWithPast( |
nothing calls this directly
no outgoing calls
no test coverage detected