Automating Custom Object Restoration
To automate the mapping of custom objects without manually passing a dictionary to every load_model call, TensorFlow provides the @tf.keras.utils.register_keras_serializable decorator. When a custom layer, loss, or metric is decorated with this attribute, Keras adds it to a global registry, allowing the loader to resolve the object by its name automatically during deserialization.
Implementation via Registration
Instead of maintaining a custom_objects dictionary, apply the decorator to your class definition:
import tensorflow as tf
@tf.keras.utils.register_keras_serializable(package="MyCustomLayers")
class CustomDense(tf.keras.layers.Layer):
def __init__(self, units=32, **kwargs):
super().__init__(**kwargs)
self.units = units
def build(self, input_shape):
self.w = self.add_weight(shape=(input_shape[-1], self.units))
def call(self, inputs):
return tf.matmul(inputs, self.w)
def get_config(self):
config = super().get_config()
config.update({"units": self.units})
return config
Once registered, tf.keras.models.load_model("path/to/model") will resolve CustomDense without additional arguments, provided the class definition is imported into the current execution environment before the load call.
SavedModel Format and Versioning
The SavedModel format handles custom layers differently than the legacy H5 format. While H5 relies heavily on Python class names and get_config(), SavedModel stores the concrete computation graph (via TensorFlow's GraphDef). This means the model can often be loaded for inference (prediction) without the original Python code, as the operations are baked into the graph.
Handling Configuration Incompatibility
When loading for training or fine-tuning, Keras still requires the Python class to reconstruct the layer object. To prevent incompatibility when configurations evolve, follow these standards:
- Strict
get_config() implementation: Always include all hyperparameters in get_config() and ensure the __init__ method can handle those exact keys.
- Default Arguments: Use default values in
__init__ to maintain backward compatibility if new parameters are added to a layer in newer versions of your code.
- Avoid Lambdas: Lambda functions are not reliably serializable; always use named functions or classes.
Verification Steps
To verify that registration is working and the model is portable:
- Save a model containing a registered custom layer using
model.save("my_model").
- In a fresh Python session, import the module containing the decorated class.
- Run
model = tf.keras.models.load_model("my_model") without the custom_objects parameter.
Diagnostic Note: If you are encountering a ValueError despite registration, please specify if you are loading a .h5 file or a SavedModel directory, as the serialization paths differ significantly.