How to load keras "history" object with custom loss?

Viewed 468

So I defined my keras model and have used a custom_loss function to train the model:

model.compile(optimizer='adam', loss=custom_loss, metrics=[custom_loss])

Then I am training the model:

history = model.fit(X_train, y_train, batch_size=1024, epochs=125, validation_split=0.2, shuffle=True)

Then I save this history object using the following code:

with open('history.pkl', 'wb') as file:  
   pickle.dump(history, file)

Now, when I am trying to read the history object as follows:

with open('history.pkl', 'rb') as file:
    history = pickle.load(file)

I get the following error:

ValueError: Unknown loss function:custom_loss

How can I read the history object? I don't get this error when I am not using custom_loss function. I am using keras 2.2.4 and tensorflow 1.15.5

Edit: Complete error traceback as requested: enter image description here

2 Answers

For most use cases, you don't want to serialize the history object. What you are usually interested in is history.history, which is a dict of the logs / metrics / losses / etc.

Try that:

pickle.dump(history.history, file)

The fuller answer is that the history object returned is a tf.keras.callbacks.History, which subclasses tf.keras.callbacks.Callback. Callback itself has a ref to the model, which then has refs to all kinds of stuff including custom objects like your custom loss. Serialization of Keras custom objects is a whole other big topic... tldr the recommended way to serialize Keras models is not to use pickle.

+1 to @Yaoshiang's answer. That's the right answer.

This is just a big note.

Reading that trace, it looks like keras has custom pickle logic that uses the standard keras save/load logic. keras doesn't know about your loss function unless you tell it.

Search for "custom objects" in this guide: https://keras.io/guides/serialization_and_saving/

Try something like:

custom_objects = {"custom_loss": custom_loss}
with keras.utils.custom_object_scope(custom_objects):
  with open('history.pkl', 'rb') as file:
    history = pickle.load(file)
Related