如何在Databricks中集成并使用Azure训练的已有Pickle格式ML模型?
将Pickle格式的外部模型集成到Databricks工作流的最佳实践
一、模型上传与存储
1. 上传至DBFS(基础存储方案)
- UI操作:左侧导航栏选「数据」→「DBFS」,创建专属文件夹(如
/models/external/),直接上传本地Pickle文件;若模型存于Azure Blob Storage,可通过dbutils.fs.mount命令挂载存储容器到DBFS,直接访问远程文件:dbutils.fs.mount( source = "wasbs://<container-name>@<storage-account-name>.blob.core.windows.net", mount_point = "/mnt/azure-models", extra_configs = {"fs.azure.account.key.<storage-account-name>.blob.core.windows.net": "<account-key>"} ) - CLI操作:安装Databricks CLI后,执行命令直接上传:
databricks fs cp local/path/your-model.pickle dbfs:/models/external/your-model.pickle
2. 用Unity Catalog管理(生产环境推荐)
创建模型目录,将Pickle文件关联到Unity Catalog的外部存储位置或托管位置,统一管控权限、版本和生命周期,适配企业级数据治理需求。
二、加载模型并验证兼容性
先在Databricks笔记本中加载模型,确认环境兼容性:
import pickle import pandas as pd # 从DBFS加载模型 model_path = "/dbfs/models/external/your-model.pickle" with open(model_path, "rb") as f: model = pickle.load(f) # 用测试数据验证模型功能 test_data = pd.DataFrame({ "feature1": [1.2, 3.4], "feature2": [5.6, 7.8] }) predictions = model.predict(test_data) print(predictions)
注意:若模型依赖特定版本的库(如scikit-learn、pandas),需在集群中安装对应版本:
%pip install scikit-learn==1.2.2 pandas==1.5.3
三、接入Databricks模型工作流
1. 注册到MLflow Model Registry
MLflow是Databricks默认的模型管理工具,外部模型也可纳入统一管理:
import mlflow import mlflow.sklearn # 绑定实验(可选,便于追踪) mlflow.set_experiment("/Users/your-email@domain.com/external-model-tracking") with mlflow.start_run(): # 记录训练环境信息(可选) mlflow.log_param("training_platform", "Azure VM") # 注册模型到Registry,若为非sklearn模型,可使用mlflow.pyfunc.log_model自定义 mlflow.sklearn.log_model( model, artifact_path="model", registered_model_name="azure-trained-external-model" )
注册后,在Databricks「模型」页面可查看模型版本、设置部署阶段(Staging/Production)。
2. 部署为实时推理服务
注册后的模型可直接部署到Databricks Model Serving:
- 进入模型注册表,选择目标版本→点击「部署」→「部署到Model Serving」
- 配置集群规格(如CPU/GPU类型),启动服务后即可获得REST API端点,通过HTTP请求调用实时预测。
3. 构建批量预测工作流
将模型封装为Spark UDF,实现分布式批量预测:
import mlflow.pyfunc # 加载生产版本模型为Spark UDF spark_udf = mlflow.pyfunc.spark_udf( spark, model_uri="models:/azure-trained-external-model/Production" ) # 读取批量数据并生成预测结果 batch_data = spark.read.table("your_catalog.your_schema.raw_data") result_data = batch_data.withColumn("prediction", spark_udf(*batch_data.columns)) # 保存结果到指定表 result_data.write.mode("overwrite").saveAsTable("your_catalog.your_schema.prediction_results")
四、环境一致性保障
- 导出训练环境的
requirements.txt,上传至DBFS后在集群中批量安装依赖:%pip install -r /dbfs/models/external/requirements.txt - 若存在自定义依赖或特殊环境,可构建专属Docker镜像,以此创建Databricks集群,确保训练与推理环境完全一致。
内容的提问来源于stack exchange,提问作者zzzz
相关产品推荐
相关产品推荐

