MCPcopy Create free account
hub / github.com/Tele-AI/Telechat / forward

Method forward

models/12B/modeling_telechat.py:636–738  ·  view source on GitHub ↗
(
            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,
    )

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 1

_prepare_attn_maskMethod · 0.95

Tested by

no test coverage detected