Using Keras Model as a Broadcast variable with Apache Spark & Elephas

Viewed 1990

I have a keras model with pretrained weights [h5df] of about 700mb. I would like to use it with Apache Spark as a broadcast variable. 1. This does not seem to be possible as the keras model itself is not spark aware and not serializable. 2. When googled it a little, I have found Elephas library that does the work. So tried wrapping up the Keras pretrained model in ElephasTransformer. This is throwing multiple errors as ( I use python 2.7 ). For example in the file ml_model.py of Elephas, "from pyspark.ml.util import keyword_only" , the import is not available. Even when I try to comment this out and make appropriate modifications in the code, it seems to be unable to handle the broadcasting of the keras model as it throws a tensor error in one of the dense layers.

How to use a pretrained keras model as a broadcast variable in apache spark ?

2 Answers

I'd distribute the models using SparkFiles

spark.sparkContext.addFile("model_file.h5")

and load locally:

from pyspark import SparkFiles
from keras.models import load_model

def f(it):
    path = SparkFiles.get("mode_file.h5")
    model =  load.model(path)

    for i in it:
        yield ... # Do something


rdd.mapPartitions(f)

In Elephas, the way I approached this problem was making the weights a broadcast variable, providing the yaml string as an argument to the mapper function, and just create the model inside the mapper function using the loaded yaml file and the weights - something to the effect of:

from tensorflow.keras.models import model_from_yaml

weights = rdd.context.broadcast(model.get_weights())

def mapper_function(yaml_file, ...):
    model = model_from_yaml(yaml_file)
    model.set_weights(weights.value)
Related