Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 77 additions & 21 deletions flyvis/solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,14 +7,15 @@

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

from flyvis.network.directories import NetworkDir
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,
Expand All @@ -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."""

Expand Down Expand Up @@ -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)
Comment thread
Copilot marked this conversation as resolved.
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

Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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.")
Expand Down Expand Up @@ -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
Comment thread
Copilot marked this conversation as resolved.

# 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.
Expand Down
6 changes: 6 additions & 0 deletions flyvis_cli/training/train_single.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
89 changes: 89 additions & 0 deletions tests/test_solver.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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))
52 changes: 52 additions & 0 deletions tests/test_train_single.py
Original file line number Diff line number Diff line change
@@ -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