如何使用本地checkpoint加载dbmdz/bert-base-italian-xxl-cased BERT模型
解决方案
你拿到的是原生TensorFlow Checkpoint格式权重(由.data/.index/.meta三个文件组成),而Hugging Face TFBertModel.from_pretrained默认查找的是封装好的tf_model.h5或pytorch_model.bin格式,参数配置错误才会触发报错。
方法1:直接加载原生Checkpoint
核心注意两个配置:
- 权重路径写Checkpoint前缀(不带
.index/.data/.meta后缀) - 新增
from_tf=True参数,明确告知库加载的是原生TensorFlow权重
代码示例:
from transformers import BertConfig, TFBertModel, BertTokenizer # 替换为你的本地解压文件夹路径 bert_folder = "../../models/pretrained/bert-base-italian-xxl-cased" # 加载配置 config = BertConfig.from_pretrained(bert_folder) # 加载原生TF Checkpoint model = TFBertModel.from_pretrained( f"{bert_folder}/model.ckpt", config=config, from_tf=True, local_files_only=True )
方法2:转换为Hugging Face标准TF格式(方便后续复用)
加载完成后直接调用save_pretrained即可导出tf_model.h5,后续无需再处理原生Checkpoint:
# 保存为标准格式,会在bert_folder下生成tf_model.h5 model.save_pretrained(bert_folder) # 后续直接加载即可 # model = TFBertModel.from_pretrained(bert_folder, local_files_only=True)
功能验证
运行以下代码确认模型加载正常:
# 加载词表 tokenizer = BertTokenizer.from_pretrained(bert_folder) # 测试输入 test_text = "Ciao, questo è un test di funzionamento." inputs = tokenizer(test_text, return_tensors="tf") outputs = model(**inputs) # 正常输出维度为(1, 序列长度, 隐藏层维度)即加载成功 print(outputs.last_hidden_state.shape)
注意事项
如果加载报错,先升级transformers库到最新版本:
pip install --upgrade transformers
内容的提问来源于stack exchange,提问作者Gerardo Zinno
相关产品推荐
相关产品推荐

