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

Class CodeGeeXModel

codegeex/paddle/codegeex_model.py:947–1010  ·  view source on GitHub ↗

CodeGeeX: A Multilingual Code Generation Model.

Source from the content-addressed store, hash-verified

945
946
947class CodeGeeXModel(paddle.nn.Layer):
948 """CodeGeeX: A Multilingual Code Generation Model."""
949
950 def __init__(
951 self,
952 hidden_size,
953 num_layers,
954 num_attention_heads,
955 padded_vocab_size,
956 max_position_embeddings,
957 ):
958 super(CodeGeeXModel, self).__init__()
959
960 self.language_model = TransformerLanguageModel(hidden_size,
961 num_layers,
962 num_attention_heads,
963 padded_vocab_size,
964 max_position_embeddings)
965 self._language_model_key = "language_model"
966
967 def forward(
968 self,
969 input_ids,
970 position_ids,
971 attention_mask,
972 layer_past=None,
973 get_key_value=False,
974 prompt_length=None,
975 context_length=None,
976 ):
977 # Language model.
978 lm_output = self.language_model(input_ids,
979 position_ids,
980 attention_mask,
981 layer_past=layer_past,
982 get_key_value=get_key_value,
983 prompt_length=prompt_length,
984 context_length=context_length)
985
986 if get_key_value:
987 lm_output, presents = lm_output
988
989 output = F.linear(lm_output, self.language_model.embedding.word_embeddings.weight.cast("float16").transpose([1, 0]))
990
991 if get_key_value:
992 output = [output, presents]
993
994 return output
995
996 def state_dict_for_save_checkpoint(self, destination=None, prefix='',
997 keep_vars=False):
998
999 state_dict_ = {}
1000 state_dict_[self._language_model_key] \
1001 = self.language_model.state_dict_for_save_checkpoint(
1002 destination, prefix, keep_vars)
1003 return state_dict_
1004

Callers 1

model_providerFunction · 0.90

Calls

no outgoing calls

Tested by 1

model_providerFunction · 0.72