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
11 changes: 6 additions & 5 deletions flyvis/network/ensemble.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
4 changes: 2 additions & 2 deletions flyvis/utils/cache_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand All @@ -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)
Expand Down
28 changes: 28 additions & 0 deletions tests/test_cache_utils.py
Original file line number Diff line number Diff line change
@@ -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})
44 changes: 44 additions & 0 deletions tests/test_ensemble.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading