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
| 850 | |
| 851 | |
| 852 | class QueryEmbedding(torch.nn.Module): |
| 853 | """Language model embeddings. |
| 854 | Arguments: |
| 855 | hidden_size: hidden size |
| 856 | vocab_size: vocabulary size |
| 857 | max_sequence_length: maximum size of sequence. This |
| 858 | is used for positional embedding |
| 859 | """ |
| 860 | |
| 861 | def __init__( |
| 862 | self, |
| 863 | hidden_size, |
| 864 | vocab_size, |
| 865 | max_sequence_length, |
| 866 | ): |
| 867 | super(QueryEmbedding, self).__init__() |
| 868 | |
| 869 | self.hidden_size = hidden_size |
| 870 | self.vocab_size = vocab_size |
| 871 | self.max_sequence_length = max_sequence_length |
| 872 | |
| 873 | # Top query position embedding (serial). |
| 874 | self.top_query_embeddings = torch.nn.Embedding(self.max_sequence_length, self.hidden_size) |
| 875 | self.top_query_embeddings = self.top_query_embeddings.half() |
| 876 | self._top_query_embeddings_key = 'top_query_embeddings' |
| 877 | |
| 878 | def forward(self, position_ids): |
| 879 | # Embeddings. |
| 880 | embeddings = self.top_query_embeddings(position_ids) |
| 881 | |
| 882 | return embeddings |
| 883 | |
| 884 | def state_dict_for_save_checkpoint(self, destination=None, prefix='', |
| 885 | keep_vars=False): |
| 886 | """For easy load.""" |
| 887 | |
| 888 | state_dict_ = {} |
| 889 | state_dict_[self._top_query_embeddings_key] \ |
| 890 | = self.top_query_embeddings.state_dict( |
| 891 | destination, prefix, keep_vars) |
| 892 | |
| 893 | return state_dict_ |
| 894 | |
| 895 | def load_state_dict(self, state_dict, strict=True): |
| 896 | """Customized load.""" |
| 897 | |
| 898 | # Position embedding. |
| 899 | if self._top_query_embeddings_key in state_dict: |
| 900 | state_dict_ = state_dict[self._top_query_embeddings_key] |
| 901 | else: |
| 902 | # for backward compatibility. |
| 903 | state_dict_ = {} |
| 904 | for key in state_dict.keys(): |
| 905 | if 'top_query_embeddings' in key: |
| 906 | state_dict_[key.split('top_query_embeddings.')[1]] \ |
| 907 | = state_dict[key] |
| 908 | self.top_query_embeddings.load_state_dict(state_dict_, strict=strict) |
| 909 |