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

Class CodeGeeXModel

codegeex/oneflow/codegeex_model.py:1039–1102  ·  view source on GitHub ↗

CodeGeeX: A Multilingual Code Generation Model.

Source from the content-addressed store, hash-verified

1037
1038
1039class CodeGeeXModel(torch.nn.Module):
1040 """CodeGeeX: A Multilingual Code Generation Model."""
1041
1042 def __init__(
1043 self,
1044 hidden_size,
1045 num_layers,
1046 num_attention_heads,
1047 padded_vocab_size,
1048 max_position_embeddings,
1049 ):
1050 super(CodeGeeXModel, self).__init__()
1051
1052 self.language_model = TransformerLanguageModel(hidden_size,
1053 num_layers,
1054 num_attention_heads,
1055 padded_vocab_size,
1056 max_position_embeddings)
1057 self._language_model_key = "language_model"
1058
1059 def forward(
1060 self,
1061 input_ids,
1062 position_ids,
1063 attention_mask,
1064 layer_past=None,
1065 get_key_value=False,
1066 prompt_length=None,
1067 context_length=None,
1068 ):
1069 # Language model.
1070 lm_output = self.language_model(input_ids,
1071 position_ids,
1072 attention_mask,
1073 layer_past=layer_past,
1074 get_key_value=get_key_value,
1075 prompt_length=prompt_length,
1076 context_length=context_length)
1077
1078 if get_key_value:
1079 lm_output, presents = lm_output
1080
1081 output = F.linear(lm_output, self.language_model.embedding.word_embeddings.weight.half())
1082
1083 if get_key_value:
1084 output = [output, presents]
1085
1086 return output
1087
1088 def state_dict_for_save_checkpoint(self, destination=None, prefix='',
1089 keep_vars=False):
1090
1091 state_dict_ = {}
1092 state_dict_[self._language_model_key] \
1093 = self.language_model.state_dict_for_save_checkpoint(
1094 destination, prefix, keep_vars)
1095 return state_dict_
1096

Callers 1

model_providerFunction · 0.90

Calls

no outgoing calls

Tested by 1

model_providerFunction · 0.72