MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / Embedding

Class Embedding

codegeex/oneflow/codegeex_model.py:773–849  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

771
772
773class 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.

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected