from __future__ import annotations
import logging
import re
from abc import ABC, abstractmethod
from collections import defaultdict
from datetime import timedelta
from pathlib import Path
from typing import TYPE_CHECKING, Any
import numpy as np
import pandas as pd
from ibl_alignment_gui.handlers.shank_handler import ShankHandler
from ibl_alignment_gui.loaders.alignment_loader import (
AlignmentLoaderDocDB,
AlignmentLoaderLocal,
AlignmentLoaderOne,
)
from ibl_alignment_gui.loaders.alignment_uploader import (
AlignmentUploaderDocDB,
AlignmentUploaderLocal,
AlignmentUploaderOne,
)
from ibl_alignment_gui.loaders.data_loader import (
DataLoaderLocal,
DataLoaderOne,
FeatureLoaderLocal,
FeatureLoaderOne,
SpikeGLXLoaderLocal,
SpikeGLXLoaderOne,
)
from ibl_alignment_gui.loaders.geometry_loader import (
GeometryLoaderLocal,
GeometryLoaderOne,
)
from ibl_alignment_gui.loaders.histology_loader import (
SliceLoader,
build_anatomical_atlas,
download_histology_data,
make_slice_loader,
)
from ibl_alignment_gui.loaders.plot_loader import PlotLoader
from ibl_alignment_gui.loaders.transform_loader import TransformLoaderAllen
from ibl_alignment_gui.utils.parse_yaml import DatasetPaths, load_alignment_yaml
from iblatlas.atlas import AllenAtlas, BrainAtlas
from iblutil.util import Bunch
from one import params
from one.api import ONE
if TYPE_CHECKING:
from collections.abc import Callable
from ibl_alignment_gui.backends.allen.docdb_api import DocDB
logger = logging.getLogger(__name__)
try:
import ephysatlas.data
EPHYS_ATLAS = True
except ImportError:
EPHYS_ATLAS = False
[docs]
class ProbeHandler(ABC):
"""
Abstract base class for handling alignment and data loading for a probe.
This class provides access to loader methods that handle different aspects of the
alignment process. Where applicable the probe is split into shanks and each shank is handled
separately.
It can also handle multiple configurations for each shank, for example if two different channel
maps are used to record data.
Parameters
----------
brain_atlas: BrainAtlas or None
A pre-built BrainAtlas instance (AllenAtlas or BrainAtlasAnatomical). If None, the atlas
is built lazily by :meth:`build_atlas` (run on a background thread as part of
:meth:`load_all`), so that the slow atlas construction does not block the GUI.
"""
def __init__(self, brain_atlas: BrainAtlas | None = None):
# Built lazily by build_atlas() (off the GUI thread) when not supplied up front.
self.brain_atlas: BrainAtlas | None = brain_atlas
self.shanks: dict[str, Bunch] = defaultdict(Bunch)
# Configuration state
self.default_config: str = 'default'
self.non_default_config: str | None = None
self.configs: list[str] = ['default']
self.possible_configs: list[str] = ['default']
self.selected_config: str = 'default'
# Active shank indices
self.selected_shank: str | None = None
self.selected_idx: int | None = None
# -------------------------------------------------------------------------
# Shank access methods
# -------------------------------------------------------------------------
[docs]
def get_current_shank(self, shank, config) -> ShankHandler:
"""Return the currently active shank."""
return self.shanks[shank][config]
[docs]
def get_selected_shank(self) -> Bunch:
"""Return the currently selected shank."""
return self.shanks[self.selected_shank]
[docs]
def get_config(self, idx: int) -> None:
"""
Select a configuration by index.
Parameters
----------
idx : int
Index in the list of possible configurations.
"""
self.selected_config = self.possible_configs[idx]
# -------------------------------------------------------------------------
# Alignment methods - methods in loaders['align']
# -------------------------------------------------------------------------
[docs]
def load_previous_alignments(self) -> dict:
"""
Load previous alignments for the selected shank.
Always returns the alignments from the default configuration.
Returns
-------
dict
Previous alignments for the selected shank.
"""
# Load the previous alignment from the default configuration
self.get_selected_shank()[self.default_config].loaders['align'].load_previous_alignments()
# Set the previous alignments from the non default configuration to the default
# one (if it exists)
if self.non_default_config is not None:
self.get_selected_shank()[self.non_default_config].loaders['align'].alignments = (
self.get_selected_shank()[self.default_config].loaders['align'].alignments
)
self.get_selected_shank()[self.non_default_config].loaders[
'align'
].get_previous_alignments()
return (
self.get_selected_shank()[self.default_config]
.loaders['align']
.get_previous_alignments()
)
[docs]
def get_previous_alignments(self) -> dict:
"""
Get previous alignments for the selected shank.
Always returns the alignments from the default configuration.
Returns
-------
dict
Previous alignments for the selected shank
"""
return (
self.get_selected_shank()[self.default_config]
.loaders['align']
.get_previous_alignments()
)
[docs]
def get_starting_alignment(self, idx: int) -> None:
"""
Set the index of the starting alignment for the selected shank and each configuration.
Parameters
----------
idx: int
The index of the previous alignment to load.
"""
for config in self.configs:
self.get_selected_shank()[config].loaders['align'].get_starting_alignment(idx)
[docs]
def get_start_alignment_idx(self) -> int:
"""
Return the index of the alignment to display for the selected shank when data is loaded.
Delegates to the default configuration's alignment loader. A recovered alignment takes
precedence over the stored (resolved) alignment, which in turn takes precedence over the
most recent one.
Returns
-------
int
Index of the alignment in the alignment keys list.
"""
return (
self.get_selected_shank()[self.default_config]
.loaders['align']
.get_start_alignment_idx()
)
[docs]
def set_init_alignment(self) -> None:
"""Initialise the alignment for the selected shank and each configuration."""
for config in self.configs:
self.get_selected_shank()[config].set_init_alignment()
[docs]
def sync_config_alignments(self) -> None:
"""
Make the configurations of a shank share one set of alignments.
The alignments of the default configuration are shared with the non default one, so that
the alignments, and the keys used to choose between them, can never differ between the
configurations of a shank. Both configurations refer to the same dictionary, so any later
change is seen by both. Does nothing when there is only one configuration.
The default configuration is the source of truth, so where it is the online one its
alignments take precedence over any that were found locally. Where it has no alignments of
its own it takes on those of the other configuration, so that none are lost.
The starting alignment of both configurations is set, as the alignments they can choose
between may have changed.
"""
if self.non_default_config is None:
return
for shank in self.shanks:
default = self.shanks[shank].get(self.default_config)
other = self.shanks[shank].get(self.non_default_config)
if default is None or other is None:
continue
default_align = default.loaders['align']
other_align = other.loaders['align']
# Where the default configuration has nothing of its own, take on the alignments of
# the other configuration rather than discarding them
if not default_align.alignments and other_align.alignments:
default_align.add_extra_alignments(other_align.alignments)
other_align.alignments = default_align.alignments
other_align.stored_alignment_key = default_align.stored_alignment_key
# The alignments are shared, so the same index applies to both configurations
default_align.get_previous_alignments()
other_align.get_previous_alignments()
idx = default_align.get_start_alignment_idx()
default_align.get_starting_alignment(idx)
other_align.get_starting_alignment(idx)
# -------------------------------------------------------------------------
# Alignment handler methods - methods in align_handle
# -------------------------------------------------------------------------
[docs]
def next_idx(self) -> int:
"""
Return the index of the next available alignment for the selected shank.
Returns
-------
int
The index of the next available alignment stored in th circular buffer.
"""
if self.non_default_config is not None:
self.get_selected_shank()[self.non_default_config].align_handle.next_idx()
return self.get_selected_shank()[self.default_config].align_handle.next_idx()
[docs]
def prev_idx(self) -> int:
"""
Return the index of the previously available alignment for the selected shank.
Returns
-------
int
The index of the previous available alignment stored in th circular buffer.
"""
if self.non_default_config is not None:
self.get_selected_shank()[self.non_default_config].align_handle.prev_idx()
return self.get_selected_shank()[self.default_config].align_handle.prev_idx()
@property
def current_idx(self) -> int:
"""
Return the current index of the alignment stored in the buffer for the selected shank.
Returns
-------
int
The index of the current alignment
"""
try:
return self.get_selected_shank()[self.default_config].align_handle.current_idx
except AttributeError:
return 0
@property
def total_idx(self) -> int:
"""
Return the total index of the alignments stored in the buffer for the selected shank.
Returns
-------
int
The total number of alignments stored in the circular buffer
"""
try:
return self.get_selected_shank()[self.default_config].align_handle.total_idx
except AttributeError:
return 0
[docs]
def get_plot(self, shank: str, plot: str, key: str, config: str | None = None) -> Any:
"""
Access a specific plot for a specific shank and configuration.
Parameters
----------
shank: str
The shank label to access.
plot: str
The plot type to access. One of 'image', 'scatter', 'line', 'probe' or 'slice'
key: str
The plot key to access.
config: str
The configuration to access. If None, uses the default configuration.
Returns
-------
Any
The requested plot, or None if not found.
"""
config = config or self.default_config
return getattr(self.shanks[shank][config].loaders['plots'], plot).get(key, None)
[docs]
def get_plot_keys(self, plot: str) -> list[str]:
"""
Find a list of available keys across all shanks and configurations for a given plot type.
Parameters
----------
plot
The plot type to get the keys for. One of 'image', 'scatter', 'line', 'probe'
or 'slice'
Returns
-------
list
A list of unique plot keys.
"""
keys = []
for shank in self.shanks:
for config in self.configs:
keys += getattr(self.shanks[shank][config].loaders['plots'], plot).keys()
return sorted(set(keys))
@property
def image_keys(self) -> list[str]:
"""
Find the list of available image plot keys across all shanks and configurations.
Returns
-------
list:
A list of unique image plot keys.
"""
return self.get_plot_keys('image_plots')
@property
def scatter_keys(self) -> list[str]:
"""
Find the list of available scatter plot keys across all shanks and configurations.
Returns
-------
list:
A list of unique scatter plot keys.
"""
return self.get_plot_keys('scatter_plots')
@property
def line_keys(self) -> list[str]:
"""
Find the list of available line plot keys across all shanks and configurations.
Returns
-------
list:
A list of unique line plot keys.
"""
return self.get_plot_keys('line_plots')
@property
def probe_keys(self) -> list[str]:
"""
Find the list of available probe plot keys across all shanks and configurations.
Returns
-------
tuple:
A tuple of unique probe plot keys.
"""
return self.get_plot_keys('probe_plots')
@property
def feature_keys(self) -> list[str]:
"""
Find the list of available feature plot keys across all shanks and configurations.
Returns
-------
tuple:
A tuple of unique probe plot keys.
"""
return self.get_plot_keys('feature_plots')
@property
def slice_keys(self) -> list[str]:
"""
Find the list of available slice plot keys across all shanks and configurations.
Returns
-------
list:
A list of unique slice plot keys.
"""
return self.get_plot_keys('slice_plots')
# -------------------------------------------------------------------------
# Data loading & upload
# -------------------------------------------------------------------------
def _make_atlas(self) -> BrainAtlas:
"""
Build the brain atlas to use for this probe handler.
The base implementation returns an Allen CCF atlas; subclasses override this to select a
different atlas (e.g. an anatomical atlas built from the session histology).
Returns
-------
BrainAtlas
The brain atlas instance.
"""
return AllenAtlas()
[docs]
def build_atlas(
self, progress_callback: Callable[[str, int, int], None] | None = None
) -> None:
"""
Build the brain atlas if not already available, and share it with the shank uploaders.
Building the atlas (downloading/reading volumes) is slow, so this is called from
:meth:`load_all` on a background thread rather than in ``__init__``. The uploaders created
in ``initialise_shanks`` hold a reference to the atlas; because the atlas may not exist yet
at that point, this method (re)assigns the freshly built atlas onto each of them.
"""
if progress_callback is not None:
progress_callback('Building atlas…', 0, 0)
if self.brain_atlas is None:
self.brain_atlas = self._make_atlas()
for probe in self.shanks:
for config in self.configs:
upload = self.shanks[probe][config].loaders.get('upload')
if upload is not None:
upload.brain_atlas = self.brain_atlas
[docs]
def load_all(self, progress_callback: Callable[[str, int, int], None] | None = None) -> None:
"""
Build the atlas, then load all data and plots for the session.
This is the single entry point run on the background loading thread, so the slow atlas
construction, data loading and plot building all happen off the GUI thread.
Parameters
----------
progress_callback : Callable or None
Optional callback invoked as ``progress_callback(message, current, total)`` to report
progress. No-op when None.
"""
self.build_atlas(progress_callback=progress_callback)
self.load_data(progress_callback=progress_callback)
[docs]
def load_data(self, progress_callback: Callable[[str, int, int], None] | None = None) -> None:
"""
Download and load data for all configs and shanks.
Parameters
----------
progress_callback : Callable or None
Optional callback invoked as ``progress_callback(message, current, total)`` before
each loading step to report progress (e.g. to a GUI progress dialog). No-op when
None, so headless callers are unaffected.
"""
total = len(self.shanks) * len(self.configs) + 1
if progress_callback is not None:
progress_callback('Downloading histology…', 0, total)
slice_loader = self.download_histology()
idx = 1
for probe in self.shanks:
for config in self.configs:
if progress_callback is not None:
progress_callback(f'Loading {probe} ({config})…', idx, total)
self.shanks[probe][config].loaders['hist'] = slice_loader
self.shanks[probe][config].load_data()
idx += 1
[docs]
def load_plots(self, progress_callback: Callable[[str, int, int], None] | None = None) -> None:
"""
Load plots for all configs and shanks.
Parameters
----------
progress_callback : Callable or None
Optional callback invoked as ``progress_callback(message, current, total)`` before
each loading step to report progress. No-op when None.
"""
total = len(self.shanks) * len(self.configs)
idx = 0
for probe in self.shanks:
for config in self.configs:
if progress_callback is not None:
progress_callback(f'Computing plots {probe} ({config})…', idx, total)
self.shanks[probe][config].load_plots()
idx += 1
[docs]
def upload_data(self) -> str:
"""
Upload data for the selected shank for each configuration.
Always returns the upload result from the default configuration.
Returns
-------
str
Upload result from the default config.
"""
info = Bunch()
for config in self.configs:
info[config] = self.get_selected_shank()[config].upload_data()
return info[self.default_config]
[docs]
def upload_shanks(
self,
shanks: list[str],
progress_callback: Callable[[str, int, int], None] | None = None,
) -> dict[str, str]:
"""
Upload data for several shanks in turn.
Saving channels and alignments (and, online, registering tracks, running alignment QC and
writing to flatiron) is slow, so this is run on a background thread. Any per-shank user
input (QC, upload confirmation) must be gathered on the main thread beforehand; this method
only performs the saving.
Parameters
----------
shanks : list of str
The shanks to upload, in order.
progress_callback : Callable or None
Optional callback invoked as ``progress_callback(message, current, total)`` before each
shank is uploaded. No-op when None.
Returns
-------
dict[str, str]
A mapping of shank label to the upload result message for that shank.
"""
info: dict[str, str] = {}
total = len(shanks)
# The shank is switched to reach each one in turn, so keep the one the user had selected
selected_shank = self.selected_shank
try:
for idx, shank in enumerate(shanks):
if progress_callback is not None:
progress_callback(f'Saving {shank}…', idx, total)
self.selected_shank = shank
info[shank] = self.upload_data()
self.load_previous_alignments()
self.get_starting_alignment(0)
finally:
self.selected_shank = selected_shank
return info
[docs]
def save_progress(self, shanks: list[str]) -> dict[str, str]:
"""
Save the current alignment of several shanks to file.
The alignments are saved so that they can be recovered if the GUI crashes before they
have been uploaded. Only the default configuration is saved.
Parameters
----------
shanks : list of str
The shanks to save the alignment for.
Returns
-------
dict[str, str]
A mapping of shank label to the save result message for that shank.
"""
info: dict[str, str] = {}
for shank in shanks:
info[shank] = self.shanks[shank][self.default_config].save_progress()
return info
# -------------------------------------------------------------------------
# Utility
# -------------------------------------------------------------------------
[docs]
@staticmethod
def normalize_shank_label(shank_label: str) -> str:
"""
Normalize a shank label to the form 'probe0X'.
Parameters
----------
shank_label : str
Input shank label.
Returns
-------
str
Normalized label.
"""
match = re.match(r'(probe\d+)', shank_label)
return match.group(1) if match else shank_label
# -------------------------------------------------------------------------
# Abstract methods
# -------------------------------------------------------------------------
[docs]
@abstractmethod
def set_info(self, *args):
"""Set probe information."""
[docs]
@abstractmethod
def download_histology(self):
"""Load histology data."""
[docs]
@abstractmethod
def get_shanks(self, *args):
"""Return shank information."""
[docs]
@abstractmethod
def initialise_shanks(self):
"""Initialize shank data."""
[docs]
class ProbeHandlerONE(ProbeHandler):
"""
ONE implementation of ProbeHandler.
For this ProbeHandler all ephys and alignment data is downloaded and accessed via
ONE and Alyx.
The data for all shanks on a probe will be loaded at once.
Parameters
----------
one : ONE
An ONE instance used to upload results to Alyx
brain_atlas : AllenAtlas
An AllenAtlas object.
spike_collection : str, optional
Spike sorting algorithm to load (e.g. 'pykilosort', 'iblsorter').
"""
def __init__(
self,
one: ONE = None,
brain_atlas: BrainAtlas | None = None,
spike_collection: str | None = None,
):
self.one = one or ONE()
self.spike_collection = spike_collection
if EPHYS_ATLAS:
self.ea_model = ephysatlas.data.get_latest_label(one=self.one, project='ea_active')
else:
self.ea_model = None
super().__init__(brain_atlas)
[docs]
def get_subjects(self) -> np.ndarray:
"""
Find all subjects that have probe insertions with spikesorting data.
Returns
-------
np.ndarray
An array of subject names
"""
self.sess_ins = self.one.alyx.rest(
'insertions', 'list', dataset_types='spikes.times', expires=timedelta(days=1)
)
self.subj_ins = [sess['session_info']['subject'] for sess in self.sess_ins]
self.subjects = np.unique(self.subj_ins)
return self.subjects
[docs]
def get_sessions(self, idx: int) -> np.ndarray:
"""
Find all probes for a given subject.
Note if multi-shank data it will return probe00 rather than probe00a, the individual shank
is chosen using the shank dropdown.
Parameters
----------
idx : idx
The index of the chosen subject
Returns
-------
np.ndarray
All probes with spikesorting data for the chosen subject
"""
self.chosen_sess = self.subjects[idx]
sess_idx = [i for i, e in enumerate(self.subj_ins) if e == self.chosen_sess]
self.sess = [self.sess_ins[idx] for idx in sess_idx]
self.sessions = [self.get_session_probe_name(sess) for sess in self.sess]
self.sessions = np.unique(self.sessions)
return self.sessions
[docs]
def get_shanks(self, idx: int) -> list:
"""
Find all shanks for a given probe and initialise the loaders.
Parameters
----------
idx : idx
The index of the chosen probe
Returns
-------
np.ndarray
All shanks for the chosen probe
"""
self.chosen_probe = self.sessions[idx]
sess_idx = [
i
for i, e in enumerate(self.sess)
if self.get_session_probe_name(e) == self.chosen_probe
]
self.shank_labels = [self.sess[idx] for idx in sess_idx]
shanks = [s['name'] for s in self.shank_labels]
idx = np.argsort(shanks)
self.shank_labels = np.array(self.shank_labels)[idx]
shanks = np.array(shanks)[idx]
self.lab = self.shank_labels[0]['session_info']['lab']
self.initialise_shanks()
return list(shanks)
[docs]
def resolve_pid(self, pid: str) -> tuple[int, int, int]:
"""
Resolve a probe insertion id to subject, session and shank dropdown indices.
The internal session and shank state is populated as a side effect (via
:meth:`get_sessions` and :meth:`get_shanks`) so that the dropdowns can be
configured to point at the requested insertion.
Parameters
----------
pid : str
The probe insertion id (UUID) to resolve.
Returns
-------
tuple[int, int, int]
The subject, session and shank dropdown indices for the insertion.
Raises
------
ValueError
If no insertion exists for `pid`, or its subject has no spikesorted
insertions (and so is absent from the subject dropdown).
"""
ins = self.one.alyx.rest('insertions', 'list', id=pid)
if len(ins) == 0:
raise ValueError(f'No probe insertion found for pid {pid}')
ins = ins[0]
subject = ins['session_info']['subject']
subj_match = np.where(self.subjects == subject)[0]
if len(subj_match) == 0:
raise ValueError(f'Subject {subject} for pid {pid} has no spikesorted insertions')
subj_idx = int(subj_match[0])
sessions = self.get_sessions(subj_idx)
sess_idx = int(np.where(sessions == self.get_session_probe_name(ins))[0][0])
shanks = self.get_shanks(sess_idx)
shank_idx = shanks.index(ins['name'])
return subj_idx, sess_idx, shank_idx
[docs]
def get_session_probe_name(self, ins: dict) -> str:
"""
Make a string containing the combination of session information and probe name.
Removes the shank identifiers from the probe names.
Parameters
----------
ins: dict
A dict containing insertion data
Returns
-------
str:
A string with the session info and probe name
"""
return (
ins['session_info']['start_time'][:10]
+ ' '
+ f'{ins["session_info"]["number"]:03}'
+ ' '
+ self.normalize_shank_label(ins['name'])
)
[docs]
def set_info(self, idx):
"""
Set the information about the selected shank.
Parameters
----------
idx: int
The index of the selected shank
"""
self.selected_shank = self.shank_labels[idx]['name']
self.selected_idx = idx
self.subj = self.shank_labels[idx]['session_info']['subject']
self.lab = self.shank_labels[idx]['session_info']['lab']
self.pid = self.shank_labels[idx]['id']
[docs]
def download_histology(self) -> SliceLoader:
"""Download and load in the histology slice data."""
_, hist_path = download_histology_data(self.subj, self.lab)
return make_slice_loader(hist_path, self.brain_atlas, 'ccf')
[docs]
def initialise_shanks(self):
"""Initialise each shank with the loaders."""
self.shanks = defaultdict(Bunch)
for ins in self.shank_labels:
loaders = Bunch()
loaders['data'] = DataLoaderOne(ins, self.one, spike_collection=self.spike_collection)
loaders['geom'] = GeometryLoaderOne(
ins, self.one, probe_collection=loaders['data'].probe_collection
)
# Work in progress alignments are saved next to the spike sorting data
spike_path = loaders['data'].spike_sorting_path
loaders['align'] = AlignmentLoaderOne(
ins, self.one, user=params.get().ALYX_LOGIN, data_path=spike_path
)
loaders['upload'] = AlignmentUploaderOne(
ins, self.one, self.brain_atlas, data_path=spike_path
)
loaders['ephys'] = SpikeGLXLoaderOne(ins, self.one)
if EPHYS_ATLAS:
loaders['features'] = FeatureLoaderOne(
ins, self.one, self.ea_model, multi_area=self.lab == 'steinmetzlab'
)
loaders['plots'] = PlotLoader()
self.shanks[ins['name']][self.default_config] = ShankHandler(loaders, 0)
[docs]
def load_data(self, progress_callback: Callable[[str, int, int], None] | None = None) -> None:
"""Load data for all configs and shanks."""
logger.info(
'Loading session %s %s (pid: %s)', self.chosen_sess, self.chosen_probe, self.pid
)
super().load_data(progress_callback=progress_callback)
[docs]
class ProbeHandlerCSV(ProbeHandler):
"""
ProbeHandler where data from two channel maps has been recorded on the shanks.
The data for the dense configuration is available via ONE whereas the data for the quarter
configuration is only available on the local file system. Reads in a csv file that contains
information about where to read the relevant data from.
"""
def __init__(
self, csv_file: str | Path, one: ONE = None, brain_atlas: BrainAtlas | None = None
):
super().__init__(brain_atlas)
csv_file = Path(csv_file)
assert csv_file.exists()
self.root_path = csv_file.parent
self.df = pd.read_csv(csv_file, keep_default_na=False)
self.df['session_strip'] = self.df['session'].str.rsplit('/', n=1).str[0]
self.one = one or ONE()
self.possible_configs = ['quarter', 'dense', 'both']
self.configs = ['quarter', 'dense']
self.default_config = 'dense'
self.non_default_config = 'quarter'
self.selected_config = 'quarter'
if EPHYS_ATLAS:
self.ea_model = ephysatlas.data.get_latest_label(one=self.one, project='ea_active')
else:
self.ea_model = None
[docs]
def get_subjects(self) -> np.ndarray:
"""
Find all sessions with spike sorting data.
Returns
-------
np.ndarray
All sessions with spikesorting data.
"""
# Returns sessions
self.subjects = self.df['session_strip'].unique()
return self.subjects
[docs]
def get_sessions(self, idx) -> np.ndarray:
"""
Find all probes for a given session.
Note if multi-shank data it will return probe00 rather than probe00a, the individual shank
is chosen using the shank dropdown.
Parameters
----------
idx : idx
The index of the chosen subject
Returns
-------
np.ndarray
All probes with spikesorting data for the chosen session
"""
self.session_df = self.df.loc[self.df['session_strip'] == self.subjects[idx]]
self.sessions = np.unique(
[self.normalize_shank_label(pr) for pr in self.session_df['probe'].values]
)
return self.sessions
[docs]
def get_shanks(self, idx: int) -> np.ndarray:
"""
Find all shanks for a given probe and initialise the loaders.
Parameters
----------
idx : idx
The index of the chosen probe
Returns
-------
np.ndarray
All shanks for the chosen probe
"""
shank = self.sessions[idx]
self.shank_df = self.session_df.loc[
self.session_df['probe'].str.contains(shank)
].sort_values('probe')
self.initialise_shanks()
self.shank_labels = self.shank_df['probe'].unique()
return self.shank_labels
[docs]
def set_info(self, idx: int) -> None:
"""
Set the information about the selected shank.
Parameters
----------
idx: int
The index of the selected shank
"""
self.selected_shank = self.shank_labels[idx]
self.selected_idx = idx
[docs]
def download_histology(self) -> SliceLoader:
"""Download and load in the histology slice data."""
_, hist_path = download_histology_data(self.subj, self.lab)
return make_slice_loader(hist_path, self.brain_atlas, 'ccf')
[docs]
def initialise_shanks(self) -> None:
"""Initialise each shank and config with the selected loaders."""
self.shanks = defaultdict(Bunch)
user = params.get().ALYX_LOGIN
for _, shank in self.shank_df.iterrows():
loaders = Bunch()
local_path = self.root_path.joinpath(shank.local_path)
data_paths = DatasetPaths(
spike_sorting=local_path.joinpath(shank.spike_collection or ''),
processed_ephys=local_path.joinpath(shank.ephys_collection or ''),
raw_ephys=local_path.joinpath(shank.meta_collection or ''),
task=local_path.joinpath(shank.task_collection or ''),
raw_task=local_path.joinpath(shank.raw_task_collection or ''),
)
ins = self.get_insertion(shank)
xyz_picks = ins['json'].get('xyz_picks', None)
xyz_picks = np.array(xyz_picks) / 1e6 if xyz_picks is not None else None
if shank.is_quarter: # Quarter is offline
loaders['data'] = DataLoaderLocal(data_paths)
loaders['geom'] = GeometryLoaderLocal(data_paths)
loaders['align'] = AlignmentLoaderLocal(
data_paths.spike_sorting, 0, 1, user=user, xyz_picks=xyz_picks
)
loaders['upload'] = AlignmentUploaderLocal(
data_paths.spike_sorting, 0, 1, self.brain_atlas, user=user
)
loaders['ephys'] = SpikeGLXLoaderLocal(data_paths.raw_ephys)
loaders['plots'] = PlotLoader()
self.shanks[shank.probe]['quarter'] = ShankHandler(loaders, 0)
else: # Dense is online
# If we don't have the data locally we download it
if data_paths.spike_sorting == local_path:
loaders['data'] = DataLoaderOne(ins, self.one)
loaders['geom'] = GeometryLoaderOne(
ins, self.one, probe_collection=loaders['data'].probe_collection
)
# Otherwise we load from local
else:
loaders['data'] = DataLoaderLocal(data_paths)
loaders['geom'] = GeometryLoaderLocal(data_paths)
# Work in progress alignments are saved next to the spike sorting data
spike_path = loaders['data'].spike_sorting_path
loaders['align'] = AlignmentLoaderOne(
ins, self.one, user=user, data_path=spike_path
)
loaders['upload'] = AlignmentUploaderOne(
ins, self.one, self.brain_atlas, data_path=spike_path
)
loaders['ephys'] = SpikeGLXLoaderOne(ins, self.one)
if EPHYS_ATLAS:
loaders['features'] = FeatureLoaderOne(
ins, self.one, self.ea_model, multi_area=True
)
loaders['plots'] = PlotLoader()
self.shanks[shank.probe]['dense'] = ShankHandler(loaders, 0)
# The dense configuration is the default one and is read from Alyx, so its alignments take
# precedence over the local ones of the quarter configuration
self.sync_config_alignments()
self.subj = shank['subject']
self.lab = shank['lab']
[docs]
def get_insertion(self, shank: pd.Series) -> dict:
"""Get the alyx probe insertion for the shank."""
ins = self.one.alyx.rest('insertions', 'list', id=shank.pid, expires=timedelta(days=1))
return ins[0]
[docs]
class ProbeHandlerLocal(ProbeHandler):
"""
Local file system implementation of ProbeHandler.
For this ProbeHandler, all ephys and alignment data must be stored in a single folder on disk.
"""
def __init__(self, brain_atlas: BrainAtlas | None = None):
super().__init__(brain_atlas)
[docs]
def get_shanks(self, folder_path: Path) -> list[str]:
"""
Find the number of shanks on the probes.
Loads the channels or ap meta data from the folder path and initialises the loaders
for each shank.
Parameters
----------
folder_path : Path
A path to the folder on the local disk that contains the data
"""
self.data_paths = DatasetPaths(
spike_sorting=folder_path,
processed_ephys=folder_path,
raw_ephys=folder_path,
histology=folder_path,
picks=folder_path,
output=folder_path,
)
# Load in the geometry and find the number of shnaks
self.geom = GeometryLoaderLocal(self.data_paths)
self.geom.get_geometry()
self.n_shanks = self.geom.channels.n_shanks
if self.n_shanks == 1:
self.shank_labels = ['shank_1']
else:
self.shank_labels = [f'shank_{iShank + 1}' for iShank in range(self.n_shanks)]
self.initialise_shanks()
return self.shank_labels
[docs]
def set_info(self, idx: int) -> None:
"""
Set the information about the selected shank.
Parameters
----------
idx: int
The index of the selected shank
"""
self.selected_shank = self.shank_labels[idx]
self.selected_idx = idx
[docs]
def download_histology(self) -> SliceLoader:
"""Load in the histology slice data."""
return make_slice_loader(
self.data_paths.histology, self.brain_atlas, self.data_paths.histology_space
)
[docs]
def initialise_shanks(self) -> None:
"""Initialise each shank with the loaders."""
self.shanks = defaultdict(Bunch)
for ish, ishank in enumerate(self.shank_labels):
loaders = Bunch()
loaders['geom'] = self.geom
loaders['data'] = DataLoaderLocal(self.data_paths)
loaders['align'] = AlignmentLoaderLocal(
self.data_paths.output, ish, self.n_shanks, picks_path=self.data_paths.picks
)
loaders['upload'] = AlignmentUploaderLocal(
self.data_paths.output, ish, self.n_shanks, self.brain_atlas
)
if self.data_paths.raw_ephys is not None:
loaders['ephys'] = SpikeGLXLoaderLocal(self.data_paths.raw_ephys)
loaders['plots'] = PlotLoader()
self.shanks[ishank][self.default_config] = ShankHandler(loaders, ish)
[docs]
class ProbeHandlerLocalYaml(ProbeHandler):
"""
Local file system ProbeHandler driven by a session yaml file.
The yaml (see :func:`ibl_alignment_gui.utils.parse_yaml.load_alignment_yaml`) specifies, per
probe/config, where each dataset lives (spike sorting, raw/processed ephys, picks, histology,
output, and optional per-channel features). The resolved ``DatasetPaths`` for each probe/config
are wired directly into the local loaders (``DataLoaderLocal``, ``GeometryLoaderLocal`` etc.).
Parameters
----------
yaml_file : str or Path
Path to the session yaml configuration file.
brain_atlas : AllenAtlas or None
An AllenAtlas instance (created if None).
"""
def __init__(self, yaml_file: str | Path, brain_atlas: BrainAtlas | None = None):
self.configs, self.probes, self.data_paths, self.histology_space = load_alignment_yaml(
yaml_file
)
# The atlas (anatomical or Allen, see _make_atlas) is built lazily by build_atlas() on the
# background loading thread rather than here, so it does not block GUI construction.
super().__init__(brain_atlas)
# The base sets a single 'default' config; mirror it to the yaml config name and, when the
# yaml carries two configs, expose both (plus 'both') as in the multi-config workflows.
self.default_config = self.configs[0]
if len(self.configs) > 1:
self.non_default_config = self.configs[1]
self.possible_configs = self.configs + ['both']
else:
self.possible_configs = [self.default_config]
self.selected_config = self.default_config
[docs]
def get_shanks(self, _) -> list[str]:
"""
Determine the shanks from the yaml and initialise the loaders.
If a single probe is specified we load its geometry to detect whether it is a multi-shank
recording. Otherwise each probe entry in the yaml is treated as an individual shank.
Parameters
----------
_ : Any
Ignored — the yaml path was supplied at construction time. The signature matches the
other ProbeHandlers so the controller can call it uniformly.
"""
if len(self.probes) == 1:
data_path = self.data_paths[self.default_config][self.probes[0]]
geom = GeometryLoaderLocal(data_path)
geom.get_geometry()
# Shank count comes from the ALF channels object when present, else from the SpikeGLX
# meta (e.g. external datasets with no spike sorting), mirroring the fallback used in
# GeometryLoader.get_sites_for_shank.
sites = geom.channels if geom.channels is not None else geom.electrodes
self.n_shanks = sites.n_shanks
if self.n_shanks == 1:
self.shank_labels = list(self.probes)
else:
self.shank_labels = [f'shank_{ishank + 1}' for ishank in range(self.n_shanks)]
else:
self.shank_labels = list(self.probes)
self.n_shanks = 1
self.initialise_shanks()
return self.shank_labels
[docs]
def set_info(self, idx: int) -> None:
"""
Set the information about the selected shank.
Parameters
----------
idx : int
The index of the selected shank.
"""
self.selected_shank = self.shank_labels[idx]
self.selected_idx = idx
[docs]
def download_histology(self) -> SliceLoader:
"""Load in the histology slice data."""
data_paths = self.data_paths[self.selected_config][self.shank_labels[0]]
return make_slice_loader(data_paths.histology, self.brain_atlas, self.histology_space)
[docs]
def initialise_shanks(self) -> None:
"""Initialise each shank and config with loaders pointing at the resolved yaml paths."""
self.shanks = defaultdict(Bunch)
# A single probe entry may still be multi-shank; in that case all shanks share that one
# probe's dataset paths and are told apart by their shank index (the geometry is split per
# shank inside ShankHandler.load_data via get_sites_for_shank).
single_probe = len(self.probes) == 1
for ish, shank in enumerate(self.shank_labels):
ishank = ish if self.n_shanks > 1 else 0
probe = self.probes[0] if single_probe else shank
for config in self.configs:
data_path = self.data_paths[config][probe]
loaders = Bunch()
loaders['geom'] = GeometryLoaderLocal(data_path)
loaders['data'] = DataLoaderLocal(data_path)
loaders['align'] = self._build_align_loader(data_path, ishank)
loaders['upload'] = self._build_upload_loader(data_path, ishank)
loaders['ephys'] = SpikeGLXLoaderLocal(data_path.raw_ephys)
# Per-session features (if the yaml specifies them) load via the existing
# shank_handler.load_data -> loaders['features'] path, so the session is
# self-contained and switching yaml switches the features too.
# TODO pass in geometry
if data_path.features is not None:
loaders['features'] = FeatureLoaderLocal(data_path.features)
loaders['plots'] = PlotLoader()
self.shanks[shank][config] = ShankHandler(loaders, ishank)
# Each configuration has loaded the alignments from its own paths, so share them to keep
# the configurations of each shank in step
self.sync_config_alignments()
def _build_align_loader(self, data_path: DatasetPaths, ishank: int) -> AlignmentLoaderLocal:
"""
Build the alignment loader for a shank.
Reads xyz picks and previous alignments from the local file system;
:class:`ProbeHandlerAllenYaml` overrides it to use the DocDB backend.
The previous alignments are read from the output path, which is where the uploader writes
them to, while the xyz picks are read from the picks path.
Parameters
----------
data_path : DatasetPaths
The resolved dataset paths for the probe/config.
ishank : int
Index of the shank (0-based).
Returns
-------
AlignmentLoaderLocal
The alignment loader for the shank.
"""
return AlignmentLoaderLocal(
data_path.output,
ishank,
self.n_shanks,
histology_space=self.histology_space,
picks_path=data_path.picks or data_path.spike_sorting,
)
def _build_upload_loader(self, data_path: DatasetPaths, ishank: int) -> AlignmentUploaderLocal:
"""
Build the alignment uploader for a shank.
Writes channel locations and alignments to the local file system;
:class:`ProbeHandlerAllenYaml` overrides it to additionally post the results to DocDB.
Parameters
----------
data_path : DatasetPaths
The resolved dataset paths for the probe/config.
ishank : int
Index of the shank (0-based).
Returns
-------
AlignmentUploaderLocal
The alignment uploader for the shank.
"""
return AlignmentUploaderLocal(
data_path.output,
ishank,
self.n_shanks,
self.brain_atlas,
)
[docs]
class ProbeHandlerAllenYaml(ProbeHandlerLocalYaml):
"""
Probe handler for the Allen/Code Ocean (anatomical) workflow with DocDB support.
Extends :class:`ProbeHandlerLocalYaml` (so all the yaml/anatomical/data/geometry/histology/
transform wiring is reused) and, mirroring how :class:`ProbeHandlerONE` owns a ``one``
instance, owns a :class:`~ibl_alignment_gui.backends.allen.docdb_api.DocDB` instance that is
injected into the DocDB alignment loader and uploader (overriding the local factory hooks
:meth:`ProbeHandler._build_align_loader` / :meth:`ProbeHandler._build_upload_loader`).
The ``use_docdb`` flag selects the alignment backend: when True the DocDB-backed loader and
uploader are used (previous alignments read from DocDB with a local fallback; results written
locally and posted to DocDB); when False the plain local variants are used. It can be flipped
at runtime with :meth:`set_use_docdb` (e.g. from the DocDB checkbox).
Parameters
----------
yaml_file : str or Path
Path to the session yaml configuration file.
brain_atlas : BrainAtlas or None
A pre-built brain atlas. If None, it is built lazily (anatomical or Allen, per the yaml).
docdb : DocDB or None
The DocDB client to inject. A default :class:`DocDB` is created if None.
use_docdb : bool
Whether to use the DocDB alignment backend (True) or the local one (False).
"""
def __init__(
self,
yaml_file: str | Path,
brain_atlas: BrainAtlas | None = None,
docdb: DocDB | None = None,
use_docdb: bool = True,
):
# Imported lazily so the base install (offline / IBL modes) does not require the allen
# extra; the alignment-gui-allen launcher checks the extra is installed up front.
from ibl_alignment_gui.backends.allen.docdb_api import DocDB # noqa: PLC0415
self.docdb: DocDB = docdb or DocDB()
self.use_docdb: bool = use_docdb
super().__init__(yaml_file, brain_atlas=brain_atlas)
def _make_atlas(self) -> BrainAtlas:
"""Return the appropriate atlas based on the histology space in the YAML config."""
histology_path = self.data_paths[self.selected_config][self.shank_labels[0]].histology
if self.histology_space == 'anatomical' and histology_path:
return build_anatomical_atlas(histology_path)
return AllenAtlas()
def _build_align_loader(self, data_path: DatasetPaths, ishank: int) -> AlignmentLoaderDocDB:
"""Build a DocDB alignment loader (falling back to local when ``use_docdb`` is False)."""
# The output path is used, as the docdb record and the local prev_alignments file are both
# written relative to it by the uploader
return AlignmentLoaderDocDB(
data_path.output,
ishank,
self.n_shanks,
self.docdb,
use_db=self.use_docdb,
histology_space=self.histology_space,
picks_path=data_path.picks or data_path.spike_sorting,
)
def _build_upload_loader(self, data_path: DatasetPaths, ishank: int) -> AlignmentUploaderDocDB:
"""Build a DocDB alignment uploader (falling back to local when ``use_docdb`` is False)."""
return AlignmentUploaderDocDB(
data_path.output,
ishank,
self.n_shanks,
self.brain_atlas,
self.docdb,
transform_loader=self._build_transform_loader(data_path),
use_db=self.use_docdb,
)
@staticmethod
def _build_transform_loader(data_path: DatasetPaths) -> TransformLoaderAllen | None:
"""Return a SmartSPIM -> CCF transform loader for the probe, or None if not configured."""
return (
TransformLoaderAllen(data_path.transforms)
if data_path.transforms is not None
else None
)
[docs]
def set_use_docdb(self, use_docdb: bool) -> None:
"""
Switch the alignment backend and refresh previous alignments for every shank.
Flips the ``use_db`` flag on each shank's existing DocDB alignment loader and uploader
(leaving the loaded ephys/geometry/histology untouched) and re-reads the previous
alignments so the alignment dropdown reflects the new source.
Parameters
----------
use_docdb : bool
Whether to use the DocDB alignment backend (True) or the local one (False).
"""
self.use_docdb = use_docdb
for shank in self.shanks:
for config in self.configs:
handler = self.shanks[shank][config]
handler.loaders['align'].use_db = use_docdb
handler.loaders['upload'].use_db = use_docdb
handler.loaders['align'].load_previous_alignments()
# Each configuration has re-read the alignments from its own source, so share them again
self.sync_config_alignments()