How can I save a Tensorflow 2.2.0 model with a custom training loop?

Viewed 733

I am struggling to save a tf.keras model to easily load and be able to use it. I have used the tf.keras.Model subclass method to construct a MLP model with a custom loss function, as you can see below:

class MyModel(tf.keras.Model):
    def __init__(self):
        super(MyModel, self).__init__()
        self.dense1 = Dense(400, activation='relu', kernel_initializer=initializers.glorot_uniform(), input_dim=5)
        self.dense2 = Dense(400, activation='relu', kernel_initializer=initializers.glorot_uniform())
        self.dense3 = Dense(400, activation='relu', kernel_initializer=initializers.glorot_uniform())
        self.dense4 = Dense(400, activation='relu', kernel_initializer=initializers.glorot_uniform())
        self.dense_out = Dense(1, activation='relu', kernel_initializer=initializers.glorot_uniform())

    @tf.function(input_signature=[tf.TensorSpec(shape=(None, 5), dtype=tf.float32, name='inputs')])   #CHECK tf.saved_model.save docs!
    def call(self, inputs, **kwargs):
        x = self.dense1(inputs)
        x = self.dense2(x)
        x = self.dense3(x)
        x = self.dense4(x)
        return self.dense_out(x)

    def get_loss(self, X, Y):
        with tf.GradientTape() as tape:
            tape.watch(tf.convert_to_tensor(X))
            Y_pred = self.call(X)
        return tf.reduce_mean(tf.math.square(Y_pred-Y)) + tf.reduce_mean(tf.maximum(0, tape.gradient(Y_pred, X)[:, 2]))

    def get_grad_and_loss(self, X, Y):
        with tf.GradientTape() as tape:
            tape.watch(tf.convert_to_tensor(X))
            L = self.get_loss(X, Y)
        g = tape.gradient(L, self.trainable_weights)
        return g, L

I then make an instance of the model and proceed with a standard training loop:

model = MyModel()
epochs = 5
optimizer = tf.keras.optimizers.Adam(learning_rate=1e-5, beta_1=0.9, beta_2=0.999)
x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.2)
x_train, x_val, y_train, y_val = train_test_split(x_train, y_train, test_size=0.25)
train_dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train))
train_dataset = train_dataset.shuffle(buffer_size=1024).batch(batch)
val_dataset = tf.data.Dataset.from_tensor_slices((x_val, y_val))
val_dataset = val_dataset.shuffle(buffer_size=1024).batch(batch)
val_acc_metric = tf.keras.metrics.MeanAbsoluteError()


## TRAINING LOOP
losses = []
for epoch in range(epochs):
    print(f'############ START OF EPOCH {epoch + 1} ################')
    for step, (x_batch_train, y_batch_train) in enumerate(train_dataset):
        grads, L = model.get_grad_and_loss(x_batch_train, y_batch_train)
        losses.append(float(L))
        optimizer.apply_gradients(zip(grads, model.trainable_weights))

        if step % 100 == 0:
            print('Training loss (for one batch) at step %s: %s' % (step, float(L)))
            print(f'Seen so far: {(step+1)*batch} samples')

    # Run a validation loop at the end of each epoch.
    for val_step, (x_batch_val, y_batch_val) in enumerate(val_dataset):
        val_logits = model.call(x_batch_val)
        # Update val metrics
        val_acc_metric(y_batch_val, val_logits)
    val_acc = val_acc_metric.result()
    val_acc_metric.reset_states()
    print(f'Validation acc: {val_acc}')

I have tried to follow the steps outlined here. I call the model on a random input in order to trigger model.build internally, and then I attempt to save the model using the following:

model.save('mymodel', signatures=model.call.get_concrete_function([tf.TensorSpec(shape=(None, 5), dtype=tf.float32, name='inputs')]))

I then get this following error:

Traceback (most recent call last):
File "<input>", line 1, in <module>
File "/Users/Maximocravero/opt/miniconda3/envs/finance_research/lib/python3.8/site- 
packages/tensorflow/python/eager/def_function.py", line 959, in get_concrete_function
concrete = self._get_concrete_function_garbage_collected(*args, **kwargs)
File "/Users/Maximocravero/opt/miniconda3/envs/finance_research/lib/python3.8/site- 
packages/tensorflow/python/eager/def_function.py", line 871, in 
_get_concrete_function_garbage_collected
return self._stateless_fn._get_concrete_function_garbage_collected(  # pylint: 
disable=protected-access
File "/Users/Maximocravero/opt/miniconda3/envs/finance_research/lib/python3.8/site- 
packages/tensorflow/python/eager/function.py", line 2480, in 
_get_concrete_function_garbage_collected
raise ValueError("Structure of Python function inputs does not match "
ValueError: Structure of Python function inputs does not match input_signature.

I don't understand this issue as I specify the same TensorSpec as in the tf.function above the model.call attribute. I attempted this without including the tf.function above the model call, which leads to an error relating to the input dimensions having to be set. I am able to address this by calling the model on an arbitrary input, which does allow me to save the model but I have to compile it prior to using it and I get the following warning:

WARNING:tensorflow:From 
/Users/Maximocravero/opt/miniconda3/envs/finance_research/lib/python3.8/site- 
packages/tensorflow/python/ops/resource_variable_ops.py:1813: calling 
BaseResourceVariable.__init__ (from tensorflow.python.ops.resource_variable_ops) with 
constraint is deprecated and will be removed in a future version.
Instructions for updating:
If using Keras pass *_constraint arguments to layers.

My question is whether or not I am completely missing something or if it is usual to have to compile loaded custom models? I am running TensorFlow 2.2.0 with Python 3.8.2, and based on the documentation for saving models this should really be quite simple. I am new to TensorFlow so it may well be that it's a silly mistake, but ultimately it's still a basic model with 5 inputs and a single output. Any help would be greatly appreciated.

0 Answers
Related