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
| 771 | |
| 772 | |
| 773 | class Embedding(torch.nn.Module): |
| 774 | """Language model embeddings. |
| 775 | Arguments: |
| 776 | hidden_size: hidden size |
| 777 | vocab_size: vocabulary size |
| 778 | max_sequence_length: maximum size of sequence. This |
| 779 | is used for positional embedding |
| 780 | """ |
| 781 | |
| 782 | def __init__( |
| 783 | self, |
| 784 | hidden_size, |
| 785 | vocab_size, |
| 786 | max_sequence_length, |
| 787 | ): |
| 788 | super(Embedding, self).__init__() |
| 789 | self.hidden_size = hidden_size |
| 790 | self.vocab_size = vocab_size |
| 791 | self.max_sequence_length = max_sequence_length |
| 792 | |
| 793 | # Word embeddings. |
| 794 | self.word_embeddings = torch.nn.Embedding(self.vocab_size, self.hidden_size) |
| 795 | self._word_embeddings_key = 'word_embeddings' |
| 796 | |
| 797 | # Position embedding. |
| 798 | self.position_embeddings = torch.nn.Embedding(self.max_sequence_length, self.hidden_size) |
| 799 | self.position_embeddings = self.position_embeddings.half() |
| 800 | self._position_embeddings_key = 'position_embeddings' |
| 801 | |
| 802 | def forward(self, input_ids, position_ids): |
| 803 | # Embeddings. |
| 804 | words_embeddings = self.word_embeddings(input_ids) |
| 805 | position_embeddings = self.position_embeddings(position_ids) |
| 806 | embeddings = words_embeddings + position_embeddings |
| 807 | |
| 808 | return embeddings |
| 809 | |
| 810 | def state_dict_for_save_checkpoint(self, destination=None, prefix='', |
| 811 | keep_vars=False): |
| 812 | """For easy load.""" |
| 813 | |
| 814 | state_dict_ = {} |
| 815 | state_dict_[self._word_embeddings_key] \ |
| 816 | = self.word_embeddings.state_dict(destination, prefix, keep_vars) |
| 817 | state_dict_[self._position_embeddings_key] \ |
| 818 | = self.position_embeddings.state_dict( |
| 819 | destination, prefix, keep_vars) |
| 820 | |
| 821 | return state_dict_ |
| 822 | |
| 823 | def load_state_dict(self, state_dict, strict=True): |
| 824 | """Customized load.""" |
| 825 | |
| 826 | # Word embedding. |
| 827 | if self._word_embeddings_key in state_dict: |
| 828 | state_dict_ = state_dict[self._word_embeddings_key] |
| 829 | else: |
| 830 | # for backward compatibility. |