Fix bert position ids in DPR convert script (#7776)
* fix bert position ids in DPR convert script * style
This commit is contained in:
@@ -44,7 +44,8 @@ class DPRContextEncoderState(DPRState):
|
||||
print("Loading DPR biencoder from {}".format(self.src_file))
|
||||
saved_state = load_states_from_checkpoint(self.src_file)
|
||||
encoder, prefix = model.ctx_encoder, "ctx_model."
|
||||
state_dict = {}
|
||||
# Fix changes from https://github.com/huggingface/transformers/commit/614fef1691edb806de976756d4948ecbcd0c0ca3
|
||||
state_dict = {"bert_model.embeddings.position_ids": model.ctx_encoder.bert_model.embeddings.position_ids}
|
||||
for key, value in saved_state.model_dict.items():
|
||||
if key.startswith(prefix):
|
||||
key = key[len(prefix) :]
|
||||
@@ -61,7 +62,8 @@ class DPRQuestionEncoderState(DPRState):
|
||||
print("Loading DPR biencoder from {}".format(self.src_file))
|
||||
saved_state = load_states_from_checkpoint(self.src_file)
|
||||
encoder, prefix = model.question_encoder, "question_model."
|
||||
state_dict = {}
|
||||
# Fix changes from https://github.com/huggingface/transformers/commit/614fef1691edb806de976756d4948ecbcd0c0ca3
|
||||
state_dict = {"bert_model.embeddings.position_ids": model.question_encoder.bert_model.embeddings.position_ids}
|
||||
for key, value in saved_state.model_dict.items():
|
||||
if key.startswith(prefix):
|
||||
key = key[len(prefix) :]
|
||||
@@ -77,7 +79,10 @@ class DPRReaderState(DPRState):
|
||||
model = DPRReader(DPRConfig(**BertConfig.get_config_dict("bert-base-uncased")[0]))
|
||||
print("Loading DPR reader from {}".format(self.src_file))
|
||||
saved_state = load_states_from_checkpoint(self.src_file)
|
||||
state_dict = {}
|
||||
# Fix changes from https://github.com/huggingface/transformers/commit/614fef1691edb806de976756d4948ecbcd0c0ca3
|
||||
state_dict = {
|
||||
"encoder.bert_model.embeddings.position_ids": model.span_predictor.encoder.bert_model.embeddings.position_ids
|
||||
}
|
||||
for key, value in saved_state.model_dict.items():
|
||||
if key.startswith("encoder.") and not key.startswith("encoder.encode_proj"):
|
||||
key = "encoder.bert_model." + key[len("encoder.") :]
|
||||
|
||||
Reference in New Issue
Block a user