CodeGeeX: A Multilingual Code Generation Model.
| 1037 | |
| 1038 | |
| 1039 | class 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 |
no outgoing calls