@@ -1447,16 +1447,23 @@ def load(cls, pretrained_model_name_or_path, language=None, **kwargs):
14471447 dpr_question_encoder .model = transformers .DPRQuestionEncoder .from_pretrained (farm_lm_model , config = dpr_config , ** kwargs )
14481448 dpr_question_encoder .language = dpr_question_encoder .model .config .language
14491449 else :
1450- model_type = AutoConfig .from_pretrained (pretrained_model_name_or_path ). model_type
1451- if model_type == "dpr" :
1450+ original_model_config = AutoConfig .from_pretrained (pretrained_model_name_or_path )
1451+ if original_model_config . model_type == "dpr" :
14521452 # "pretrained dpr model": load existing pretrained DPRQuestionEncoder model
14531453 dpr_question_encoder .model = transformers .DPRQuestionEncoder .from_pretrained (
14541454 str (pretrained_model_name_or_path ), ** kwargs )
14551455 else :
14561456 # "from scratch": load weights from different architecture (e.g. bert) into DPRQuestionEncoder
1457- dpr_question_encoder .model = transformers .DPRQuestionEncoder (config = transformers .DPRConfig (** kwargs ))
1457+ # but keep config values from original architecture
1458+ # TODO test for architectures other than BERT, e.g. Electra
1459+ if original_model_config .model_type != "bert" :
1460+ logger .warning (f"Using a model of type '{ original_model_config .model_type } ' which might be incompatible with DPR encoders."
1461+ f"Bert based encoders are supported that need input_ids,token_type_ids,attention_mask as input tensors." )
1462+ original_config_dict = vars (original_model_config )
1463+ original_config_dict .update (kwargs )
1464+ dpr_question_encoder .model = transformers .DPRQuestionEncoder (config = transformers .DPRConfig (** original_config_dict ))
14581465 dpr_question_encoder .model .base_model .bert_model = AutoModel .from_pretrained (
1459- str (pretrained_model_name_or_path ), ** kwargs )
1466+ str (pretrained_model_name_or_path ), ** original_config_dict )
14601467 dpr_question_encoder .language = cls ._get_or_infer_language_from_name (language , pretrained_model_name_or_path )
14611468
14621469 return dpr_question_encoder
@@ -1541,16 +1548,25 @@ def load(cls, pretrained_model_name_or_path, language=None, **kwargs):
15411548 dpr_context_encoder .language = dpr_context_encoder .model .config .language
15421549 else :
15431550 # Pytorch-transformer Style
1544- model_type = AutoConfig .from_pretrained (pretrained_model_name_or_path ). model_type
1545- if model_type == "dpr" :
1551+ original_model_config = AutoConfig .from_pretrained (pretrained_model_name_or_path )
1552+ if original_model_config . model_type == "dpr" :
15461553 # "pretrained dpr model": load existing pretrained DPRContextEncoder model
15471554 dpr_context_encoder .model = transformers .DPRContextEncoder .from_pretrained (
15481555 str (pretrained_model_name_or_path ), ** kwargs )
15491556 else :
15501557 # "from scratch": load weights from different architecture (e.g. bert) into DPRContextEncoder
1551- dpr_context_encoder .model = transformers .DPRContextEncoder (config = transformers .DPRConfig (** kwargs ))
1558+ # but keep config values from original architecture
1559+ # TODO test for architectures other than BERT, e.g. Electra
1560+ if original_model_config .model_type != "bert" :
1561+ logger .warning (
1562+ f"Using a model of type '{ original_model_config .model_type } ' which might be incompatible with DPR encoders."
1563+ f"Bert based encoders are supported that need input_ids,token_type_ids,attention_mask as input tensors." )
1564+ original_config_dict = vars (original_model_config )
1565+ original_config_dict .update (kwargs )
1566+ dpr_context_encoder .model = transformers .DPRContextEncoder (
1567+ config = transformers .DPRConfig (** original_config_dict ))
15521568 dpr_context_encoder .model .base_model .bert_model = AutoModel .from_pretrained (
1553- str (pretrained_model_name_or_path ), ** kwargs )
1569+ str (pretrained_model_name_or_path ), ** original_config_dict )
15541570 dpr_context_encoder .language = cls ._get_or_infer_language_from_name (language , pretrained_model_name_or_path )
15551571
15561572 return dpr_context_encoder
0 commit comments