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
| 677 | |
| 678 | |
| 679 | class Embedding(torch.nn.Module): |
| 680 | """Language model embeddings. |
| 681 | |
| 682 | Arguments: |
| 683 | hidden_size: hidden size |
| 684 | vocab_size: vocabulary size |
| 685 | max_sequence_length: maximum size of sequence. This |
| 686 | is used for positional embedding |
| 687 | """ |
| 688 | |
| 689 | def __init__( |
| 690 | self, |
| 691 | hidden_size, |
| 692 | vocab_size, |
| 693 | max_sequence_length, |
| 694 | ): |
| 695 | super(Embedding, self).__init__() |
| 696 | self.hidden_size = hidden_size |
| 697 | self.vocab_size = vocab_size |
| 698 | self.max_sequence_length = max_sequence_length |
| 699 | |
| 700 | # Word embeddings. |
| 701 | self.word_embeddings = torch.nn.Embedding(self.vocab_size, self.hidden_size) |
| 702 | self._word_embeddings_key = 'word_embeddings' |
| 703 | |
| 704 | # Position embedding. |
| 705 | self.position_embeddings = torch.nn.Embedding(self.max_sequence_length, self.hidden_size) |
| 706 | self.position_embeddings = self.position_embeddings.half() |
| 707 | self._position_embeddings_key = 'position_embeddings' |
| 708 | |
| 709 | def forward(self, input_ids, position_ids): |
| 710 | # Embeddings. |
| 711 | words_embeddings = self.word_embeddings(input_ids) |
| 712 | position_embeddings = self.position_embeddings(position_ids) |
| 713 | embeddings = words_embeddings + position_embeddings |
| 714 | |
| 715 | return embeddings |
| 716 | |
| 717 | def state_dict_for_save_checkpoint(self, destination=None, prefix='', |
| 718 | keep_vars=False): |
| 719 | """For easy load.""" |
| 720 | |
| 721 | state_dict_ = {} |
| 722 | state_dict_[self._word_embeddings_key] \ |
| 723 | = self.word_embeddings.state_dict(destination, prefix, keep_vars) |
| 724 | state_dict_[self._position_embeddings_key] \ |
| 725 | = self.position_embeddings.state_dict( |
| 726 | destination, prefix, keep_vars) |
| 727 | |
| 728 | return state_dict_ |
| 729 | |
| 730 | def load_state_dict(self, state_dict, strict=True): |
| 731 | """Customized load.""" |
| 732 | |
| 733 | # Word embedding. |
| 734 | if self._word_embeddings_key in state_dict: |
| 735 | state_dict_ = state_dict[self._word_embeddings_key] |
| 736 | else: |