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

Class CodeGeeXModel

codegeex/torch/codegeex_model.py:948–1011  ·  view source on GitHub ↗

CodeGeeX: A Multilingual Code Generation Model.

Source from the content-addressed store, hash-verified

946
947
948class CodeGeeXModel(torch.nn.Module):
949 """CodeGeeX: A Multilingual Code Generation Model."""
950
951 def __init__(
952 self,
953 hidden_size,
954 num_layers,
955 num_attention_heads,
956 padded_vocab_size,
957 max_position_embeddings,
958 ):
959 super(CodeGeeXModel, self).__init__()
960
961 self.language_model = TransformerLanguageModel(hidden_size,
962 num_layers,
963 num_attention_heads,
964 padded_vocab_size,
965 max_position_embeddings)
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):
999
1000 state_dict_ = {}
1001 state_dict_[self._language_model_key] \
1002 = self.language_model.state_dict_for_save_checkpoint(
1003 destination, prefix, keep_vars)
1004 return state_dict_
1005

Callers 3

model_providerFunction · 0.90
model_providerFunction · 0.90
model_providerFunction · 0.90

Calls

no outgoing calls

Tested by 2

model_providerFunction · 0.72
model_providerFunction · 0.72