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
4 changes: 2 additions & 2 deletions spikeinterface_gui/layout_presets.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,12 +93,12 @@ def get_layout_description(preset_name, layout=None):

merge_focus_layout = dict(
zone1=['merge', 'unitlist'],
zone2=['waveform'],
zone2=['curation'],
zone3=['spikeamplitude'],
zone4=['ndscatter'],
zone5=['probe'],
zone6=[],
zone7=['spikerate'],
zone7=['waveform'],
zone8=['correlogram'],
)

Expand Down
186 changes: 186 additions & 0 deletions spikeinterface_gui/tests/test_mainwindow_merge_focus.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,186 @@
from argparse import ArgumentParser
from spikeinterface_gui import run_mainwindow, run_launcher

from spikeinterface_gui.tests.testingtools import clean_all, make_analyzer_folder, make_curation_dict

from spikeinterface import load_sorting_analyzer


from pathlib import Path

import numpy as np
import sys


# yep is for testing
yep_layout = dict(
zone1=['curation', 'spikelist'],
zone2=['unitlist', 'mergelist'],
zone3=['trace', 'tracemap', 'spikeamplitude'],
zone4=['similarity'],
zone5=['probe'],
zone6=['ndscatter', ],
zone7=['waveform', 'waveformheatmap', ],
zone8=['correlogram', 'isi'],
)


def setup_module():
global test_folder
case = test_folder.stem.split('_')[-1]
make_analyzer_folder(test_folder, case=case, unit_dtype="int")


def teardown_module():
clean_all(test_folder)


def test_mainwindow(start_app=False, verbose=True, curation=False, only_some_extensions=False, events=False):

analyzer = load_sorting_analyzer(test_folder / "sorting_analyzer")
# analyzer = load_analyzer(test_folder / "sorting_analyzer.zarr")

tm = analyzer.get_extension("template_metrics").get_data().iloc[0, :]
# print(tm)
# return

print(analyzer)

if curation:
curation_dict = make_curation_dict(analyzer)
else:
curation_dict = None

if only_some_extensions:
analyzer = analyzer.copy()
# analyzer._recording = None
for k in ("principal_components", "template_similarity", "spike_amplitudes"):
analyzer.delete_extension(k)
print(analyzer)

n = analyzer.unit_ids.size
analyzer.sorting.set_property(
key='yep', values=np.array([f"yep{i}" for i in range(n)]))

extra_unit_properties = dict(
yop=np.array([f"yop{i}" for i in range(n)]),
yip=np.array([f"yip{i}" for i in range(n)]),
)

for segment_index in range(analyzer.get_num_segments()):
shift = (segment_index + 1) * 100
# add a gap to times
gap = 5
times = analyzer.recording.get_times(segment_index)
times = times + shift
times[len(times)//2:] += gap # add a gap in the middle
analyzer.recording.set_times(
times,
segment_index=segment_index
)

events_dict = None
if events:
events_dict = {"event1": {"times": []}, "event2": {"times": []}}
for segment_index in range(analyzer.get_num_segments()):
times = analyzer.recording.get_times(segment_index)
events_dict["event1"]["times"].append(
np.random.choice(times, 30)
)
events_dict["event2"]["times"].append(
np.random.choice(times, 50)
)
# add some events outside of recording times to test filtering
events_dict["event1"]["times"][-1] = np.concatenate(
[events_dict["event1"]["times"][-1],
[times[0] - 10, times[-1] + 20]]
)
events_dict["event2"]["times"][-1] = np.concatenate(
[events_dict["event2"]["times"][-1],
[times[0] - 5, times[-1] + 15]]
)

win = run_mainwindow(
analyzer,
mode="desktop",
start_app=start_app,
verbose=verbose,
curation=curation, curation_dict=curation_dict,
displayed_unit_properties=None,
extra_unit_properties=extra_unit_properties,
layout_preset='default',
events=events_dict
# user_settings={"mainsettings": {"color_mode": "color_by_visibility", "max_visible_units": 5}}
)


def test_launcher(verbose=True):

# case 1
analyzer_folders = None
root_folder = None

# case 2 : explore parent
analyzer_folders = None
root_folder = Path(__file__).parent

# case 3 : list
# analyzer_folders = [
# Path(__file__).parent / 'my_dataset_small/sorting_analyzer',
# Path(__file__).parent / 'my_dataset_big/sorting_analyzer',
# ]
# root_folder = None

# case 4 : dict
# analyzer_folders = {
# 'small' : Path(__file__).parent / 'my_dataset_small/sorting_analyzer',
# 'big' : Path(__file__).parent / 'my_dataset_big/sorting_analyzer',
# }
# root_folder = None

win = run_launcher(mode="desktop", analyzer_folders=analyzer_folders,
root_folder=root_folder, verbose=verbose)


def test_main_window_merge_focus(start_app=False, Verbose=True):
analyzer = load_sorting_analyzer(test_folder/"sorting_analyzer")
if not analyzer.has_extension("template_similarity"):
analyzer.compute_one_extension("template_similarity")
merge_unit_groups = sc.compute_merge_unit_groups(
analyzer,
preset="slay",
)
print("Computed merge unit groups:", merge_unit_groups)
win = run_mainwindow(
analyzer,
mode="desktop",
start_app=start_app,
verbose=verbose,
curation=True,
layout_preset="merge_focus",
merge_unit_groups=merge_unit_groups,
)
return win


parser = ArgumentParser()
parser.add_argument('--dataset', default="small",
help='Path to the dataset folder')
parser.add_argument('--events', action="store_true",
help='Simulate and add events')

if __name__ == '__main__':
args = parser.parse_args()
dataset = args.dataset
global test_folder
if dataset is not None:
test_folder = Path(__file__).parents[2] / f"my_dataset_{dataset}"

if not test_folder.is_dir():
setup_module()

win = test_mainwindow(start_app=True, verbose=True,
curation=True, events=args.events)
# win = test_mainwindow(start_app=True, verbose=True, curation=False)

# test_launcher(verbose=True)
2 changes: 1 addition & 1 deletion spikeinterface_gui/tests/testingtools.py
Original file line number Diff line number Diff line change
Expand Up @@ -138,9 +138,9 @@ def make_analyzer_folder(test_folder, case="small", unit_dtype="str"):
sorting_analyzer.compute("correlograms", window_ms=50., bin_ms=1.)
sorting_analyzer.compute("template_similarity", method="l1")
sorting_analyzer.compute("principal_components", n_components=3, mode='by_channel_global', whiten=True, **job_kwargs)
sorting_analyzer.compute(["spike_amplitudes", "spike_locations"], **job_kwargs)
sorting_analyzer.compute("quality_metrics", metric_names=["snr", "firing_rate"])
sorting_analyzer.compute("template_metrics")
sorting_analyzer.compute(["spike_amplitudes", "spike_locations"], **job_kwargs)


def make_curation_dict(analyzer):
Expand Down
Loading