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

如何在Palantir Foundry模型集成仓库发布pyspark.ml.PipelineModel

发布PySpark PipelineModel到Palantir Foundry时出现"cannot pickle 'socket' object"错误

我需要在Palantir Foundry的模型集成代码仓库发布一个pyspark.ml.PipelineModel对象,官方监督学习教程仅针对sklearn模型与pandas DataFrame,我已将逻辑适配为Spark DataFrame输入,训练代码在工作簿中可正常运行,但发布时触发cannot pickle 'socket' object错误,推测问题出在模型适配器代码中。

相关代码

测试数据集

在代码工作簿中创建的测试数据集:

def toy_data():
    import pandas as pd
    pandas_df = pd.DataFrame({
        "x1":["a", "a", "a", "a", "a", "a", "a", "a", "a", "a", "b", "b", "b", "b", "b", "b", "b", "b", "b", "b"],
        "x2":["z", "v", "z", "v", "z", "v", "z", "v", "z", "v", "z", "v", "z", "v", "z", "v", "z", "v", "z", "v"],
        "y":[1, 3, 2, 3, 2, 1, 2, 3, 4, 3, 11, 17, 13, 12, 17, 22, 21, 7, 14, 10]
    })
    spark_df = spark.createDataFrame(pandas_df)
    return spark_df

模型训练逻辑(model_training.py)

from transforms.api import transform, Input
from palantir_models.transforms import ModelOutput
from main.model_adapters.adapter import ExampleModelAdapter

@transform(
    training_data_input=Input("INPUT_PATH"),
    model_output=ModelOutput("MODEL_OUTPUT_PATH"),
)
def compute(training_data_input, model_output):
    training_df = training_data_input.dataframe()
    model = train_model(training_df)
    foundry_model = ExampleModelAdapter(model)
    model_output.publish(model_adapter=foundry_model)

def train_model(training_df):
    from pyspark.ml.feature import StringIndexer, VectorAssembler
    from pyspark.ml.pipeline import Pipeline
    from pyspark.ml.regression import DecisionTreeRegressor
    from pyspark.ml.evaluation import RegressionEvaluator
    from pyspark.ml.tuning import CrossValidator, ParamGridBuilder
    train_set = training_df
    cat_vars = ['x1', 'x2']
    indexer = StringIndexer(inputCols=cat_vars, outputCols=[var + "_indexed" for var in cat_vars], handleInvalid='keep')
    assembler = VectorAssembler(inputCols=[var + "_indexed" for var in cat_vars], outputCol='features')
    decision_tree = DecisionTreeRegressor(featuresCol='features', labelCol='y', seed=123)
    regression_evaluator = RegressionEvaluator(labelCol='y')
    pipeline=Pipeline(stages=[indexer, assembler, decision_tree])
    params = (ParamGridBuilder().addGrid(param=decision_tree.maxDepth, values=[1, 2]).addGrid(param=decision_tree.minInstancesPerNode, values=[1, 2, 3]).build())
    cv = CrossValidator(estimator=pipeline, estimatorParamMaps=params, evaluator=regression_evaluator, numFolds=2, seed=477)
    cv_trees = cv.fit(train_set)
    best_model = cv_trees.bestModel
    return best_model

原模型适配器(adapter.py)

import palantir_models as pm
from palantir_models_serializers import DillSerializer

class ExampleModelAdapter(pm.ModelAdapter):
    @pm.auto_serialize(model=DillSerializer())
    def __init__(self, model):
        self.model = model
    @classmethod
    def api(cls):
        inputs = {"df_in": pm.Spark()}
        outputs = {"df_out": pm.Spark()}
        return inputs, outputs
    def predict(self, df_in):
        return self.model.transform(df_in)

错误输出

发布时触发的错误:

[module version: 1.1008.0]
TypeError: cannot pickle 'socket' object
...

依赖版本

已列出Python、transforms、palantir_models等相关库的版本信息。

解决方案

问题原因

pyspark.ml.PipelineModel内部包含与Spark上下文绑定的对象(如socket连接),无法通过DillSerializer这类通用序列化工具持久化,必须使用Spark原生的模型保存/加载机制。

修改后的模型适配器(adapter.py)

替换原适配器代码,实现自定义的save和load方法,利用Spark原生API处理模型持久化:

import palantir_models as pm
from pyspark.ml.pipeline import PipelineModel

class ExampleModelAdapter(pm.ModelAdapter):
    def __init__(self, model):
        self.model = model

    @classmethod
    def api(cls):
        inputs = {"df_in": pm.Spark()}
        outputs = {"df_out": pm.Spark()}
        return inputs, outputs

    def predict(self, df_in):
        return self.model.transform(df_in)

    @classmethod
    def load(cls, artifact_context):
        # 从Foundry指定路径加载模型
        model_path = artifact_context.artifact_uri("spark_model")
        model = PipelineModel.load(model_path)
        return cls(model)

    def save(self, artifact_context):
        # 将模型保存到Foundry指定路径
        model_path = artifact_context.artifact_uri("spark_model")
        self.model.save(model_path)

验证步骤

  1. 替换原适配器代码为上述版本
  2. 重新运行模型发布流程
  3. 验证模型可正常加载并对Spark DataFrame做出预测

内容的提问来源于stack exchange,提问作者Chris A.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 17:00:53