如何使用cloudpickle将模型保存至Databricks DBFS并加载?
将模型持久化到DBFS的解决方案
保存模型到DBFS
你可以通过两种方式将序列化后的模型写入DBFS,避免仅存在于会话变量中:
方法1:利用DBFS本地挂载路径(推荐)
Databricks将DBFS挂载到了本地文件系统的/dbfs/路径下,直接用Python文件操作即可读写:
import cloudpickle # 序列化模型 pickled_model = cloudpickle.dumps(final_model) # 确保目标目录存在(不存在则创建) dbutils.fs.mkdirs('/ml_models') # 写入DBFS文件 with open('/dbfs/ml_models/causal_forest_model.pkl', 'wb') as f: f.write(pickled_model)
方法2:使用dbutils.fs工具类
适合不想依赖本地挂载路径的场景,通过二进制流写入:
import cloudpickle import io # 序列化模型 pickled_model = cloudpickle.dumps(final_model) # 确保目标目录存在 dbutils.fs.mkdirs('/ml_models') # 将二进制数据写入DBFS model_bytes = io.BytesIO(pickled_model) dbutils.fs.put('/ml_models/causal_forest_model.pkl', model_bytes.getvalue(), overwrite=True)
从DBFS加载模型
无论会话是否过期,只要有权限访问目标路径,就能重新加载模型:
对应方法1的加载方式
import cloudpickle # 从DBFS读取并加载模型 with open('/dbfs/ml_models/causal_forest_model.pkl', 'rb') as f: loaded_model = cloudpickle.load(f) # 验证加载结果 print(loaded_model) # 输出示例:<econml.dml.causal_forest.CausalForestDML at 0x...>
对应方法2的加载方式
import cloudpickle import io # 读取DBFS中的模型二进制数据 model_data = dbutils.fs.cat('/ml_models/causal_forest_model.pkl') # 反序列化加载模型 loaded_model = cloudpickle.load(io.BytesIO(model_data)) print(loaded_model)
关键注意事项
- 提前用
dbutils.fs.mkdirs('<目标目录路径>')创建存储目录,避免写入失败 - 确认当前用户对目标DBFS路径有读写权限
- 大型模型推荐使用本地挂载路径的文件操作,性能更稳定
内容的提问来源于stack exchange,提问作者titutubs
相关产品推荐
相关产品推荐

