MCPcopy Create free account
hub / github.com/THUDM/LongWriter / forward

Method forward

train/patch/modeling_chatglm.py:786–849  ·  view source on GitHub ↗
(
            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,
    )

Source from the content-addressed store, hash-verified

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(

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected