[Keras] Deserializion failed
- Dominant language
- Python
- Stars
- 1.7k
- Forks
- 778
- Avg merge
- 2h 50m
- Merged PRs (30d)
- 3
Description
**What happened**:
Deserializion of a simple Keras model failed on distributed mode of dask
**What you expected to happen**:
As deserialize_keras_model and serialize_keras_model works locally I was expecting to be able to embed it into a delayed function.
Am I missing something ?
**Minimal Complete Verifiable Example**:
```python
import numpy as np
import dask
import distributed
from distributed import Client
from distributed.protocol.keras import deserialize_keras_model, serialize_keras_model
import keras
def build_model():
inp = keras.layers.Input(shape=(1,))
out = keras.layers.Dense(1, activation='linear')(inp)
model = keras.models.Model(inputs=inp, outputs=out)
return model
@dask.delayed
def get_keras_model_delayed():
return build_model()
def get_keras_model():
return build_model()
client = Client()
data = np.random.randint(0, 100, size=(15, 1))
model = get_keras_model()
headers, frames = serialize_keras_model(model)
deserialized_model = deserialize_keras_model(headers, frames)
assert (deserialized_model.predict(data) == model.predict(data)).all()
model = get_keras_model_delayed().compute() # this line raise an exception
```
Stacktrace:
```
2022-08-25 16:13:35,064 - distributed.protocol.core - CRITICAL - Failed to deserialize
Traceback (most recent call last):
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\distributed\protocol\core.py", line 158, in loads
return msgpack.loads(
File "msgpack\_unpacker.pyx", line 194, in msgpack._cmsgpack.unpackb
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\distributed\protocol\core.py", line 138, in _decode_default
return merge_and_deserialize(
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\distributed\protocol\serialize.py", line 497, in merge_and_deserialize
return deserialize(header, merged_frames, deserializers=deserializers)
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\distributed\protocol\serialize.py", line 426, in deserialize
return loads(header, frames)
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\distributed\protocol\serialize.py", line 59, in dask_loads
return loads(header["sub-header"], frames)
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\distributed\protocol\keras.py", line 41, in deserialize_keras_model
model = model_from_config(header)
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\keras\saving\model_config.py", line 51, in model_from_config
return deserialize(config, custom_objects=custom_objects)
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\keras\layers\serialization.py", line 205, in deserialize
return generic_utils.deserialize_keras_object(
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\keras\utils\generic_utils.py", line 679, in deserialize_keras_object
deserialized_obj = cls.from_config(
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\keras\engine\training.py", line 2720, in from_config
inputs, outputs, layers = functional.reconstruct_from_config(
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\keras\engine\functional.py", line 1312, in reconstruct_from_config
if process_node(layer, node_data):
File "C:\Users\FabienAulaire\workspace\jupyter\venv_jupyter\lib\site-packages\keras\engine\functional.py", line 1209, in process_node
input_data = input_data.as_list()
AttributeError: 'str' object has no attribute 'as_list'
---------------------------------------------------------------------------
```
**Environment**:
- Dask version: 2022.8.1
- Python version: 3.10.5
- Operating System: Windows
- Install method (conda, pip, source): pip
Contributor guide
Assessment
This issue has not been assessed yet.