tf.keras.Model training accuracy jumps after running Model.evaluate() during training

Viewed 164

I'm training a TensorFlow Keras Model using Model.fit(). I'm also using callbacks to log my training accuracy metrics after every batch using TensorFlow's on_train_batch_end() syntax. In addition, I'm using another callback to run Model.evaluate() every 1,000 batches to compute validation set accuracy and update the logs dict passed around the callbacks during Model.fit().

Looking at the logged metrics vs. batch number shows very perplexing results. After the Model.evaluate() run, the training accuracy experiences a significant 'jolt', initially triggering a rapid increase in the logged training accuracy and subsequently triggering a significant drop training accuracy followed by a slower recovery (see attached images).

My guess is that it's something to do with the Model.evaluate()'s call to reset_metrics(), which loops through and calls the reset_states() method on each metric. I can't work out what reset_states() is doing and if this is relevant to the behaviour I'm observing. It seems to relate to the Mean parent class of CategoricalAccuracy. I haven't been able to find anything helpful in the TensorFlow docs yet.

Are the metrics shown during Model.fit() actually some form of moving averages rather than the batch-wise metric? In that case, the reset_states() method would be resetting the moving average, possibly producing the jolting behaviour.

Can anyone with a better grasp of TensorFlow's inner workings help?

Jolting accuracy.

Jolting accuracy (zoomed in).

0 Answers
Related