如何用pickle或dill保存存储了已训练DistillBert模型的字典?
保存包含训练后DistillBert模型的字典的方法
方法1:使用pickle实现保存和加载
大多数常规训练场景下,pickle可直接完成该需求,操作如下:
- 保存代码
import pickle # 假设model_dict为你存储模型的字典,格式为 {数据集名称: 训练完成的DistillBert模型实例} model_dict = {"dataset_a": fitted_model_a, "dataset_b": fitted_model_b} # 序列化保存到本地文件 with open("distillbert_model_dict.pkl", "wb") as f: pickle.dump(model_dict, f)
- 加载代码
import pickle # 加载前必须导入对应模型类,保证当前命名空间存在类定义,否则会抛出序列化错误 from transformers import DistilBertModel with open("distillbert_model_dict.pkl", "rb") as f: loaded_model_dict = pickle.load(f) # 直接调用字典内的模型即可 model_output = loaded_model_dict["dataset_a"](input_ids, attention_mask)
如果你的模型绑定了自定义钩子、自定义层或者特殊上下文对象,pickle可能会抛出PicklingError,此时可使用dill完成操作。
方法2:使用dill实现保存和加载(pickle失效时适用)
dill是pickle的扩展库,支持序列化更多类型的Python对象,适配复杂场景的模型存储需求。
首先执行安装命令:pip install dill
- 保存代码
import dill model_dict = {"dataset_a": fitted_model_a, "dataset_b": fitted_model_b} with open("distillbert_model_dict.dill", "wb") as f: dill.dump(model_dict, f)
- 加载代码
import dill # 跨环境加载时建议提前导入对应模型类,避免依赖版本差异导致的报错 from transformers import DistilBertModel with open("distillbert_model_dict.dill", "rb") as f: loaded_model_dict = dill.load(f) # 调用逻辑和原生模型完全一致 model_output = loaded_model_dict["dataset_b"](input_ids, attention_mask)
注意:跨设备/跨环境加载模型时,需保证两端的transformers、torch等依赖库版本一致,避免出现模型结构不匹配、运行异常的问题。
内容的提问来源于stack exchange,提问作者user16894190
相关产品推荐
相关产品推荐

