You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 13:25:03