diff --git a/src/extra/components/__init__.py b/src/extra/components/__init__.py index 0989e485..d69d881d 100644 --- a/src/extra/components/__init__.py +++ b/src/extra/components/__init__.py @@ -1,4 +1,6 @@ +from .utils import TrainData, align + from .scantool import Scantool # noqa from .pulses import XrayPulses, OpticalLaserPulses, MachinePulses, \ PumpProbePulses, DldPulses # noqa diff --git a/src/extra/components/pulses.py b/src/extra/components/pulses.py index 0bacb958..f8f2c09c 100644 --- a/src/extra/components/pulses.py +++ b/src/extra/components/pulses.py @@ -194,6 +194,12 @@ def select_trains(self, trains): return res + def data_counts(self, labelled=True): + if self._key is not None: + return self._key.data_counts(labelled=labelled) + + raise NotImplementedError('data_counts') + def pulse_ids(self, labelled=True, copy=True): """Get pulse IDs. diff --git a/src/extra/components/utils.py b/src/extra/components/utils.py index 5a8667f7..5a84ca65 100644 --- a/src/extra/components/utils.py +++ b/src/extra/components/utils.py @@ -1,3 +1,4 @@ + from ..utils.misc import _isinstance_no_import # Source prefixes in use at each SASE. @@ -75,3 +76,35 @@ def _select_subcomponent_trains(src, keys, dst=None): setattr(dst, key, prop.select_trains(trains)) return dst + + +from collections.abc import Sequence +from typing import Protocol, Self, runtime_checkable +import pandas as pd +from extra_data import by_id + + +@runtime_checkable +class TrainData(Protocol): + train_ids: Sequence[int] + + def select_trains(self, train_sel) -> Self: + raise NotImplementedError('select_trains') + + def data_counts(self, labelled=True) -> pd.Series: + raise NotImplementedError('data_counts') + + +def align(*unaligned): + counts = pd.concat([ + obj.data_counts() for obj in unaligned + if isinstance(obj, TrainData) + ], axis=1, join='inner') + + if len(counts.columns) != len(unaligned): + raise TypeError('one or more passed objects do not implement ' + 'the TrainData protocol') + + train_ids = counts.index[(counts > 0).all(axis=1)].to_numpy() + + return [obj.select_trains(by_id[train_ids]) for obj in unaligned]