From 187397f177875e3c9b45faf7626ec0fef767b6bf Mon Sep 17 00:00:00 2001 From: Maris Basha Date: Thu, 1 Oct 2026 10:50:55 -0400 Subject: [PATCH 1/2] fix: make solver resume work and keep the training history Reopening a network directory: MultiTaskSolver put delete_if_exists into the config it passes to NetworkDir. datamate stores that key when it creates the directory but removes it from the passed config when the directory exists, so opening the same directory again with the same config raised FileExistsError (incompatible config). The flag is now passed through datamate's delete_if_exists context instead, imported from datamate or, for datamate < 1.0, from datamate.directory. The context is only entered when deletion is requested, so an enclosing delete_if_exists context still applies. train-single: the FileExistsError above also stopped an accidental rerun of a trained network from training over its checkpoints. The script now raises a FileExistsError that names resume=true and delete_if_exists=true when the directory already has checkpoints and resume is not set. recover(): the method called resolve_checkpoints with four arguments against a one-argument signature (#23), then read checkpoint.index, checkpoints.index and checkpoints.path, none of which exist. It now resolves "best" through best_checkpoint_default_fn and an int as a position into the sorted checkpoints, e.g. -1 for the last one. Resumed iteration: checkpoints store the last completed iteration, self.iteration - 1, and recover() assigned it back unchanged. A resume repeated one iteration, and resuming from the initial checkpoint set the iteration to -1, where the scheduler picked the final learning rate. recover() now restores self.iteration as the stored value plus one. Training history: train() started the loss and activity lists empty on every call and then overwrote dir.loss, dir.activity* and dir.loss_, so continuing training or resuming kept only the iterations of the last call. The lists now start from the stored history up to the current iteration. --- flyvis/solver.py | 70 +++++++++++++++++++++-------- flyvis_cli/training/train_single.py | 6 +++ tests/test_solver.py | 64 ++++++++++++++++++++++++++ tests/test_train_single.py | 52 +++++++++++++++++++++ 4 files changed, 173 insertions(+), 19 deletions(-) create mode 100644 tests/test_train_single.py diff --git a/flyvis/solver.py b/flyvis/solver.py index cccc40d..786c1a7 100644 --- a/flyvis/solver.py +++ b/flyvis/solver.py @@ -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,6 +24,11 @@ ) 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"] @@ -106,9 +112,16 @@ 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: + self.dir = NetworkDir(name, config) self.path = self.dir.path @@ -283,12 +296,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): @@ -588,35 +609,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..831e1f9 100644 --- a/tests/test_solver.py +++ b/tests/test_solver.py @@ -1,4 +1,6 @@ +import numpy as np import pytest +import torch from datamate import set_root_context from flyvis.solver import MultiTaskSolver @@ -54,3 +56,65 @@ 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 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 From 6ba1755f28d4ae6c16000bc213747ec2e0f5f8d5 Mon Sep 17 00:00:00 2001 From: Maris Basha Date: Thu, 1 Oct 2026 14:22:57 -0400 Subject: [PATCH 2/2] solver: open older network directories and keep the checkpoint index current Older directories: directories created before the previous commit store delete_if_exists in their config, so opening them with the same config still raised FileExistsError. MultiTaskSolver now opens such a directory by path when the stored flag is the only difference to the passed config. Any other difference still raises. Checkpoint index: checkpoint() incremented _last_chkpt_ind and _curr_chkpt_ind separately. After recovering an older or the best checkpoint, the next checkpoint was saved under the last index but recorded the current index as the recovered one plus one, and a later recover() of that index returned early. The current index is now set to the newly saved checkpoint. --- flyvis/solver.py | 30 +++++++++++++++++++++++++++--- tests/test_solver.py | 25 +++++++++++++++++++++++++ 2 files changed, 52 insertions(+), 3 deletions(-) diff --git a/flyvis/solver.py b/flyvis/solver.py index 786c1a7..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 @@ -34,6 +34,25 @@ __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.""" @@ -121,7 +140,12 @@ def __init__( with delete_existing_directory(): self.dir = NetworkDir(name, config) else: - self.dir = NetworkDir(name, config) + 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 @@ -441,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.") diff --git a/tests/test_solver.py b/tests/test_solver.py index 831e1f9..4b4b9c3 100644 --- a/tests/test_solver.py +++ b/tests/test_solver.py @@ -3,6 +3,7 @@ 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 @@ -118,3 +119,27 @@ def test_solver_recover(mock_sintel_data, tmp_path): 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))