Skip to content

Commit 05cf038

Browse files
Add test coverage for Experiment class and data generation utilities (#113)
* test: add comprehensive coverage for Experiment class and data generation utilities * refactor: remove redundant helpers and use setup_test_experiment * fix: enforce non-NaN rows in mock data and restore return values --------- Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
1 parent 8e23c8c commit 05cf038

4 files changed

Lines changed: 380 additions & 184 deletions

File tree

experanto/experiment.py

Lines changed: 13 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -76,7 +76,7 @@ def __init__(
7676
def _load_devices(self) -> None:
7777
# Populate devices by going through subfolders
7878
# Assumption: blocks are sorted by start time
79-
device_folders = [d for d in self.root_folder.iterdir() if (d.is_dir())]
79+
device_folders = [d for d in self.root_folder.iterdir() if d.is_dir()]
8080

8181
for d in device_folders:
8282
if d.name not in self.modality_config:
@@ -95,14 +95,14 @@ def _load_devices(self) -> None:
9595
dev = instantiate(
9696
interp_conf, root_folder=d, cache_data=self.cache_data
9797
)
98+
9899
# Check if instantiated object is proper Interpolator
99100
if not isinstance(dev, Interpolator):
100101
raise ValueError(
101-
"Please provide an Interpolator which inherits from experantos Interpolator class."
102+
"Instantiated object must inherit from Interpolator class."
102103
)
103104

104105
elif isinstance(interp_conf, Interpolator):
105-
# Already instantiated Interpolator
106106
dev = interp_conf
107107

108108
else:
@@ -207,26 +207,22 @@ def interpolate(
207207
dict_keys(['screen', 'responses', 'eye_tracker'])
208208
"""
209209
if device is None:
210-
values = {}
211-
valid = {}
210+
values, valid = {}, {}
212211
for d, interp in self.devices.items():
213212
res = interp.interpolate(times, return_valid=return_valid)
214213
if return_valid:
215214
vals, vlds = res
216-
values[d] = vals
217-
valid[d] = vlds
215+
values[d], valid[d] = vals, vlds
218216
else:
219217
values[d] = res
220-
if return_valid:
221-
return values, valid
222-
else:
223-
return values
218+
return (values, valid) if return_valid else values
219+
224220
elif isinstance(device, str):
225-
assert device in self.devices, f"Unknown device '{device}'"
226-
res = self.devices[device].interpolate(times, return_valid=return_valid)
227-
return res
228-
else:
229-
raise ValueError(f"Unsupported device type: {type(device)}")
221+
if device not in self.devices:
222+
raise KeyError(f"Unknown device '{device}'")
223+
return self.devices[device].interpolate(times, return_valid=return_valid)
224+
225+
raise ValueError(f"Unsupported device type: {type(device)}")
230226

231227
def get_valid_range(self, device_name: str) -> tuple[float, float]:
232228
"""Get the valid time range for a specific device.
@@ -239,7 +235,7 @@ def get_valid_range(self, device_name: str) -> tuple[float, float]:
239235
Returns
240236
-------
241237
tuple
242-
A tuple ``(start_time, end_time)`` representing the valid
238+
A tuple `(start_time, end_time)` representing the valid
243239
time interval in seconds.
244240
"""
245241
return tuple(self.devices[device_name].valid_interval)

tests/create_experiment.py

Lines changed: 17 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -1,48 +1,30 @@
11
import shutil
22
from contextlib import contextmanager
33

4-
import numpy as np
5-
import yaml
4+
from .create_sequence_data import _generate_sequence_data
65

76

87
@contextmanager
9-
def make_sequence_device(
10-
root, name, start, end, sampling_rate=10.0, n_signals=5, override_meta=None
8+
def setup_test_experiment(
9+
tmp_path,
10+
n_devices=2,
11+
devices_kwargs=None,
12+
default_sampling_rate=1.0,
1113
):
12-
"""Create a single sequence device folder under root."""
13-
device_root = root / name
14-
try:
15-
(device_root / "meta").mkdir(parents=True, exist_ok=True)
16-
17-
n_samples = (
18-
int((end - start) * sampling_rate) + 1
19-
) # +1 to include both start and end as sample points
20-
timestamps = np.linspace(start, end, n_samples)
21-
data = np.random.rand(n_samples, n_signals)
22-
23-
np.save(device_root / "timestamps.npy", timestamps)
24-
np.save(device_root / "data.npy", data)
14+
devices_kwargs = devices_kwargs or [{}] * n_devices
15+
default_params = {"sampling_rate": default_sampling_rate}
2516

26-
meta = {
27-
"start_time": start,
28-
"end_time": end,
29-
"modality": "sequence",
30-
"sampling_rate": sampling_rate,
31-
"phase_shift_per_signal": False,
32-
"is_mem_mapped": False,
33-
"n_signals": n_signals,
34-
"n_timestamps": n_samples,
35-
"dtype": "float64",
36-
}
37-
if override_meta:
38-
meta.update(override_meta)
39-
with open(device_root / "meta.yml", "w") as f:
40-
yaml.safe_dump(meta, f)
41-
42-
yield device_root
17+
devices_kwargs = [default_params | kwargs for kwargs in devices_kwargs]
4318

19+
try:
20+
tmp_path.mkdir(parents=True, exist_ok=True)
21+
for device_id, device_kwargs in enumerate(devices_kwargs):
22+
device_path = tmp_path / f"device_{device_id}"
23+
_generate_sequence_data(device_path, **device_kwargs)
24+
yield tmp_path
4425
finally:
45-
shutil.rmtree(device_root)
26+
if tmp_path.exists():
27+
shutil.rmtree(tmp_path)
4628

4729

4830
def make_modality_config(*device_names, sampling_rates=None, offsets=None):

tests/create_sequence_data.py

Lines changed: 81 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -10,74 +10,110 @@
1010
SEQUENCE_ROOT = Path("tests/sequence_data")
1111

1212

13-
@contextmanager
14-
def create_sequence_data(
13+
def _generate_sequence_data(
14+
sequence_root,
1515
n_signals=10,
1616
shifts_per_signal=False,
1717
use_mem_mapped=False,
18+
start_time=0.0,
1819
t_end=10.0,
1920
sampling_rate=10.0,
2021
contain_nans=False,
2122
):
22-
try:
23-
SEQUENCE_ROOT.mkdir(parents=True, exist_ok=True)
24-
(SEQUENCE_ROOT / "meta").mkdir(parents=True, exist_ok=True)
25-
26-
meta = {
27-
"start_time": 0,
28-
"end_time": t_end,
29-
"modality": "sequence",
30-
"sampling_rate": sampling_rate,
31-
"phase_shift_per_signal": shifts_per_signal,
32-
"is_mem_mapped": use_mem_mapped,
33-
"n_signals": n_signals,
34-
}
35-
36-
timestamps = np.linspace(
37-
meta["start_time"],
38-
meta["end_time"],
39-
int((meta["end_time"] - meta["start_time"]) * meta["sampling_rate"]) + 1,
23+
"""Generates synthetic sequence data folders for testing interpolator logic."""
24+
25+
sequence_root = Path(sequence_root)
26+
sequence_root.mkdir(parents=True, exist_ok=True)
27+
(sequence_root / "meta").mkdir(parents=True, exist_ok=True)
28+
29+
meta = {
30+
"start_time": start_time,
31+
"end_time": t_end,
32+
"modality": "sequence",
33+
"sampling_rate": sampling_rate,
34+
"phase_shift_per_signal": shifts_per_signal,
35+
"is_mem_mapped": use_mem_mapped,
36+
"n_signals": n_signals,
37+
}
38+
39+
# Determine number of samples based on duration and rate
40+
duration = meta["end_time"] - meta["start_time"]
41+
n_samples = int(round(duration * meta["sampling_rate"])) + 1
42+
43+
timestamps = np.linspace(meta["start_time"], meta["end_time"], n_samples)
44+
45+
data = np.random.rand(len(timestamps), n_signals)
46+
47+
if contain_nans:
48+
nan_indices = np.random.choice(
49+
data.size, size=int(0.1 * data.size), replace=False
4050
)
41-
np.save(SEQUENCE_ROOT / "timestamps.npy", timestamps)
42-
meta["n_timestamps"] = len(timestamps)
51+
data.flat[nan_indices] = np.nan
52+
# ensure each row has at least one non-NaN
53+
if n_signals > 0:
54+
row_all_nan = np.isnan(data).all(axis=1)
55+
data[row_all_nan, 0] = 0.0
4356

44-
data = np.random.rand(len(timestamps), n_signals)
57+
if not use_mem_mapped:
58+
np.save(sequence_root / "data.npy", data)
59+
else:
60+
filename = sequence_root / "data.mem"
61+
fp = np.memmap(filename, dtype=data.dtype, mode="w+", shape=data.shape)
62+
fp[:] = data[:]
63+
fp.flush()
64+
del fp
4565

46-
if contain_nans:
47-
nan_indices = np.random.choice(
48-
data.size, size=int(0.1 * data.size), replace=False
49-
)
50-
data.flat[nan_indices] = np.nan
66+
np.save(sequence_root / "timestamps.npy", timestamps)
67+
meta["n_timestamps"] = len(timestamps)
68+
meta["dtype"] = str(data.dtype)
5169

52-
if not use_mem_mapped:
53-
np.save(SEQUENCE_ROOT / "data.npy", data)
54-
else:
55-
filename = SEQUENCE_ROOT / "data.mem"
70+
# Handle per-signal phase shifts if required by the test case
71+
shifts = None
72+
if shifts_per_signal:
73+
shifts = np.random.rand(n_signals) / meta["sampling_rate"] * 0.9
74+
np.save(sequence_root / "meta" / "phase_shifts.npy", shifts)
5675

57-
fp = np.memmap(filename, dtype=data.dtype, mode="w+", shape=data.shape)
58-
fp[:] = data[:]
59-
fp.flush() # Ensure data is written to disk
60-
del fp
61-
meta["dtype"] = str(data.dtype)
76+
with open(sequence_root / "meta.yml", "w") as f:
77+
yaml.safe_dump(meta, f)
6278

63-
if shifts_per_signal:
64-
shifts = np.random.rand(n_signals) / meta["sampling_rate"] * 0.9
65-
np.save(SEQUENCE_ROOT / "meta" / "phase_shifts.npy", shifts)
79+
return timestamps, data, shifts
6680

67-
with open(SEQUENCE_ROOT / "meta.yml", "w") as f:
68-
yaml.safe_dump(meta, f)
6981

70-
yield timestamps, data, shifts if shifts_per_signal else None
82+
@contextmanager
83+
def create_sequence_data(
84+
n_signals=10,
85+
shifts_per_signal=False,
86+
use_mem_mapped=False,
87+
t_end=10.0,
88+
sampling_rate=10.0,
89+
contain_nans=False,
90+
start_time=0.0,
91+
):
92+
"""Context manager for temporary sequence data creation and cleanup."""
93+
try:
94+
yield _generate_sequence_data(
95+
sequence_root=SEQUENCE_ROOT,
96+
n_signals=n_signals,
97+
shifts_per_signal=shifts_per_signal,
98+
use_mem_mapped=use_mem_mapped,
99+
t_end=t_end,
100+
sampling_rate=sampling_rate,
101+
contain_nans=contain_nans,
102+
start_time=start_time,
103+
)
71104
finally:
72-
shutil.rmtree(SEQUENCE_ROOT)
105+
if SEQUENCE_ROOT.exists():
106+
shutil.rmtree(SEQUENCE_ROOT)
73107

74108

75109
@contextmanager
76110
def sequence_data_and_interpolator(data_kwargs=None, interp_kwargs=None):
77111
data_kwargs = data_kwargs or {}
78112
interp_kwargs = interp_kwargs or {}
79113
with create_sequence_data(**data_kwargs) as (timestamps, data, shifts):
114+
# Restore the helper expected by the rest of the test suite
115+
80116
with closing(
81-
Interpolator.create("tests/sequence_data", **interp_kwargs)
117+
Interpolator.create(str(SEQUENCE_ROOT), **interp_kwargs)
82118
) as seq_interp:
83119
yield timestamps, data, shifts, seq_interp

0 commit comments

Comments
 (0)