Language model embeddings. Arguments: hidden_size: hidden size vocab_size: vocabulary size max_sequence_length: maximum size of sequence. This is used for positional embedding
| 756 | |
| 757 | |
| 758 | class QueryEmbedding(paddle.nn.Layer): |
| 759 | """Language model embeddings. |
| 760 | |
| 761 | Arguments: |
| 762 | hidden_size: hidden size |
| 763 | vocab_size: vocabulary size |
| 764 | max_sequence_length: maximum size of sequence. This |
| 765 | is used for positional embedding |
| 766 | """ |
| 767 | |
| 768 | def __init__( |
| 769 | self, |
| 770 | hidden_size, |
| 771 | vocab_size, |
| 772 | max_sequence_length, |
| 773 | ): |
| 774 | super(QueryEmbedding, self).__init__() |
| 775 | |
| 776 | self.hidden_size = hidden_size |
| 777 | self.vocab_size = vocab_size |
| 778 | self.max_sequence_length = max_sequence_length |
| 779 | |
| 780 | # Top query position embedding (serial). |
| 781 | self.top_query_embeddings = paddle.nn.Embedding(self.max_sequence_length, self.hidden_size) |
| 782 | self.top_query_embeddings = self.top_query_embeddings.to(dtype="float16") |
| 783 | self._top_query_embeddings_key = 'top_query_embeddings' |
| 784 | |
| 785 | def forward(self, position_ids): |
| 786 | # Embeddings. |
| 787 | embeddings = self.top_query_embeddings(position_ids) |
| 788 | |
| 789 | return embeddings |
| 790 | |
| 791 | def state_dict_for_save_checkpoint(self, destination=None, prefix='', |
| 792 | keep_vars=False): |
| 793 | """For easy load.""" |
| 794 | |
| 795 | state_dict_ = {} |
| 796 | state_dict_[self._top_query_embeddings_key] \ |
| 797 | = self.top_query_embeddings.state_dict( |
| 798 | destination, prefix, keep_vars) |
| 799 | |
| 800 | return state_dict_ |
| 801 | |
| 802 | def set_state_dict(self, state_dict, use_structured_name=True): |
| 803 | """Customized load.""" |
| 804 | |
| 805 | # Position embedding. |
| 806 | if self._top_query_embeddings_key in state_dict: |
| 807 | state_dict_ = state_dict[self._top_query_embeddings_key] |
| 808 | else: |
| 809 | # for backward compatibility. |
| 810 | state_dict_ = {} |
| 811 | for key in state_dict.keys(): |
| 812 | if 'top_query_embeddings' in key: |
| 813 | state_dict_[key.split('top_query_embeddings.')[1]] \ |
| 814 | = state_dict[key] |
| 815 | self.top_query_embeddings.set_state_dict(state_dict_, use_structured_name=use_structured_name) |