From cc90adbc03be283fa395795169f3696299f8d442 Mon Sep 17 00:00:00 2001 From: Maris Basha Date: Wed, 23 Sep 2026 11:07:24 -0400 Subject: [PATCH] fix: keep ensemble order and arguments in cache keys and subsets Response cache: make_hashable sorted lists, so context_aware_cache gave the same key for the same names in a different order. After sort(), inside rank_by_validation_error() or select_items(), the Ensemble response methods returned the responses cached for the previous order, and network_id no longer matched ensemble.names. Arguments such as speeds=[25, 19] and speeds=[19, 25] also shared a key. Lists and tuples now keep their order; sets and dicts are still sorted. Subsets: ens[0:2] and ens[[0, 1]] rebuilt the subset from the model names with default arguments. root_dir, best_checkpoint_fn(_kwargs) and the other constructor arguments were dropped, and the names were resolved under flyvis.results_dir, so a subset of an ensemble loaded from another root pointed at different models. Subsets are now built from the model paths with the original constructor arguments. simulate_from_dataset: the per-batch responses were combined with np.stack, which raised whenever the number of stimuli was not a multiple of batch_size. They are now concatenated. --- flyvis/network/ensemble.py | 11 +++++----- flyvis/utils/cache_utils.py | 4 ++-- tests/test_cache_utils.py | 28 +++++++++++++++++++++++ tests/test_ensemble.py | 44 +++++++++++++++++++++++++++++++++++++ 4 files changed, 80 insertions(+), 7 deletions(-) create mode 100644 tests/test_cache_utils.py diff --git a/flyvis/network/ensemble.py b/flyvis/network/ensemble.py index b58f813..c5c8b55 100644 --- a/flyvis/network/ensemble.py +++ b/flyvis/network/ensemble.py @@ -203,10 +203,11 @@ def __getitem__( """ if isinstance(key, (int, np.integer)): return dict.__getitem__(self, self.names[key]) - elif isinstance(key, slice): - return self.__class__(self.names[key]) - elif isinstance(key, (np.ndarray, list)): - return self.__class__(np.array(self.names)[key]) + elif isinstance(key, (slice, np.ndarray, list)): + names = np.array(self.names)[key] + paths = [dict.__getitem__(self, name).dir.path for name in names] + # same constructor arguments, but no try_sort to keep the order of key + return self.__class__(paths, *self._init_args[1:-1]) elif key in self.names: return dict.__getitem__(self, key) else: @@ -375,7 +376,7 @@ def handle_network(network: Network): else: yield resp - r = np.stack(list(handle_network(network))) + r = np.concatenate(list(handle_network(network))) yield r.reshape(-1, r.shape[-2], r.shape[-1]) progress_bar.update(1) diff --git a/flyvis/utils/cache_utils.py b/flyvis/utils/cache_utils.py index a5e5b21..700ea5f 100644 --- a/flyvis/utils/cache_utils.py +++ b/flyvis/utils/cache_utils.py @@ -64,7 +64,7 @@ def make_hashable(obj: Any) -> Any: """Recursively converts an object into a hashable type.""" if isinstance(obj, (int, float, str, bool, type(None))): return obj - elif isinstance(obj, (list, set)): + elif isinstance(obj, (set, frozenset)): try: # Try direct sorting first return tuple(make_hashable(e) for e in sorted(obj)) @@ -86,7 +86,7 @@ def make_hashable(obj: Any) -> Any: key=lambda x: hash(make_hashable(x[0])), ) ) - elif isinstance(obj, (tuple, frozenset)): + elif isinstance(obj, (list, tuple)): return tuple(make_hashable(e) for e in obj) elif isinstance(obj, slice): return (obj.start, obj.stop, obj.step) diff --git a/tests/test_cache_utils.py b/tests/test_cache_utils.py new file mode 100644 index 0000000..f136082 --- /dev/null +++ b/tests/test_cache_utils.py @@ -0,0 +1,28 @@ +from flyvis.utils.cache_utils import context_aware_cache, make_hashable + + +class Responses: + def __init__(self, names): + self.names = names + self.cache = {} + + @context_aware_cache(context=lambda self: self.names) + def responses(self, *args): + return list(self.names) + + +def test_context_aware_cache_respects_order(): + responses = Responses(["a", "b", "c"]) + assert responses.responses() == ["a", "b", "c"] + responses.names = ["c", "b", "a"] + assert responses.responses() == ["c", "b", "a"] + assert responses.responses([2, 1]) == ["c", "b", "a"] + responses.names = ["a", "b", "c"] + assert responses.responses([1, 2]) == ["a", "b", "c"] + + +def test_make_hashable(): + assert make_hashable([1, 2]) != make_hashable([2, 1]) + assert make_hashable((1, 2)) != make_hashable((2, 1)) + assert make_hashable({1, 2}) == make_hashable({2, 1}) + assert make_hashable({"a": 1, "b": 2}) == make_hashable({"b": 2, "a": 1}) diff --git a/tests/test_ensemble.py b/tests/test_ensemble.py index 110f7b2..60844ac 100644 --- a/tests/test_ensemble.py +++ b/tests/test_ensemble.py @@ -1,13 +1,16 @@ +import shutil from copy import deepcopy from pathlib import Path import matplotlib.pyplot as plt import numpy as np +import pandas as pd import pytest import torch from datamate import Directory, Namespace from flyvis import results_dir +from flyvis.datasets.datasets import SequenceDataset from flyvis.network.ensemble import Ensemble, TaskError from flyvis.network.ensemble_view import EnsembleView from flyvis.network.network import IntegrationWarning, Network @@ -55,6 +58,47 @@ def test_simulate(ensemble: Ensemble): activity = np.array(list(ensemble.simulate(torch.ones(1, 2, 721).random_(2), 1))) +def test_getitem_subset_keeps_arguments(ensemble, tmp_path): + for name in ensemble.names[:3]: + shutil.copytree(results_dir / name, tmp_path / name) + other_root = Ensemble( + tmp_path / "flow/0000", + root_dir=tmp_path, + best_checkpoint_fn_kwargs=ensemble[0].best_checkpoint_fn_kwargs, + ) + + for subset, index in [(other_root[0:2], [0, 1]), (other_root[[2, 0]], [2, 0])]: + assert isinstance(subset, Ensemble) + assert [nv.dir.path for nv in subset.values()] == [ + other_root[i].dir.path for i in index + ] + for nv in subset.values(): + assert nv.best_checkpoint_fn_kwargs == other_root[0].best_checkpoint_fn_kwargs + + +class ConstantStimuli(SequenceDataset): + dt = 1 / 50 + t_pre = 0.0 + t_post = 0.0 + + def __init__(self, n_stimuli, n_frames=4, n_hexals=721): + self.arg_df = pd.DataFrame({"stimulus": np.arange(n_stimuli)}) + self.n_frames = n_frames + self.n_hexals = n_hexals + + def get_item(self, key): + return torch.full((self.n_frames, self.n_hexals), 0.5) + + +def test_simulate_from_dataset_last_batch_smaller(ensemble): + dataset = ConstantStimuli(n_stimuli=3) + responses = list( + ensemble[0:1].simulate_from_dataset(dataset, dt=1 / 50, t_pre=0.1, batch_size=2) + ) + assert len(responses) == 1 + assert responses[0].shape[:2] == (3, dataset.n_frames) + + def test_validation_losses(ensemble): losses = ensemble.validation_losses() assert len(losses) == len(ensemble)