Export Pix2Pix generator to tflite model

Viewed 402

I trained a Pix2Pix generator from the Tensorflow 2.0 tutorial and I exported it in tflite this way :

converter = tf.lite.TFLiteConverter.from_keras_model(generator)
tflite_model = converter.convert()
open("facades.tflite", "wb").write(tflite_model)

Unfortunately, I have problems that seem to come from tf.keras.layers.BatchNormalization when I try to infer it.

First, the result of an inference only returns Nan values. This can be resolved by disabling the fused implementation.

Secondly, the BatchNormalization layer behaves differently depending on whether we are in training or prediction. The tutorial explicitly states to make a prediction in training=True mode. I don't know how to do this with the tflite model.

One solution talks about replacing the BatchNormalization layer by an InstanceNormalization, which can be found in the tensorflow_addons. The conversion to tflite is done without any problem, but there is still a problem with the inference. when I call invoke on the interpreter it crashes by returning me a SEGFAULT. According to the stackcall it would come from SquaredDifference operator of the InstanceNormalization layer.

Has anyone managed to convert this TensorFlow 2.0 model into a tflite and infer it correctly ? How ? Thank you.

PS : I would prefer a solution with BatchNormalization because it is a standard layer in Keras and can therefore also work with TensorFlow javascript.

0 Answers
Related