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