如何将BERT权重加载到HuggingFace DPRQuestionEncoder用于RAG微调
问题背景
希望将BERT(或其他Transformer模型)的权重加载到DPRQuestionEncoder架构中,以便使用HuggingFace的save_pretrained方法,将保存后的模型接入RAG架构完成端到端微调。
运行的测试代码如下:
from transformers import DPRQuestionEncoder model = DPRQuestionEncoder.from_pretrained('bert-base-uncased')
触发报错如下:
You are using a model of type bert to instantiate a model of type dpr. This is not supported for all configurations of models and can yield errors. NotImplementedErrorTraceback (most recent call last) <ipython-input-27-1f1b990b906b> in <module> ----> 1 model = DPRQuestionEncoder.from_pretrained(model_name) 2 # https://github.com/huggingface/transformers/blob/41cd52a768a222a13da0c6aaae877a92fc6c783c/src/transformers/models/dpr/modeling_dpr.py#L520 /opt/conda/lib/python3.8/site-packages/transformers/modeling_utils.py in from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs) 1211 ) 1212 -> 1213 model, missing_keys, unexpected_keys, error_msgs = cls._load_state_dict_into_model( 1214 model, state_dict, pretrained_model_name_or_path, _fast_init=_fast_init 1215 ) /opt/conda/lib/python3.8/site-packages/transformers/modeling_utils.py in _load_state_dict_into_model(cls, model, state_dict, pretrained_model_name_or_path, _fast_init) 1286 ) 1287 for module in unintialized_modules: -> 1288 model._init_weights(module) 1289 1290 # copy state_dict so _load_from_state_dict can modify it /opt/conda/lib/python3.8/site-packages/transformers/modeling_utils.py in _init_weights(self, module) 515 Initialize the weights. This method should be overridden by derived class. 516 """ --> 517 raise NotImplementedError(f"Make sure `_init_weigths` is implemented for {self.__class__}") 518 519 def tie_weights(self): NotImplementedError: Make sure `_init_weigths` is implemented for <class 'transformers.models.dpr.modeling_dpr.DPRQuestionEncoder'>
原因说明
- DPRQuestionEncoder和原生BERT的权重命名空间不一致,直接调用
from_pretrained加载BERT权重会出现大量key不匹配,未匹配到的权重会触发模型初始化逻辑 - DPRQuestionEncoder本身没有实现
_init_weights方法,初始化未匹配模块时就会抛出对应的NotImplementedError
解决方案
采用手动权重映射的方式加载,代码如下:
from transformers import BertModel, DPRQuestionEncoder, DPRConfig # 加载原生BERT模型 bert_model = BertModel.from_pretrained('bert-base-uncased') # 初始化DPRQuestionEncoder,配置参数和BERT对齐 dpr_config = DPRConfig( hidden_size=bert_model.config.hidden_size, num_hidden_layers=bert_model.config.num_hidden_layers, num_attention_heads=bert_model.config.num_attention_heads, intermediate_size=bert_model.config.intermediate_size, hidden_act=bert_model.config.hidden_act, layer_norm_eps=bert_model.config.layer_norm_eps ) dpr_model = DPRQuestionEncoder(dpr_config) # 权重映射:BERT主干权重对应DPR的question_encoder子模块 bert_state_dict = bert_model.state_dict() dpr_state_dict = dpr_model.state_dict() for k, v in bert_state_dict.items(): dpr_key = f"question_encoder.{k}" if dpr_key in dpr_state_dict: dpr_state_dict[dpr_key] = v # 加载映射后的权重 dpr_model.load_state_dict(dpr_state_dict) # 保存后即可接入RAG架构使用 dpr_model.save_pretrained("./dpr-question-encoder-from-bert-base")
注意事项
- 需保证DPR和BERT的模型配置参数完全对齐,否则会出现维度不匹配的加载错误
- 如果替换为其他Transformer模型作为底座,只需调整对应的权重前缀映射规则即可
- 加载完成后可先输入测试样本验证输出逻辑正常,再进行后续微调操作
内容的提问来源于stack exchange,提问作者albero
相关产品推荐
相关产品推荐

