Saving and loading a tf.keras model that uses custom models as layers

Viewed 142

I've created a model architecture that uses two layers of custom keras models. I can train and predict with it and save it, but cannot load it.

Originally, I used custom tf.keras.layers.Layer sub-classes. I was able to train, predict, save, reload, etc. But I want the ability to access the weights and biases of all layers, e.g. when I used tf.keras.layers.LSTM two levels down from the actual model. So I switched to using tf.keras.Model as the parent class. Now in an interactive kernel I can access the weights etc. of a trained model.

Here is how it is being saved, pretty standard (apologies, I can't share too much code):

self.model.save(model_path, save_format="h5")
self.model.save_weights(model_weights_path)

I am facing errors when trying to reload the model. I've defined all the necessary custom objects.

self.model = tf.keras.models.load_model(
    self.model,
    custom_objects=self.custom_objects)

This worked using custom tf.keras.layers.Layer classes, but not anymore.

I seems to be that in the config for the lowest custom models (that calls actual layer classes), there is a 'layers' key class in the config. However, in the intermediate custom models, there is no 'layers' key.

Traceback:


  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\saving\save.py", line 146, in load_model
    return hdf5_format.load_model_from_hdf5(filepath, custom_objects, compile)

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\saving\hdf5_format.py", line 169, in load_model_from_hdf5
    custom_objects=custom_objects)

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\saving\model_config.py", line 55, in model_from_config
    return deserialize(config, custom_objects=custom_objects)

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\layers\serialization.py", line 108, in deserialize
    printable_module_name='layer')

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\utils\generic_utils.py", line 303, in deserialize_keras_object
    list(custom_objects.items())))

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\engine\network.py", line 937, in from_config
    config, custom_objects)

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\engine\network.py", line 1895, in reconstruct_from_config
    process_layer(layer_data)

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\engine\network.py", line 1875, in process_layer
    layer = deserialize_layer(layer_data, custom_objects=custom_objects)

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\layers\serialization.py", line 108, in deserialize
    printable_module_name='layer')

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\utils\generic_utils.py", line 303, in deserialize_keras_object
    list(custom_objects.items())))

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\engine\network.py", line 937, in from_config
    config, custom_objects)

  File "C:\Users\<>\AppData\Local\Continuum\anaconda3\envs\<>\lib\site-packages\tensorflow_core\python\keras\engine\network.py", line 1894, in reconstruct_from_config
    for layer_data in config['layers']:

KeyError: 'layers'
0 Answers
Related