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

如何结合PySpark实现猴子补丁式Keras模型的Pickle序列化

我之前刚好踩过这个坑,把Keras模型放到Spark Worker的UDF里跑,核心就是解决模型的Pickle序列化问题,结合猴子补丁和Spark的广播机制就能搞定,给你一步步拆解:

核心思路

Keras原生Model类不支持Pickle序列化,所以我们用猴子补丁给Model类添加__getstate__和__setstate__方法,让它能被序列化;然后通过Spark的广播变量把序列化后的模型分发到所有Worker节点;最后在UDF里反序列化模型并调用,同时用全局变量缓存模型避免重复加载。

具体步骤

1. 给Keras Model打Pickle支持的猴子补丁

这个补丁会让模型在序列化时保存结构(YAML格式)和权重,反序列化时重建模型并加载权重:

import io
import pickle
import yaml
from keras.models import Model

def make_keras_model_picklable():
    # 序列化时保存模型结构和权重
    def __getstate__(self):
        model_str = ""
        with io.StringIO() as stream:
            yaml.dump(self.to_yaml(), stream)
            model_str = stream.getvalue()
        return {"model_str": model_str, "weights": self.get_weights()}

    # 反序列化时重建模型并加载权重
    def __setstate__(self, state):
        from keras.models import model_from_yaml
        model = model_from_yaml(state["model_str"])
        model.set_weights(state["weights"])
        self.__dict__.update(model.__dict__)

    # 给Model类绑定这两个方法
    Model.__getstate__ = __getstate__
    Model.__setstate__ = __setstate__

# 必须在加载Keras模型之前调用这个补丁函数
make_keras_model_picklable()

2. 在Driver端加载并序列化模型

补丁生效后,你的Keras模型就能正常被Pickle序列化了:

from keras.models import load_model

# 加载预训练好的Keras模型
model = load_model("your_trained_model.h5")

# 序列化模型为字节流
model_pkl = pickle.dumps(model)

3. 广播序列化后的模型到所有Worker

用Spark的广播变量高效分发模型,避免每个Task都重复传输:

from pyspark.sql import SparkSession

spark = SparkSession.builder.appName("KerasSparkIntegration").getOrCreate()
# 广播序列化后的模型
broadcast_model = spark.sparkContext.broadcast(model_pkl)

4. 定义调用模型的UDF

这里关键是用全局变量缓存模型,每个Worker节点只反序列化一次,避免重复加载浪费资源:

import numpy as np
from pyspark.sql.functions import udf
from pyspark.sql.types import FloatType  # 根据你的模型输出类型调整

def predict_with_input(input_features):
    # 用函数属性缓存模型,每个Worker只会初始化一次
    if not hasattr(predict_with_input, "model"):
        # 从广播变量中获取序列化模型并反序列化
        model_pkl = broadcast_model.value
        predict_with_input.model = pickle.loads(model_pkl)
    
    model = predict_with_input.model
    # 把输入数据转换成模型需要的格式(这里示例是单样本输入,根据你的模型调整)
    input_array = np.array([input_features])
    # 执行预测
    prediction = model.predict(input_array)[0][0]
    return float(prediction)

# 注册UDF,指定输出数据类型
predict_udf = udf(predict_with_input, FloatType())

5. 在DataFrame上应用UDF

现在就可以把UDF用到你的DataFrame列上了:

# 假设你的DataFrame有一个名为"input_features"的列,是模型的输入
df = df.withColumn("prediction", predict_udf(df["input_features"]))
# 查看结果
df.show()
关键注意事项
  • 依赖版本一致:所有Spark Worker节点必须安装和Driver完全相同版本的Keras、TensorFlow、NumPy等库,版本不兼容会导致模型加载失败。
  • 输入格式匹配:UDF里要确保输入数据的维度、类型和模型训练时的输入一致,比如如果模型接受的是(None, 10)的输入,你要把DataFrame里的特征转换成对应形状的数组。
  • 资源配置:如果模型很大,需要调整Spark的广播变量相关配置(比如spark.driver.maxResultSize),避免内存溢出。

内容的提问来源于stack exchange,提问作者Erp12

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 08:33:26