Spark中广播TensorFlow模型


文中是以PySpark3.1.2 + tf 1.15 的keras模块测试,tf 2若遇到同样问题,解决方法应类似

出于某些原因,想将tf模型部署在pyspark上。

首先想到使用广播变量将tf模型广播到每个executor上,但发现存在问题

model = Sequential()
model.add(Dense(1, input_dim=42, activation='sigmoid'))
model.compile(optimizer='Nadam', loss='binary_crossentropy', metrics=['accuracy'])

pickle.dumps(model)
---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
/tmp/ipykernel_46651/1468016730.py in 
      3 model.compile(optimizer='Nadam', loss='binary_crossentropy', metrics=['accuracy'])
      4 
----> 5 pickle.dumps(model)

TypeError: can't pickle _thread._local objects

pyspark中广播变量的建立依赖于pickle的序列化,但是TensorFlow中的Keras模型不支持pickle。所以暂时放弃这个思路

再想到使用广播变量来广播模型权重,在每个task中建立模型并set_weights

但如果对于每个task单独建立模型,认为性能开销比较大。

发现大佬的解决方法https://github.com/tensorflow/tensorflow/issues/34697#issuecomment-627193883

因为,模型不能广播的主要原因是默认情况下模型不能序列化

所以只需让模型支持序列化即可

方法为修改__reduce__方法

__reduce__方法告诉Python如何对模型进行pickle

方法如下

import pickle

from tensorflow.keras.models import Sequential, Model
from tensorflow.keras.layers import Dense
from tensorflow.python.keras.layers import deserialize, serialize
from tensorflow.python.keras.saving import saving_utils


def unpack(model, training_config, weights):
    restored_model = deserialize(model)
    if training_config is not None:
        restored_model.compile(
            **saving_utils.compile_args_from_training_config(
                training_config
            )
        )
    restored_model.set_weights(weights)
    return restored_model

# Hotfix function
def make_keras_picklable():

    def __reduce__(self):
        model_metadata = saving_utils.model_metadata(self)
        training_config = model_metadata.get("training_config", None)
        model = serialize(self)
        weights = self.get_weights()
        return (unpack, (model, training_config, weights))

    cls = Model
    cls.__reduce__ = __reduce__

# Run the function
make_keras_picklable()

# Create the model
model = Sequential()
model.add(Dense(1, input_dim=42, activation='sigmoid'))
model.compile(optimizer='Nadam', loss='binary_crossentropy', metrics=['accuracy'])

# Save
with open('model.pkl', 'wb') as f:
    pickle.dump(model, f)

主要部分为修改了模型的__reduce__方法,关于__reduce__可参考官方文档

采用返回一个元组的形式

元组中第一个元素是一个函数unpack,用来根据保存的状态 进行模型的还原构建

第二个元素也是个元组,表示模型相关的状态。

至此,模型新增reduce方法后,这个方法作为helper,可以帮助对模型进行pickle

有一说一,在spark的广播变量中,感觉这样也是在每个task中对模型进行构建模型并set_weights。

根据Broadcast的value属性来看,按理说获取value的时候只应该进行一次反序列化,将结果保存在_value中,从而保证每个executor上只反序列化一次模型。实际测试中只看到了driver段的广播变量只读取一次模型,传给worker的广播变量在每个task都进行了反序列化。初步认为这是因为worker是由daemon中fork而来,只对子进程中广播变量的value进行了存储。这部分还得再看一看。

所以实验中提升不大,尴尬,两种方法中对模型进行pickle的话快了一丁丁点

我这边也稍稍修改了下,无关大局

class MyModel(Model):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
    
    @staticmethod
    def unpack(model, training_config, weights):
        restored_model = deserialize(model, custom_objects=None)
        if training_config is not None:
            restored_model.compile(
                **saving_utils.compile_args_from_training_config(
                    training_config
                )
            )
        restored_model.set_weights(weights)
        return restored_model


    def __reduce__(self):
        model_metadata = saving_utils.model_metadata(self)
        training_config = model_metadata.get("training_config", None)
        model = serialize(self)
        weights = self.get_weights()
        return (MyModel.unpack, (model, training_config, weights))

注:模型别忘了单独放到一个包中,否则出现import报错,根据pyspark原理,调用了daemon管理worker进行具体计算,没单独放的话会出现daemon中找不到这个Model的问题。