MCPcopy Create free account
hub / github.com/FreedomIntelligence/CMB / forward

Method forward

workers/chatglm3_modeling.py:786–843  ·  view source on GitHub ↗
(
            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,
    )

Source from the content-addressed store, hash-verified

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 )

Callers

nothing calls this directly

Calls 2

get_promptMethod · 0.95
get_masksMethod · 0.95

Tested by

no test coverage detected