CodeGeeX: A Multilingual Code Generation Model.
| 946 | |
| 947 | |
| 948 | class 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 |
no outgoing calls