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

Class QueryEmbedding

codegeex/torch/codegeex_model.py:759–816  ·  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

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

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected