Databricks中Spark用applyInPandas加载joblib模型失败的解决方法
解决Spark applyInPandas中joblib加载自定义类失败的问题
问题核心分析
直接调用Process(df)时代码在driver节点执行,自定义类WeightedEnsembleRegressor处于当前命名空间,joblib可以正常反序列化;但使用applyInPandas时,Spark会将Process函数序列化后分发到worker节点执行,worker的执行上下文是workerwrap模块,而自定义类未被同步到worker的全局命名空间,导致pickle反序列化时找不到类,抛出AttributeError。
解决思路与步骤
1. 将自定义类封装为独立模块并分发到所有节点
把WeightedEnsembleRegressor类单独写入一个Python文件(如ensemble_models.py),上传到Databricks的DBFS或共享工作目录,再通过Spark API将文件分发到所有worker节点,确保worker能导入该类:
# 在主脚本中添加模块分发 spark.sparkContext.addPyFile("/dbfs/path/to/ensemble_models.py")
2. 在Process函数内部导入自定义类
由于applyInPandas的函数在worker上下文执行,必须在Process函数内部显式导入自定义类,而非仅在driver全局导入:
def Process(df): import traceback import xgboost import pickle import joblib import sys import pandas as pd # 从分发的模块导入自定义类 from ensemble_models import WeightedEnsembleRegressor # 后续逻辑保持不变,注意模型路径改为worker可访问的DBFS路径
3. 确保模型文件路径全局可访问
不要使用driver本地路径存储模型,改用DBFS路径(如/dbfs/path/to/final_model_v3_1.joblib),保证所有worker节点都能读取到模型文件。
4. 统一节点依赖版本
确认driver和所有worker节点的joblib、xgboost、catboost等依赖版本完全一致,版本不匹配也会导致反序列化失败。
修改后的代码示例
第一步:创建ensemble_models.py
import numpy as np import joblib class WeightedEnsembleRegressor: """ Holds models, their weights, scaler and feature order for prediction & persistence. """ def __init__(self, trained_models, model_weights, scaler, feature_order): self.trained_models = trained_models self.model_weights = model_weights self.scaler = scaler self.feature_order = feature_order def save(self, path): joblib.dump(self, path) @staticmethod def load(path): return joblib.load(path)
第二步:主脚本代码
# 分发自定义模块到所有节点 spark.sparkContext.addPyFile("/dbfs/path/to/ensemble_models.py") def Process(df): import traceback import xgboost import pickle import joblib import sys import pandas as pd from ensemble_models import WeightedEnsembleRegressor data = {'msg': [""]} try: # 使用DBFS路径加载模型 ensemble = WeightedEnsembleRegressor.load('/dbfs/path/to/final_model_v3_1.joblib') data['msg'] = [ f"success {sys.version} {joblib.__version__} {pickle.compatible_formats}" ] except Exception as e: # 处理catboost导入可能失败的情况 catboost_version = "not installed" try: import catboost catboost_version = catboost.__version__ except: pass data['msg'] = [ f"fail {sys.version} {joblib.__version__} {pickle.compatible_formats}\n" f"{traceback.format_exc()}\n" f"xgboost: {xgboost.__version__}\n" f"catboost: {catboost_version}" ] return pd.DataFrame.from_dict(data, orient='index').transpose()
内容的提问来源于stack exchange,提问作者user6386155
相关产品推荐
相关产品推荐

