diff --git a/flyvis/solver.py b/flyvis/solver.py index cccc40d..917f035 100644 --- a/flyvis/solver.py +++ b/flyvis/solver.py @@ -7,7 +7,7 @@ import numpy as np import torch -from datamate import Directory, Namespace +from datamate import Directory, Namespace, namespacify from toolz import valfilter, valmap from torch import nn @@ -15,6 +15,7 @@ from flyvis.network.network import Network from flyvis.task.tasks import Task from flyvis.utils.chkpt_utils import ( + best_checkpoint_default_fn, recover_decoder, recover_network, recover_optimizer, @@ -23,11 +24,35 @@ ) from flyvis.utils.tensor_utils import asymmetric_weighting +try: + from datamate import delete_if_exists as delete_existing_directory +except ImportError: # datamate < 1.0 + from datamate.directory import delete_if_exists as delete_existing_directory + logging = logging.getLogger(__name__) __all__ = ["MultiTaskSolver", "Penalty", "HyperParamScheduler"] +def open_legacy_network_dir(name: str, config: dict) -> Optional[NetworkDir]: + """Open a network directory that stores delete_if_exists in its config. + + Earlier versions passed delete_if_exists in the config, so directories created + by them store it and no longer match the config. Returns the directory if the + flag is the only difference, else None. + """ + existing = NetworkDir(name) + stored = existing.config + if ( + stored is not None + and "delete_if_exists" in stored + and stored.without("delete_if_exists") + == Namespace({"type": stored.type, **namespacify(config)}) + ): + return existing + return None + + class SolverProtocol(Protocol): """SolverProtocol implements training, testing, checkpointing etc. of networks.""" @@ -106,9 +131,21 @@ def __init__( ) -> None: name = name or config["network_name"] assert isinstance(name, str), "Provided name argument is not a string." - self.dir = NetworkDir( - name, {**(config or {}), **dict(delete_if_exists=delete_if_exists)} - ) + # Pass the deletion flag as context, not in the config: datamate stores it + # with a new directory but drops it when opening an existing one, so the + # configs would never match again. + config = {**(config or {})} + config.pop("delete_if_exists", None) + if delete_if_exists: + with delete_existing_directory(): + self.dir = NetworkDir(name, config) + else: + try: + self.dir = NetworkDir(name, config) + except FileExistsError: + self.dir = open_legacy_network_dir(name, config) + if self.dir is None: + raise self.path = self.dir.path @@ -283,12 +320,20 @@ def train(self, overfit: bool = False, initial_checkpoint: bool = True) -> None: logging.info("Training for %s epochs.", n_epochs) logging.info("Checkpointing every %s epochs.", chkpt_every_epoch) - # Initialize data structures to store the loss and activity over iterations. - loss_over_iters = [] - activity_over_iters = [] - activity_min_over_iters = [] - activity_max_over_iters = [] - loss_per_task = {f"loss_{task}": [] for task in self.task.dataset.tasks} + # Initialize data structures to store the loss and activity over iterations, + # starting from the history of earlier calls or of a recovered checkpoint. + def history(key): + if key in self.dir: + return list(self.dir[key][: self.iteration]) + return [] + + loss_over_iters = history("loss") + activity_over_iters = history("activity") + activity_min_over_iters = history("activity_min") + activity_max_over_iters = history("activity_max") + loss_per_task = { + f"loss_{task}": history(f"loss_{task}") for task in self.task.dataset.tasks + } start_time = time.time() with self.task.dataset.augmentation(augment): @@ -420,7 +465,7 @@ def checkpoint(self) -> None: ``` """ self._last_chkpt_ind += 1 - self._curr_chkpt_ind += 1 + self._curr_chkpt_ind = self._last_chkpt_ind # Tracking of validation loss and training batch loss. logging.info("Test on validation data.") @@ -588,35 +633,46 @@ def recover( decoder: Recover decoder parameters. optimizer: Recover optimizer parameters. penalty: Recover penalty parameters. - checkpoint: Index of the checkpoint to recover. + checkpoint: "best", or the position of the checkpoint in the sorted + list of checkpoints, e.g. -1 for the last one. validation_subdir: Name of the subdir to base the best checkpoint on. loss_file_name: Name of the loss to base the best checkpoint on. strict: Whether to load the state dict of the decoders strictly. force: Force recovery of checkpoint if _curr_chkpt_ind is already the same as the checkpoint index. """ - checkpoints = resolve_checkpoints( - self.dir, checkpoint, validation_subdir, loss_file_name - ) + checkpoints = resolve_checkpoints(self.dir) - if checkpoint.index is None or not any((network, decoder, optimizer, penalty)): + if not checkpoints.paths or not any((network, decoder, optimizer, penalty)): logging.info("No checkpoint found. Continuing with initialized parameters.") return - if checkpoints.index == self._curr_chkpt_ind and not force: + if checkpoint == "best": + path = best_checkpoint_default_fn( + self.dir.path, validation_subdir, loss_file_name + ) + position = [p.name for p in checkpoints.paths].index(path.name) + else: + position = checkpoint + path = checkpoints.paths[position] + index = checkpoints.indices[position] + + if index == self._curr_chkpt_ind and not force: logging.info("Checkpoint already recovered.") return # Set the current and last checkpoint index. New checkpoints incrementally # increase the last checkpoint index. self._last_chkpt_ind = checkpoints.indices[-1] - self._curr_chkpt_ind = checkpoints.index + self._curr_chkpt_ind = index # Load checkpoint data. - state_dict = torch.load(checkpoints.path) - logging.info(f"Checkpoint {checkpoints.path} loaded.") + state_dict = torch.load(path, weights_only=False) + logging.info(f"Checkpoint {path} loaded.") - self.iteration = state_dict.get("iteration", None) + # Checkpoints store the last completed iteration, self.iteration - 1. + iteration = state_dict.get("iteration", None) + self.iteration = 0 if iteration is None else iteration + 1 if "scheduler" in self._initialized: # Set the scheduler to the right iteration. diff --git a/flyvis_cli/training/train_single.py b/flyvis_cli/training/train_single.py index 8461e9b..eacbe7b 100644 --- a/flyvis_cli/training/train_single.py +++ b/flyvis_cli/training/train_single.py @@ -70,6 +70,12 @@ def main(args): ) logging.info("Initialized solver with NetworkDir at %s.", solver.dir.path) + if solver.checkpoints and not args.resume: + raise FileExistsError( + f"{solver.dir.path} already has checkpoints. Pass resume=true to continue " + "training or delete_if_exists=true to start over." + ) + if args.get("save_environment", False): save_env(solver.dir.path) diff --git a/tests/test_solver.py b/tests/test_solver.py index 7e0ad59..4b4b9c3 100644 --- a/tests/test_solver.py +++ b/tests/test_solver.py @@ -1,6 +1,9 @@ +import numpy as np import pytest +import torch from datamate import set_root_context +from flyvis.network.directories import NetworkDir from flyvis.solver import MultiTaskSolver from flyvis.utils.config_utils import get_default_config @@ -54,3 +57,89 @@ def test_solver_overfit(solver): solver.train(overfit=True) loss = solver.dir.loss[:] assert loss[-1] < loss[0] + + +def small_config(mock_sintel_data, n_iters): + return get_default_config( + path="../../flyvis/config/solver.yaml", + overrides=[ + "task_name=flow", + "ensemble_and_network_id=0", + f"task.n_iters={n_iters}", + f"+task.dataset.sintel_path={str(mock_sintel_data)}", + "task.original_split=false", + "task.dataset.boxfilter.extent=1", + "task.dataset.n_frames=4", + "task.dataset.dt=0.041", + "task.batch_size=2", + "network.connectome.extent=1", + "scheduler.chkpt_every_epoch=1", + ], + ) + + +def test_solver_continue_training_keeps_history(mock_sintel_data, tmp_path): + with set_root_context(str(tmp_path)): + solver = MultiTaskSolver("test", small_config(mock_sintel_data, n_iters=2)) + solver.train() + loss = solver.dir.loss[:] + assert len(loss) == solver.iteration == 2 + + solver.task.n_iters = 4 + solver.train() + assert len(solver.dir.loss[:]) == solver.iteration == 4 + np.testing.assert_array_equal(solver.dir.loss[:2], loss) + + +def test_solver_recover(mock_sintel_data, tmp_path): + config = small_config(mock_sintel_data, n_iters=2) + with set_root_context(str(tmp_path)): + solver = MultiTaskSolver("test", config) + solver.train() + state = {k: v.clone() for k, v in solver.network.state_dict().items()} + + resumed = MultiTaskSolver("test", config) + resumed.recover(checkpoint=-1) + assert resumed.iteration == solver.iteration + assert resumed._last_chkpt_ind == solver._last_chkpt_ind + for key, value in resumed.network.state_dict().items(): + torch.testing.assert_close(value, state[key]) + + # the initial checkpoint was made before the first iteration + resumed.recover(checkpoint=0) + assert resumed.iteration == 0 + + resumed.recover(checkpoint="best") + assert resumed._curr_chkpt_ind in solver.checkpoints + + resumed.task.n_iters = 4 + resumed.train() + assert len(resumed.dir.loss[:]) == resumed.iteration == 4 + + fresh = MultiTaskSolver("test", config, delete_if_exists=True) + assert fresh.checkpoints == [] + assert "loss" not in fresh.dir + + +def test_solver_recover_older_checkpoint_then_checkpoint(mock_sintel_data, tmp_path): + with set_root_context(str(tmp_path)): + solver = MultiTaskSolver("test", small_config(mock_sintel_data, n_iters=2)) + solver.train() + assert solver.checkpoints == [0, 1] + + solver.recover(checkpoint=0) + solver.checkpoint() + assert solver._last_chkpt_ind == solver._curr_chkpt_ind == 2 + + +def test_solver_opens_directory_with_stored_delete_flag(mock_sintel_data, tmp_path): + # directories created before the fix store delete_if_exists in their config + config = {**small_config(mock_sintel_data, n_iters=2), "delete_if_exists": False} + with set_root_context(str(tmp_path)): + legacy = NetworkDir("test", config) + + solver = MultiTaskSolver("test", config) + assert solver.dir.path == legacy.path + + with pytest.raises(FileExistsError): + MultiTaskSolver("test", small_config(mock_sintel_data, n_iters=3)) diff --git a/tests/test_train_single.py b/tests/test_train_single.py new file mode 100644 index 0000000..d4d05f1 --- /dev/null +++ b/tests/test_train_single.py @@ -0,0 +1,52 @@ +import os +import subprocess +import sys +from importlib import resources + +from flyvis.network.directories import NetworkDir + +TRAIN_SINGLE = str(resources.files("flyvis_cli") / "training" / "train_single.py") + + +def train_single(root_dir, sintel_path, *overrides): + return subprocess.run( + [ + sys.executable, + TRAIN_SINGLE, + "task_name=flow", + "ensemble_and_network_id=0000/000", + "description=test", + "task.n_iters=4", + f"+task.dataset.sintel_path={sintel_path}", + "task.original_split=false", + "task.dataset.boxfilter.extent=1", + "task.dataset.n_frames=4", + "task.dataset.dt=0.041", + "task.batch_size=2", + "network.connectome.extent=1", + "scheduler.chkpt_every_epoch=1", + *overrides, + ], + cwd=root_dir, + env={**os.environ, "FLYVIS_ROOT_DIR": str(root_dir)}, + capture_output=True, + text=True, + ) + + +def test_train_single_resume(mock_sintel_data, tmp_path): + # only the initial checkpoint, as left by a job stopped before training + run = train_single(tmp_path, mock_sintel_data, "train=false", "checkpoint_only=true") + assert run.returncode == 0, run.stderr + + # the same command again would train over the existing checkpoints + run = train_single(tmp_path, mock_sintel_data) + assert run.returncode != 0 + assert "already has checkpoints" in run.stderr + + run = train_single(tmp_path, mock_sintel_data, "resume=true") + assert run.returncode == 0, run.stderr + + network_dir = NetworkDir(tmp_path / "results" / "flow" / "0000" / "000") + assert len(network_dir.loss[:]) == 4 + assert network_dir.chkpt_iter[:].tolist()[-1] == 3