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

Class QueryEmbedding

codegeex/paddle/codegeex_model.py:758–815  ·  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

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

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected