diff --git a/ttools/callbacks.py b/ttools/callbacks.py index 6fd74e0..375bb4d 100644 --- a/ttools/callbacks.py +++ b/ttools/callbacks.py @@ -158,7 +158,7 @@ def __init__(self, keys=None, val_keys=None, smoothing=0.999): self.keys = keys if val_keys is None: - self.val_keys = [] + self.val_keys = [] else: self.val_keys = val_keys @@ -184,7 +184,7 @@ class VisdomLoggingCallback(KeyedCallback): 0.0 disables smoothing. """ - def __init__(self, keys=None, val_keys=None, frequency=100, server=None, + def __init__(self, keys=None, val_keys=None, frequency=100, server=None, port=8097, base_url="/", env="main", log=False, smoothing=0.99): super(VisdomLoggingCallback, self).__init__( keys=keys, val_keys=val_keys, smoothing=smoothing) @@ -460,14 +460,23 @@ class CheckpointingCallback(Callback): PERIODIC_PREFIX = "periodic_" EPOCH_PREFIX = "epoch_" + BEST_MODEL_FILENAME = "best" def __init__(self, checkpointer, interval=600, - max_files=5, max_epochs=10): + max_files=5, max_epochs=10, + best_val_key=None, best_val_value=None): super(CheckpointingCallback, self).__init__() self.checkpointer = checkpointer self.interval = interval self.max_files = max_files self.max_epochs = max_epochs + self.best_val_key = best_val_key + + if best_val_value is not None: + LOG.info("Loaded best model ({}={})".format(best_val_key, best_val_value)) + self.best_val_value = best_val_key + else: + self.best_val_value = float('inf') self.last_checkpoint_time = time.time() @@ -506,6 +515,22 @@ def batch_end(self, batch_data, train_step_data): self.checkpointer.save(filename, extras={"epoch": self.epoch}) self.__purge_old_files() + def validation_end(self, val_data): + """Save a best model checkpoint if value for best_val_key is lowest so far.""" + + super(CheckpointingCallback, self).validation_end(val_data) + + if self.best_val_key is None: + return + if val_data[self.best_val_key] > self.best_val_value: + return + + self.best_val_value = val_data[self.best_val_key] + + LOG.debug("Best model checkpoint ({}={})".format(self.best_val_key, self.best_val_value)) + self.checkpointer.save(CheckpointingCallback.BEST_MODEL_FILENAME, + extras={"epoch": self.epoch, "best_val_value": self.best_val_value}) + def __purge_old_files(self): """Delete checkpoints that are beyond the max to keep.""" @@ -635,7 +660,7 @@ def training_end(self): print("end logging experiment", self.epoch, self.batch) def _get_commit(self): - return subprocess.check_output(["git", "rev-parse", "HEAD"]) + return subprocess.check_output(["git", "rev-parse", "HEAD"]) class CSVLoggingCallback(KeyedCallback): diff --git a/ttools/training.py b/ttools/training.py index f0115e3..bdb6739 100644 --- a/ttools/training.py +++ b/ttools/training.py @@ -69,7 +69,7 @@ def training_step(self, batch): This should implement a forward pass of the model, compute gradients, take an optimizer step and return useful metrics and tensors for - visualization and training callbacks. + visualization and training callbacks. Args: batch (dict): batch of data provided by a data pipeline. @@ -168,7 +168,7 @@ def train(self, dataloader, starting_epoch=None, num_epochs=None, if starting_epoch is None: starting_epoch = 0 - LOG.info("Starting taining from epoch %d", starting_epoch) + LOG.info("Starting training from epoch %d", starting_epoch) epoch = starting_epoch