使用自定义回调,可以绘制适合特定时期所需的总时间。
class timecallback(tf.keras.callbacks.Callback):
def __init__(self):
self.times = []
# use this value as reference to calculate cummulative time taken
self.timetaken = time.clock()
def on_epoch_end(self,epoch,logs = {}):
self.times.append((epoch,time.clock() - self.timetaken))
def on_train_end(self,logs = {}):
plt.xlabel('Epoch')
plt.ylabel('Total time taken until an epoch in seconds')
plt.plot(*zip(*self.times))
plt.show()
然后将其作为回调传递给model.fit函数,如下所示
timetaken = timecallback()
model.fit(train_images, train_labels, epochs=5,callbacks = [timetaken])
如果你想绘制每个历元的时间。您可以用on_epoch_end替换on_train_end方法。
def on_epoch_end(self,epoch,logs= {}):
# same as the on_train_end function