[Current Behavior]
Tensorflow Object Detection Model is trained and performs well. It is then saved to checkpoint and re-loaded from that checkpoint. After re-loading the saved checkpoint, model no longer detects anything.
[Desired Behavior]
Model that is re-loaded from checkpoint should produce the same or very similar output to model that was used to produce checkpoint.
[Minimal code to reproduce error]
Google Colab notebook prepared which reproduces problem can be found here:
[More detailed description of the problem and what the minimal code does]
I am using Tensorflow's object detection models to train a model to detect various things in microscope images.
A model that trains well gives bad performance when re-loaded from a checkpoint/saved model format. I am training using an eager mode training loop adapted from their example code.
When inference using the model that has just been trained and has not left memory, performance is pretty great. However when I create a new model in inference mode and load the saved checkpoint - model no longer detects anything.
When I train using the TF Object Detection model_main_tf2.py command line method instead of my custom eager loop, I experience the same outcome. Model appears to be training well and loss gets very low, however when loading a model from the checkpoint that is produced, performance is terrible.
I have created a minimal example on Colab (linked above), where I have made minimal changes to their example code.
It overfits to tensorflow's small duck dataset and performs inference on the same data as training.
When run, this notebook will train a SSD model by overfitting to the small duck dataset. Inference is then performed on the same data as training.
After training, the notebook will plot two sets of images:
- Duck dataset images with overlayed results from the in-memory model.
- Duck dataset images with overlayed results from a model that is created and loaded from the most recent training checkpoint (the checkpoint produced after training has finished).
The former has very high performance, typically, everything is detected with very high confidence. The latter will detect either nothing at all (<1e-3 confidences) or it will seem to detect the object but also a bunch of other junk (Note: In rare cases, the latter actually works well. It if works just reset the notebook and run again. There seems to be some variation and in some cases it can 'accidently' perform well)
The effect is especially bad when batchnorm is trainable. Models trained with batchnorm very quickly converge to perform amazingly with a model that is not re-loaded, however re-loaded models perform extremely poor. Effect is also present when trained without batchnorm but to a lesser degree.
Currently my suspicion is batchnorm being an issue somehow - but I cannot find anything in the internet about anyone else encountering this type of issue?