MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / forward

Method forward

codegeex/torch/codegeex_model.py:968–995  ·  view source on GitHub ↗
(
        self,
        input_ids,
        position_ids,
        attention_mask,
        layer_past=None,
        get_key_value=False,
        prompt_length=None,
        context_length=None,
    )

Source from the content-addressed store, hash-verified

966 self._language_model_key = "language_model"
967
968 def forward(
969 self,
970 input_ids,
971 position_ids,
972 attention_mask,
973 layer_past=None,
974 get_key_value=False,
975 prompt_length=None,
976 context_length=None,
977 ):
978 # Language model.
979 lm_output = self.language_model(input_ids,
980 position_ids,
981 attention_mask,
982 layer_past=layer_past,
983 get_key_value=get_key_value,
984 prompt_length=prompt_length,
985 context_length=context_length)
986
987 if get_key_value:
988 lm_output, presents = lm_output
989
990 output = F.linear(lm_output, self.language_model.embedding.word_embeddings.weight.half())
991
992 if get_key_value:
993 output = [output, presents]
994
995 return output
996
997 def state_dict_for_save_checkpoint(self, destination=None, prefix='',
998 keep_vars=False):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected