Skip to content
Draft
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
2 changes: 2 additions & 0 deletions src/extra/components/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@

from .utils import TrainData, align

from .scantool import Scantool # noqa
from .pulses import XrayPulses, OpticalLaserPulses, MachinePulses, \
PumpProbePulses, DldPulses # noqa
Expand Down
6 changes: 6 additions & 0 deletions src/extra/components/pulses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
33 changes: 33 additions & 0 deletions src/extra/components/utils.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@

from ..utils.misc import _isinstance_no_import

# Source prefixes in use at each SASE.
Expand Down Expand Up @@ -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]
Loading